{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":87793,"databundleVersionId":12276181,"sourceType":"competition"}],"dockerImageVersionId":30919,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport math\nimport random\nfrom typing import List, Tuple, Optional\nfrom collections import defaultdict\n\nimport numpy as np\nimport pandas as pd\nfrom tqdm import tqdm\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.optim import AdamW\nfrom torch.utils.tensorboard import SummaryWriter","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-17T05:42:59.731437Z","iopub.execute_input":"2025-11-17T05:42:59.731771Z","iopub.status.idle":"2025-11-17T05:42:59.736897Z","shell.execute_reply.started":"2025-11-17T05:42:59.731747Z","shell.execute_reply":"2025-11-17T05:42:59.735960Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# ------------------------- Reproducibility -------------------------\ndef set_seed(seed=42):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\nset_seed(42)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# ------------------------- Utilities -------------------------\nVOCAB = {\"A\":1, \"C\":2, \"G\":3, \"U\":4, \"N\":5}  # 0 reserved for PAD\n\n\ndef seq_to_ids(seq: str, max_len: int) -> List[int]:\n    ids = [VOCAB.get(ch, VOCAB['N']) for ch in seq.upper()][:max_len]\n    if len(ids) < max_len:\n        ids += [0] * (max_len - len(ids))\n    return ids\n\n# dot-bracket to contact map (kept for completeness)\ndef dotbracket_to_contacts(dot: str) -> np.ndarray:\n    stack = []\n    n = len(dot)\n    mat = np.zeros((n, n), dtype=np.uint8)\n    pairs = {')': '(', ']': '[', '}': '{'}\n    openers = set(['(', '[', '{', '<'])\n    closers = set([')', ']', '}', '>'])\n    for i, ch in enumerate(dot):\n        if ch in openers:\n            stack.append((ch, i))\n        elif ch in closers:\n            if not stack:\n                continue\n            top_ch, top_i = stack.pop()\n            expected = pairs.get(ch, '(')\n            if top_ch != expected:\n                pass\n            mat[top_i, i] = 1\n            mat[i, top_i] = 1\n    return mat\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# ------------------------- Dataset -------------------------\nclass RNADataset(Dataset):\n    \"\"\"Dataset built from a sequences dataframe. The contact maps are attached\n    later via a custom collate function that uses precomputed contacts_by_id.\n    Expects seq_df containing column 'target_id' and 'sequence'.\"\"\"\n    def __init__(self, seq_df: pd.DataFrame, max_len: int = 512):\n        self.df = seq_df.reset_index(drop=True)\n        self.max_len = max_len\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        seq = str(row['sequence']).strip()\n        seq_ids = seq_to_ids(seq, self.max_len)\n        seq_len = min(len(seq), self.max_len)\n        return {\n            'id': row['target_id'],\n            'seq_ids': torch.LongTensor(seq_ids),\n            'seq_len': seq_len\n        }\n\n# collate that pads and attaches contacts_by_id provided externally\ndef collate_with_contacts(batch, contacts_map, max_len):\n    ids = [b['id'] for b in batch]\n    seq_ids = torch.stack([b['seq_ids'] for b in batch])\n    seq_lens = torch.tensor([b['seq_len'] for b in batch], dtype=torch.long)\n    B, L = seq_ids.shape\n    contacts = torch.zeros((B, L, L), dtype=torch.float32)\n    for i, rid in enumerate(ids):\n        key = str(rid)\n        if key in contacts_map:\n            mat = contacts_map[key]\n            # ensure shape matches L\n            if mat.shape[0] >= L:\n                contacts[i] = torch.from_numpy(mat[:L, :L]).float()\n            else:\n                tmp = np.zeros((L, L), dtype=np.uint8)\n                tmp[:mat.shape[0], :mat.shape[1]] = mat\n                contacts[i] = torch.from_numpy(tmp).float()\n        else:\n            # fallback to zeros\n            contacts[i] = torch.zeros((L, L), dtype=torch.float32)\n    return {'id': ids, 'seq_ids': seq_ids, 'seq_lens': seq_lens, 'contact': contacts}\n\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ------------------------- Model -------------------------\nclass PositionalEncoding(nn.Module):\n    def __init__(self, d_model: int, max_len: int = 512):\n        super().__init__()\n        pe = torch.zeros(max_len, d_model)\n        position = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1)\n        div_term = torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model))\n        pe[:, 0::2] = torch.sin(position * div_term)\n        pe[:, 1::2] = torch.cos(position * div_term)\n        self.register_buffer('pe', pe)\n\n    def forward(self, x):\n        L = x.size(1)\n        return x + self.pe[:L].unsqueeze(0)\n\nclass Encoder1D(nn.Module):\n    def __init__(self, vocab_size=6, emb_dim=128, n_heads=8, ff_dim=256, n_layers=3, max_len=512, dropout=0.1):\n        super().__init__()\n        self.embedding = nn.Embedding(vocab_size, emb_dim, padding_idx=0)\n        self.pos_enc = PositionalEncoding(emb_dim, max_len=max_len)\n        encoder_layer = nn.TransformerEncoderLayer(d_model=emb_dim, nhead=n_heads, dim_feedforward=ff_dim, dropout=dropout, batch_first=True)\n        self.transformer = nn.TransformerEncoder(encoder_layer, num_layers=n_layers)\n        self.layer_norm = nn.LayerNorm(emb_dim)\n\n    def forward(self, seq_ids, src_key_padding_mask=None):\n        x = self.embedding(seq_ids) * math.sqrt(self.embedding.embedding_dim)\n        x = self.pos_enc(x)\n        x = self.transformer(x, src_key_padding_mask=src_key_padding_mask)\n        x = self.layer_norm(x)\n        return x\n\nclass Pairwise2DUNet(nn.Module):\n    def __init__(self, in_channels=256, base_channels=64):\n        super().__init__()\n        self.conv1 = nn.Sequential(nn.Conv2d(in_channels, base_channels, 3, padding=1), nn.BatchNorm2d(base_channels), nn.ReLU(inplace=True))\n        self.conv2 = nn.Sequential(nn.Conv2d(base_channels, base_channels*2, 3, padding=1), nn.BatchNorm2d(base_channels*2), nn.ReLU(inplace=True))\n        self.conv3 = nn.Sequential(nn.Conv2d(base_channels*2, base_channels*4, 3, padding=1), nn.BatchNorm2d(base_channels*4), nn.ReLU(inplace=True))\n        self.up2 = nn.ConvTranspose2d(base_channels*4, base_channels*2, 2, stride=2)\n        self.dec2 = nn.Sequential(nn.Conv2d(base_channels*4, base_channels*2, 3, padding=1), nn.BatchNorm2d(base_channels*2), nn.ReLU(inplace=True))\n        self.up1 = nn.ConvTranspose2d(base_channels*2, base_channels, 2, stride=2)\n        self.dec1 = nn.Sequential(nn.Conv2d(base_channels*2, base_channels, 3, padding=1), nn.BatchNorm2d(base_channels), nn.ReLU(inplace=True))\n        self.out_conv = nn.Conv2d(base_channels, 1, 1)\n\n    def forward(self, x):\n        c1 = self.conv1(x)\n        p1 = F.max_pool2d(c1, 2)\n        c2 = self.conv2(p1)\n        p2 = F.max_pool2d(c2, 2)\n        c3 = self.conv3(p2)\n        u2 = self.up2(c3)\n        if u2.size() != c2.size():\n            u2 = F.interpolate(u2, size=c2.shape[2:], mode='bilinear', align_corners=False)\n        d2 = self.dec2(torch.cat([u2, c2], dim=1))\n        u1 = self.up1(d2)\n        if u1.size() != c1.size():\n            u1 = F.interpolate(u1, size=c1.shape[2:], mode='bilinear', align_corners=False)\n        d1 = self.dec1(torch.cat([u1, c1], dim=1))\n        out = self.out_conv(d1)\n        return out.squeeze(1)\n\nclass RNAContactPredictor(nn.Module):\n    def __init__(self, enc: Encoder1D, pairwise_channels=256):\n        super().__init__()\n        self.enc = enc\n        emb_dim = enc.embedding.embedding_dim\n        self.proj = nn.Linear(emb_dim, pairwise_channels // 2)\n        self.pair_net = Pairwise2DUNet(in_channels=pairwise_channels, base_channels=32)\n\n    def forward(self, seq_ids, seq_lens=None):\n        src_key_padding_mask = (seq_ids == 0)\n        h = self.enc(seq_ids, src_key_padding_mask=src_key_padding_mask)\n        p = self.proj(h)\n        B, L, C = p.shape\n        a = p.unsqueeze(2)\n        b = p.unsqueeze(1)\n        pair = torch.cat([a.expand(-1, -1, L, -1), b.expand(-1, L, -1, -1)], dim=-1)\n        pair = pair.permute(0, 3, 1, 2).contiguous()\n        logits = self.pair_net(pair)\n        return logits\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# ------------------------- Contact building from coordinates -------------------------\nTHRESH = 8.0  # angstroms\nMAX_COORD_PER_ROW = 100  # safe cap for flattened rows\n\n\ndef parse_struct_and_pos_from_id(id_str):\n    parts = str(id_str).rsplit('_', maxsplit=1)\n    if len(parts) == 2 and parts[1].isdigit():\n        return parts[0], int(parts[1])\n    return str(id_str), None\n\n\ndef build_contacts_from_labels(labels_df: pd.DataFrame, seq_df: Optional[pd.DataFrame] = None, max_len: int = 200, threshold: float = THRESH):\n    contacts_by_id = {}\n    coords_by_id = {}\n    cols = labels_df.columns.tolist()\n\n    # detect flattened-per-structure format (many x_k columns)\n    x_columns = sorted([c for c in cols if c.lower().startswith('x_')], key=lambda s: int(s.split('_')[1]) if '_' in s and s.split('_')[1].isdigit() else 0)\n\n    if len(x_columns) > 1 and labels_df.shape[0] > 0 and labels_df['ID'].nunique() == labels_df.shape[0]:\n        # one row per structure, columns x_1..x_n\n        for _, row in labels_df.iterrows():\n            sid = row['ID']\n            coords = []\n            k = 1\n            while True:\n                xk = f'x_{k}'\n                yk = f'y_{k}'\n                zk = f'z_{k}'\n                if xk in labels_df.columns and yk in labels_df.columns and zk in labels_df.columns:\n                    x = row[xk]; y = row[yk]; z = row[zk]\n                    if pd.isna(x) or pd.isna(y) or pd.isna(z):\n                        break\n                    coords.append((float(x), float(y), float(z)))\n                    k += 1\n                else:\n                    break\n            coords_by_id[str(sid)] = coords\n    else:\n        # one row per residue\n        temp = defaultdict(dict)\n        for _, row in labels_df.iterrows():\n            sid_full = row['ID']\n            struct_id, pos = parse_struct_and_pos_from_id(sid_full)\n            if 'resid' in labels_df.columns:\n                try:\n                    pos = int(row['resid'])\n                except:\n                    pass\n            # determine coordinate columns\n            if 'x_1' in labels_df.columns:\n                x = row['x_1']; y = row['y_1']; z = row['z_1']\n            elif 'x' in labels_df.columns and 'y' in labels_df.columns and 'z' in labels_df.columns:\n                x = row['x']; y = row['y']; z = row['z']\n            else:\n                # try to find first triplet x_k,y_k,z_k\n                found = False\n                for c in labels_df.columns:\n                    if c.lower().startswith('x_') and '_' in c:\n                        base = c.split('_',1)[1]\n                        xc, yc, zc = f'x_{base}', f'y_{base}', f'z_{base}'\n                        if xc in labels_df.columns and yc in labels_df.columns and zc in labels_df.columns:\n                            x = row[xc]; y = row[yc]; z = row[zc]\n                            found = True\n                            break\n                if not found:\n                    continue\n            if pd.isna(x) or pd.isna(y) or pd.isna(z):\n                continue\n            if pos is None:\n                pos = len(temp[struct_id]) + 1\n            temp[struct_id][int(pos)] = (float(x), float(y), float(z))\n\n        for struct_id, d in temp.items():\n            ordered = [d[k] for k in sorted(d.keys())]\n            coords_by_id[str(struct_id)] = ordered\n\n    # compute contact maps\n    for sid, coords in coords_by_id.items():\n        L = min(len(coords), max_len)\n        if L == 0:\n            contacts_by_id[sid] = np.zeros((max_len, max_len), dtype=np.uint8)\n            continue\n        arr = np.array(coords[:L])\n        dists = np.sqrt(np.sum((arr[:, None, :] - arr[None, :, :])**2, axis=-1))\n        contact = (dists <= threshold).astype(np.uint8)\n        np.fill_diagonal(contact, 0)\n        if contact.shape[0] < max_len:\n            pad = np.zeros((max_len, max_len), dtype=np.uint8)\n            pad[:L, :L] = contact\n            contact = pad\n        contacts_by_id[sid] = contact\n    return contacts_by_id, coords_by_id\n\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ------------------------- Metrics & Loss -------------------------\n\ndef contact_metrics(pred_logits: torch.Tensor, target: torch.Tensor, seq_lens: torch.Tensor, threshold: float = 0.5):\n    preds = (torch.sigmoid(pred_logits) > threshold).int()\n    t = target.int()\n    B, L, _ = preds.shape\n    results = {'tp':0, 'fp':0, 'fn':0}\n    for i in range(B):\n        n = int(seq_lens[i].item())\n        p = preds[i, :n, :n]\n        g = t[i, :n, :n]\n        diag = torch.eye(n, dtype=torch.bool, device=preds.device)\n        p = p & (~diag)\n        g = g & (~diag)\n        # upper triangle mask\n        mask = torch.triu(torch.ones((n,n), dtype=torch.bool, device=preds.device), diagonal=1)\n        p_u = p[mask]\n        g_u = g[mask]\n        tp = int(((p_u==1) & (g_u==1)).sum().item())\n        fp = int(((p_u==1) & (g_u==0)).sum().item())\n        fn = int(((p_u==0) & (g_u==1)).sum().item())\n        results['tp'] += tp\n        results['fp'] += fp\n        results['fn'] += fn\n    tp = results['tp']; fp = results['fp']; fn = results['fn']\n    precision = tp / (tp + fp) if tp + fp > 0 else 0.0\n    recall = tp / (tp + fn) if tp + fn > 0 else 0.0\n    f1 = 2*precision*recall/(precision+recall) if (precision+recall) > 0 else 0.0\n    return {'precision': precision, 'recall': recall, 'f1': f1, 'tp':tp, 'fp':fp, 'fn':fn}\n\n\ndef masked_bce_loss_logits(logits, targets, seq_lens, pos_weight_tensor=None):\n    B, L, _ = logits.shape\n    losses = []\n    for i in range(B):\n        n = int(seq_lens[i].item())\n        log = logits[i, :n, :n]\n        tar = targets[i, :n, :n]\n        mask = torch.triu(torch.ones((n, n), dtype=torch.bool, device=logits.device), diagonal=1)\n        log_u = log[mask]\n        tar_u = tar[mask]\n        if pos_weight_tensor is not None:\n            loss = F.binary_cross_entropy_with_logits(log_u, tar_u, pos_weight=pos_weight_tensor)\n        else:\n            loss = F.binary_cross_entropy_with_logits(log_u, tar_u)\n        losses.append(loss)\n    return torch.stack(losses).mean()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# ------------------------- Training Loop -------------------------\n\ndef train_one_epoch_comp(model, dataloader, optimizer, scaler, device, epoch, writer=None, accumulation_steps=1, pos_weight=None):\n    model.train()\n    total_loss = 0.0\n    pbar = tqdm(enumerate(dataloader), total=len(dataloader))\n    optimizer.zero_grad()\n    for step, batch in pbar:\n        seq_ids = batch['seq_ids'].to(device)\n        seq_lens = batch['seq_lens'].to(device)\n        contacts = batch['contact'].to(device)\n\n        with torch.cuda.amp.autocast(enabled=(scaler is not None)):\n            logits = model(seq_ids, seq_lens)\n            loss = masked_bce_loss_logits(logits, contacts, seq_lens, pos_weight_tensor=pos_weight)\n            loss = loss / accumulation_steps\n\n        if scaler is not None:\n            scaler.scale(loss).backward()\n            if (step + 1) % accumulation_steps == 0:\n                scaler.unscale_(optimizer)\n                torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n                scaler.step(optimizer)\n                scaler.update()\n                optimizer.zero_grad()\n        else:\n            loss.backward()\n            if (step + 1) % accumulation_steps == 0:\n                torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n                optimizer.step()\n                optimizer.zero_grad()\n\n        total_loss += loss.item() * accumulation_steps\n        pbar.set_description(f\"Epoch {epoch} loss: {total_loss/(step+1):.4f}\")\n\n    avg_loss = total_loss / len(dataloader)\n    if writer:\n        writer.add_scalar('train/loss', avg_loss, epoch)\n    return avg_loss\n\n\ndef validate_comp(model, dataloader, device, epoch, writer=None, pos_weight=None):\n    model.eval()\n    total_loss = 0.0\n    agg_tp = agg_fp = agg_fn = 0\n    with torch.no_grad():\n        for batch in tqdm(dataloader, desc='Valid'):\n            seq_ids = batch['seq_ids'].to(device)\n            seq_lens = batch['seq_lens'].to(device)\n            contacts = batch['contact'].to(device)\n            logits = model(seq_ids, seq_lens)\n            loss = masked_bce_loss_logits(logits, contacts, seq_lens, pos_weight_tensor=pos_weight)\n            total_loss += loss.item()\n            metrics = contact_metrics(logits, contacts, seq_lens)\n            agg_tp += metrics['tp']; agg_fp += metrics['fp']; agg_fn += metrics['fn']\n\n    precision = agg_tp / (agg_tp + agg_fp) if agg_tp + agg_fp > 0 else 0.0\n    recall = agg_tp / (agg_tp + agg_fn) if agg_tp + agg_fn > 0 else 0.0\n    f1 = 2*precision*recall/(precision+recall) if (precision+recall) > 0 else 0.0\n    avg_loss = total_loss / len(dataloader)\n    if writer:\n        writer.add_scalar('valid/loss', avg_loss, epoch)\n        writer.add_scalar('valid/precision', precision, epoch)\n        writer.add_scalar('valid/recall', recall, epoch)\n        writer.add_scalar('valid/f1', f1, epoch)\n    return {'loss': avg_loss, 'precision': precision, 'recall': recall, 'f1': f1}\n\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ------------------------- Main Script -------------------------\nif __name__ == '__main__':\n    # ------------------------- Competition-aware Main Script -------------------------\n    # Paths - adjust to your Kaggle dataset paths\n    MAX_LEN = 200  # safe default for Stanford RNA dataset\n    BATCH_SIZE = 8\n    N_EPOCHS = 20\n    LR = 3e-4\n    DEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\n    ACCUM_STEPS = 1\n    MIXED_PREC = True if DEVICE == 'cuda' else False\n\n    # competition files (adjust if mounted differently)\n    DATA_DIR = '/kaggle/input/stanford-rna-3d-folding'\n    TRAIN_SEQ = os.path.join(DATA_DIR, 'train_sequences.csv')\n    TRAIN_LABELS = os.path.join(DATA_DIR, 'train_labels.csv')\n    VAL_SEQ = os.path.join(DATA_DIR, 'validation_sequences.csv')\n    VAL_LABELS = os.path.join(DATA_DIR, 'validation_labels.csv')\n    TEST_SEQ = os.path.join(DATA_DIR, 'test_sequences.csv')\n    SAMPLE_SUB = os.path.join(DATA_DIR, 'sample_submission.csv')\n\n    OUT_DIR = 'checkpoints'\n    os.makedirs(OUT_DIR, exist_ok=True)\n\n    writer = SummaryWriter(log_dir=os.path.join(OUT_DIR, 'runs'))\n\n    # ------------------------- Load CSVs -------------------------\n    train_seq_df = pd.read_csv(TRAIN_SEQ)\n    train_lbl_df = pd.read_csv(TRAIN_LABELS)\n    val_seq_df = pd.read_csv(VAL_SEQ)\n    val_lbl_df = pd.read_csv(VAL_LABELS)\n\n    # Build contacts_by_id from coordinate labels\n    contacts_by_id, coords_by_id = build_contacts_from_labels(train_lbl_df, seq_df=train_seq_df, max_len=MAX_LEN, threshold=THRESH)\n    val_contacts_by_id, val_coords_by_id = build_contacts_from_labels(val_lbl_df, seq_df=val_seq_df, max_len=MAX_LEN, threshold=THRESH)\n    print('Built contacts for', len(contacts_by_id), 'train targets and', len(val_contacts_by_id), 'val targets')\n\n    # create dataset objects\n    train_ds = RNADataset(train_seq_df, max_len=MAX_LEN)\n    val_ds = RNADataset(val_seq_df, max_len=MAX_LEN)\n\n    # DataLoaders using our custom collate\n    train_loader = DataLoader(train_ds, batch_size=BATCH_SIZE, shuffle=True,\n                              collate_fn=lambda batch: collate_with_contacts(batch, contacts_by_id, MAX_LEN),\n                              num_workers=4, pin_memory=True)\n\n    val_loader = DataLoader(val_ds, batch_size=BATCH_SIZE, shuffle=False,\n                            collate_fn=lambda batch: collate_with_contacts(batch, val_contacts_by_id, MAX_LEN),\n                            num_workers=2, pin_memory=True)\n\n    # ------------------------- Model (small for competition) -------------------------\n    enc = Encoder1D(vocab_size=len(VOCAB)+1, emb_dim=64, n_heads=4, ff_dim=256, n_layers=2, max_len=MAX_LEN, dropout=0.1)\n    model = RNAContactPredictor(enc, pairwise_channels=128)\n    model = model.to(DEVICE)\n\n    optimizer = AdamW(model.parameters(), lr=LR, weight_decay=1e-5)\n    scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='min', factor=0.5, patience=3)\n    scaler = torch.cuda.amp.GradScaler() if MIXED_PREC else None\n\n    # Precompute pos_weight from training contacts\n    total_pos = 0\n    total_pairs = 0\n    for sid, mat in contacts_by_id.items():\n        # estimate n from matrix (count nonzero rows)\n        n = int(np.count_nonzero(np.sum(mat, axis=1)) + np.sum(np.any(mat,axis=0) | np.any(mat,axis=1)))\n        n = min(MAX_LEN, mat.shape[0])\n        tri = n * (n - 1) // 2\n        total_pairs += tri\n        total_pos += int(mat[:n, :n].sum() / 2)  # since symmetric\n    pos_weight = torch.tensor(((total_pairs - total_pos) / (total_pos + 1e-6)), dtype=torch.float32, device=DEVICE)\n    print(f\"Computed pos_weight={pos_weight.item():.4f} (pos={total_pos}, pairs={total_pairs})\")\n\n    # ------------------------- Training loop -------------------------\n    best_f1 = 0.0\n    for epoch in range(1, N_EPOCHS + 1):\n        train_loss = train_one_epoch_comp(model, train_loader, optimizer, scaler, DEVICE, epoch, writer, accumulation_steps=ACCUM_STEPS, pos_weight=pos_weight)\n        val_stats = validate_comp(model, val_loader, DEVICE, epoch, writer, pos_weight=pos_weight)\n        scheduler.step(val_stats['loss'])\n\n        print(f\"Epoch {epoch} -> train_loss: {train_loss:.4f}, val_loss: {val_stats['loss']:.4f}, f1: {val_stats['f1']:.4f}\")\n\n        # checkpoint best\n        if val_stats['f1'] > best_f1:\n            best_f1 = val_stats['f1']\n            ckpt_path = os.path.join(OUT_DIR, f'best_model_epoch{epoch}_f1{best_f1:.4f}.pt')\n            torch.save({'epoch': epoch, 'model_state': model.state_dict(), 'optimizer_state': optimizer.state_dict(), 'f1': best_f1}, ckpt_path)\n            print(f\"Saved best model to {ckpt_path}\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n    # ------------------------- Inference / submission -------------------------\n    # load best model if exists\n    ckpts = [p for p in os.listdir(OUT_DIR) if p.endswith('.pt')]\n    if ckpts:\n        latest = sorted(ckpts)[-1]\n        print('Loading checkpoint', latest)\n        ckpt = torch.load(os.path.join(OUT_DIR, latest), map_location=DEVICE)\n        model.load_state_dict(ckpt['model_state'])\n\n    # prepare test loader\n    test_seq_df = pd.read_csv(TEST_SEQ)\n    test_ds = RNADataset(test_seq_df, max_len=MAX_LEN)\n    test_loader = DataLoader(test_ds, batch_size=1, shuffle=False, collate_fn=lambda b: collate_with_contacts(b, {}, MAX_LEN), num_workers=1)\n\n    # produce submission rows incrementally to save memory\n    submission_rows = []\n    model.eval()\n    with torch.no_grad():\n        for batch in tqdm(test_loader, desc='Test'):\n            seq_ids = batch['seq_ids'].to(DEVICE)\n            seq_lens = batch['seq_lens'].to(DEVICE)\n            ids = batch['id']\n            logits = model(seq_ids, seq_lens)[0]  # [L,L]\n            prob = torch.sigmoid(logits)\n            n = int(seq_lens[0].item())\n            for i in range(n):\n                for j in range(i+1, n):\n                    submission_rows.append([ids[0], i, j, int((prob[i, j] > 0.5).item())])\n\n    submission_df = pd.DataFrame(submission_rows, columns=['id', 'position_i', 'position_j', 'is_paired'])\n    submission_df.to_csv('submission.csv', index=False)\n    print('Wrote submission.csv')\n\n    writer.close()\n\n# End of file\n","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}