{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":118765,"databundleVersionId":15231210,"sourceType":"competition"}],"dockerImageVersionId":31259,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# =============================================================================\n# STANFORD RNA 3D FOLDING: ULTIMATE FUSION PIPELINE\n# Fixes: OOM, Shape Mismatch, Scoring Errors, and Weight Loading\n# =============================================================================\n\nimport os\nimport gc\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom tqdm import tqdm\nimport warnings\n\nwarnings.filterwarnings(\"ignore\")\n\n# ==========================================\n# 1. CONFIGURATION (Safe Mode)\n# ==========================================\nclass Config:\n    TEST_CSV = '/kaggle/input/stanford-rna-3d-folding-2/test_sequences.csv'\n    MSA_DIR = '/kaggle/input/stanford-rna-3d-folding-2/MSA'\n    SUBMISSION_PATH = 'submission.csv'\n    \n    # Path to your weights (Update if different)\n    WEIGHTS_PATH = '/kaggle/input/ribonanza-weights/best_model.pth'\n    \n    DEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\n    BATCH_SIZE = 1        \n    \n    # --- MEMORY SAFETY SETTINGS ---\n    # 32 is the magic number to fit on P100 GPU without crashing\n    EMBED_DIM = 32       \n    NUM_LAYERS = 2 \n\nconfig = Config()\nprint(f\"✅ Config Loaded. Device: {config.DEVICE} | Embed Dim: {config.EMBED_DIM}\")\n\n# ==========================================\n# 2. ROBUST DATA PROCESSING\n# ==========================================\ndef parse_msa_to_pssm(msa_path, target_len):\n    \"\"\"Parses MSA safely. Returns uniform distribution if file missing/broken.\"\"\"\n    base_map = {'A':0, 'G':1, 'C':2, 'U':3, '-':4}\n    # Initialize with small value to prevent div/0\n    counts = np.ones((target_len, 5), dtype=np.float32) * 1e-6\n    \n    if not os.path.exists(msa_path):\n        return torch.tensor(counts / 5.0, dtype=torch.float32)\n\n    try:\n        with open(msa_path, 'r') as f: lines = f.readlines()\n        \n        # Filter valid sequences\n        seqs = [l.strip() for l in lines if not l.startswith('>') and len(l.strip()) == target_len]\n        \n        if len(seqs) == 0: return torch.tensor(counts / 5.0, dtype=torch.float32)\n             \n        # Process subset for speed\n        for seq in seqs[:50]: \n            for i, char in enumerate(seq):\n                if char in base_map: counts[i, base_map[char]] += 1.0\n                    \n        pssm = counts / counts.sum(axis=1, keepdims=True)\n        return torch.tensor(pssm, dtype=torch.float32)\n        \n    except Exception:\n        return torch.tensor(counts / 5.0, dtype=torch.float32)\n\nclass RNADataset(Dataset):\n    def __init__(self, csv_path):\n        if not os.path.exists(csv_path):\n            print(\"⚠️ Warning: File not found. Creating dummy data for testing.\")\n            self.df = pd.DataFrame({'target_id': ['R1107'], 'sequence': ['ACGU'*10]})\n        else:\n            self.df = pd.read_csv(csv_path)\n        self.mapper = {'A': 0, 'G': 1, 'C': 2, 'U': 3}\n\n    def __len__(self): return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        seq, seq_id = row['sequence'], row['target_id']\n        tokens = torch.tensor([self.mapper.get(c, 4) for c in seq], dtype=torch.long)\n        msa_file = os.path.join(config.MSA_DIR, f\"{seq_id}.MSA.fasta\")\n        pssm = parse_msa_to_pssm(msa_file, len(seq))\n        return {'id': seq_id, 'seq_str': seq, 'tokens': tokens, 'pssm': pssm}\n\n# ==========================================\n# 3. CRASH-PROOF MODEL ARCHITECTURE\n# ==========================================\nclass RelativePositionBias(nn.Module):\n    def __init__(self, num_bins=16, embed_dim=32):\n        super().__init__()\n        self.num_bins = num_bins\n        self.embedding = nn.Embedding(2 * num_bins + 1, embed_dim)\n    def forward(self, L, device):\n        range_vec = torch.arange(L, device=device)\n        rel_mat = range_vec[None, :] - range_vec[:, None]\n        rel_mat = torch.clamp(rel_mat, -self.num_bins, self.num_bins) + self.num_bins\n        return self.embedding(rel_mat).permute(2, 0, 1).unsqueeze(0)\n\nclass TriangleMultiplicativeUpdate(nn.Module):\n    def __init__(self, dim, hidden_dim=16):\n        super().__init__()\n        self.norm = nn.LayerNorm(dim)\n        self.proj_a = nn.Linear(dim, hidden_dim); self.proj_b = nn.Linear(dim, hidden_dim)\n        self.gate_a = nn.Linear(dim, hidden_dim); self.gate_b = nn.Linear(dim, hidden_dim)\n        self.proj_g = nn.Linear(dim, dim); self.gate_g = nn.Linear(dim, dim)\n        \n    def forward(self, x):\n        x = x.permute(0, 2, 3, 1) # (B, L, L, D)\n        x = self.norm(x)\n        a = self.proj_a(x * torch.sigmoid(self.gate_a(x)))\n        b = self.proj_b(x * torch.sigmoid(self.gate_b(x)))\n        triangle = torch.einsum('bikc,bjkc->bijc', a, b) # Memory Efficient Einsum\n        g = torch.sigmoid(self.gate_g(x)) * self.proj_g(x)\n        return (x + g * triangle).permute(0, 3, 1, 2)\n\n# [MODEL 1] 2D Structure Brain\n# [MODEL 1] 2D Structure Brain - FIXED\nclass RibonanzaUltra(nn.Module):\n    def __init__(self, embed_dim=32, num_layers=2):\n        super().__init__()\n        self.embed_dim = embed_dim\n        self.num_recycles = 1 \n        \n        self.embedding = nn.Embedding(5, embed_dim)\n        self.pssm_proj = nn.Linear(5, embed_dim)\n        self.fusion = nn.Linear(embed_dim * 2, embed_dim)\n        self.rel_pos = RelativePositionBias(num_bins=16, embed_dim=embed_dim//2)\n        \n        self.conv_stem = nn.Sequential(nn.Conv1d(embed_dim, embed_dim, 3, padding=1), nn.GELU(), nn.BatchNorm1d(embed_dim))\n        \n        self.pair_proj = nn.Linear(embed_dim, embed_dim//2)\n        self.triangle_block = TriangleMultiplicativeUpdate(embed_dim//2)\n        self.pair_act = nn.GELU()\n        self.pair_to_seq = nn.Linear(embed_dim//2, embed_dim)\n        \n        enc_layer = nn.TransformerEncoderLayer(d_model=embed_dim, nhead=4, dim_feedforward=embed_dim*2, batch_first=True, norm_first=True)\n        self.encoder = nn.TransformerEncoder(enc_layer, num_layers=num_layers)\n\n    def forward_step(self, x_tokens, x_pssm, prev_pair):\n        # 1. 1D Fusion\n        x_1d = self.fusion(torch.cat([self.embedding(x_tokens), self.pssm_proj(x_pssm)], dim=-1))\n        x_1d = self.conv_stem(x_1d.permute(0, 2, 1)).permute(0, 2, 1) + x_1d\n        \n        # 2. 2D Pair Creation\n        B, L, _ = x_1d.shape\n        x_pair = self.pair_proj(x_1d) # (B, L, D/2)\n        \n        # Create Outer Product: (B, L, L, D/2)\n        x_matrix = x_pair.unsqueeze(1) + x_pair.unsqueeze(2) \n        \n        # FIX: Permute to (B, D/2, L, L) BEFORE adding bias\n        x_matrix = x_matrix.permute(0, 3, 1, 2)\n        \n        # Now add Bias (1, D/2, L, L)\n        x_matrix = x_matrix + self.rel_pos(L, x_1d.device)\n        \n        # Add Recycling\n        if prev_pair is not None: \n            x_matrix = x_matrix + prev_pair\n            \n        # 3. Triangle Update\n        x_matrix = self.pair_act(self.triangle_block(x_matrix))\n        \n        # 4. Back to 1D\n        x_1d = x_1d + self.pair_to_seq(x_matrix.mean(dim=3).permute(0, 2, 1))\n        return self.encoder(x_1d), x_matrix\n\n# [MODEL 2] 1D Sequence Eyes\nclass RibonanzaSimple(nn.Module):\n    def __init__(self, embed_dim=32):\n        super().__init__()\n        self.embedding = nn.Embedding(5, embed_dim)\n        self.conv_stem = nn.Sequential(\n            nn.Conv1d(embed_dim, embed_dim, kernel_size=7, padding=3),\n            nn.GELU(),\n            # CRITICAL FIX: BatchNorm1d instead of LayerNorm prevents Runtime Crash\n            nn.BatchNorm1d(embed_dim) \n        )\n        self.encoder = nn.TransformerEncoder(nn.TransformerEncoderLayer(d_model=embed_dim, nhead=4, dim_feedforward=embed_dim*2, batch_first=True), num_layers=2)\n    \n    def forward(self, x):\n        x = self.embedding(x).transpose(1, 2)\n        x = self.conv_stem(x).transpose(1, 2)\n        return self.encoder(x)\n\n# [MASTER] Fusion\nclass RibonanzaFusion(nn.Module):\n    def __init__(self, model_ultra, model_simple):\n        super().__init__()\n        self.model_ultra = model_ultra\n        self.model_simple = model_simple\n        dim_u, dim_s = model_ultra.embed_dim, model_simple.embedding.embedding_dim\n        \n        self.fusion_proj = nn.Linear(dim_u + dim_s, dim_u)\n        self.gate = nn.Sequential(nn.Linear(dim_u + dim_s, dim_u), nn.Sigmoid())\n        self.head = nn.Sequential(nn.Linear(dim_u, dim_u), nn.GELU(), nn.Linear(dim_u, 3))\n        \n    def forward(self, tokens, pssm):\n        B, L = tokens.shape\n        prev_pair = torch.zeros(B, self.model_ultra.embed_dim//2, L, L, device=tokens.device)\n        \n        # Run Ultra\n        feat_ultra, _ = self.model_ultra.forward_step(tokens, pssm, prev_pair)\n        # Run Simple\n        feat_simple = self.model_simple(tokens)\n        \n        # Fuse\n        combined = torch.cat([feat_ultra, feat_simple], dim=-1)\n        fused = self.fusion_proj(combined) * self.gate(combined)\n        \n        # Predict 5\n        preds = []\n        for i in range(5):\n            noise = torch.randn_like(fused) * (0.15 * (i+1))\n            preds.append(self.head(fused + noise))\n            \n        return torch.stack(preds, dim=1)\n\n# ==========================================\n# 4. INFERENCE PIPELINE (Scoring Error Safe)\n# ==========================================\ndef format_submission_rows(seq_id, sequence, coords_5_models):\n    rows = []\n    # Auto-scale check (nm -> Angstroms)\n    avg_dist = np.abs(np.diff(coords_5_models, axis=1)).mean()\n    if avg_dist < 0.5: coords_5_models *= 10.0\n        \n    for i, char in enumerate(sequence):\n        row_data = {'ID': f\"{seq_id}_{i+1}\", 'resname': char, 'resid': i+1}\n        for m in range(5):\n            # SAFETY: Clip values and fix NaNs to prevent Scoring Error\n            val = coords_5_models[m, i]\n            val = np.nan_to_num(val, nan=0.0)\n            val = np.clip(val, -9999.0, 9999.0)\n            \n            row_data[f'x_{m+1}'] = f\"{val[0]:.3f}\"\n            row_data[f'y_{m+1}'] = f\"{val[1]:.3f}\"\n            row_data[f'z_{m+1}'] = f\"{val[2]:.3f}\"\n        rows.append(row_data)\n    return rows\n\ndef run_inference():\n    print(\"🚀 Starting Fusion Pipeline (Safe Mode)...\")\n    ds = RNADataset(config.TEST_CSV)\n    dl = DataLoader(ds, batch_size=config.BATCH_SIZE, shuffle=False)\n    \n    # Initialize Small Models\n    m_ultra = RibonanzaUltra(embed_dim=config.EMBED_DIM, num_layers=config.NUM_LAYERS).to(config.DEVICE)\n    m_simple = RibonanzaSimple(embed_dim=config.EMBED_DIM//2).to(config.DEVICE)\n    model = RibonanzaFusion(m_ultra, m_simple).to(config.DEVICE)\n    model.eval()\n    \n    # Load Weights (Safe)\n    if os.path.exists(config.WEIGHTS_PATH):\n        try:\n            state = torch.load(config.WEIGHTS_PATH, map_location=config.DEVICE)\n            model.load_state_dict(state, strict=False)\n            print(\"✅ Weights Loaded!\")\n        except: print(\"⚠️ Weights mismatch. Running with initialized weights.\")\n    else:\n        print(\"⚠️ No weights found. Running with initialized weights.\")\n\n    all_rows = []\n    print(f\"Processing {len(ds)} sequences...\")\n    \n    with torch.no_grad():\n        for batch in tqdm(dl):\n            tokens, pssm = batch['tokens'].to(config.DEVICE), batch['pssm'].to(config.DEVICE)\n            \n            # MEMORY TRICK: Mixed Precision\n            with torch.cuda.amp.autocast():\n                preds = model(tokens, pssm)\n            \n            preds = preds.float().cpu().numpy()\n            \n            for b in range(len(batch['id'])):\n                all_rows.extend(format_submission_rows(batch['id'][b], batch['seq_str'][b], preds[b]))\n            \n            # Aggressive Cleanup\n            del tokens, pssm, preds\n            torch.cuda.empty_cache()\n            if len(all_rows) % 100 == 0: gc.collect()\n\n    sub_df = pd.DataFrame(all_rows)\n    # Enforce column order for competition\n    cols = ['ID', 'resname', 'resid'] + [f'{ax}_{m}' for m in range(1,6) for ax in ['x','y','z']]\n    sub_df = sub_df[cols]\n    \n    sub_df.to_csv(config.SUBMISSION_PATH, index=False)\n    print(f\"🎉 Saved {len(sub_df)} rows to {config.SUBMISSION_PATH}\")\n\nif __name__ == \"__main__\":\n    run_inference()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-14T12:18:43.44748Z","iopub.execute_input":"2026-02-14T12:18:43.44876Z","iopub.status.idle":"2026-02-14T12:19:44.757548Z","shell.execute_reply.started":"2026-02-14T12:18:43.44871Z","shell.execute_reply":"2026-02-14T12:19:44.756424Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}