{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":87793,"databundleVersionId":12276181,"sourceType":"competition"},{"sourceId":13397805,"sourceType":"datasetVersion","datasetId":8501898}],"dockerImageVersionId":31154,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# ==================================\n# Standard Library Imports\n# ==================================\nimport glob\nimport math\nimport os\nimport random\nimport time\nimport warnings\n\n# ==================================\n# Third-Party Imports\n# ==================================\n# Scientific Computing & Data Handling\nimport numpy as np\nimport pandas as pd\nfrom tqdm.notebook import tqdm\n\n# PyTorch & Deep Learning\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.cuda.amp import GradScaler\nfrom torch.nn.utils.rnn import pad_sequence\nfrom torch.optim.lr_scheduler import _LRScheduler\nfrom torch.utils.data import DataLoader, Dataset, Sampler\n\n# Scikit-learn\nfrom sklearn.metrics import classification_report, f1_score\nfrom sklearn.model_selection import train_test_split\n\n# Bioinformatics\nimport tmtools\nfrom Bio.PDB import MMCIFParser\nfrom Bio.PDB.PDBExceptions import PDBConstructionWarning\n\n# Visualization\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\n# ==================================\n# Initial Setup\n# ==================================\n# Suppress specific warnings for cleaner output\nwarnings.simplefilter('ignore', PDBConstructionWarning)\n\n# ==================================\n# Configuration\n# ==================================\n# --- Input Data Locations (Read-Only) ---\nCIF_DIR = \"/kaggle/input/stanford-rna-3d-folding/PDB_RNA\"\nBASE_DIR = \"/kaggle/input/3d-rna-geoformer\"\nMETADATA_PATH = \"/kaggle/input/3d-rna-geoformer/full_metadata.csv\"\nCACHE_DIR = \"/kaggle/input/3d-rna-geoformer/cache/cache\"\n\n# --- Output Locations (Writable) ---\nOUTPUT_DIR = \"/kaggle/working/\"\nPRETRAINED_MODEL_PATH = os.path.join(OUTPUT_DIR, \"best_structural_model.pth\")\nFINETUNED_MODEL_PATH = os.path.join(OUTPUT_DIR, \"best_functional_model.pth\")\n\n# --- Device and General Hyperparameters ---\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nBATCH_SIZE = 16\nGRAD_ACCUMULATION_STEPS = 1\nMAX_LEN = 300\nEPOCHS = 10\nLEARNING_RATE = 5e-5\nCLIP_GRAD_NORM = 1.0\nWARMUP_STEPS = 100\n\n# --- Tuned Loss Weights ---\nFAPE_WEIGHT = 1.5\nTORSIONAL_WEIGHT = 0.5\nSTERIC_CLASH_WEIGHT = 0.1\nSECONDARY_STRUCTURE_WEIGHT = 0.3\nTRIPLET_WEIGHT = 1.0\nFAPE_CLAMP_DIST = 2.0\nINITIAL_FX_WEIGHT = 0.5\nFINAL_FX_WEIGHT = 4.0\n\n# --- Model Hyperparameters ---\nN_BLOCKS = 4\nD_MODEL = 128\nD_POINT = 4\nN_HEADS = 4\nD_HEAD_SCALAR = 8\nD_HEAD_POINT = 2\nFF_DIM = 256\nDROPOUT_RATE = 0.1\n\n# ==================================\n# Verification\n# ==================================\nprint(f\"✅ Setup complete.\")\nprint(f\"Using device: {DEVICE}\")\nprint(f\"Training with Smart Batching (Batch Size: {BATCH_SIZE})\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-21T06:00:11.212677Z","iopub.execute_input":"2025-11-21T06:00:11.212962Z","iopub.status.idle":"2025-11-21T06:00:17.267080Z","shell.execute_reply.started":"2025-11-21T06:00:11.212936Z","shell.execute_reply":"2025-11-21T06:00:17.266404Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def extract_info_from_cif(cif_path):\n    SITE_ANNOTATION_TAGS = [\"_struct_site.id\", \"_pdbx_struct_binding_site.id\"]\n    try:\n        has_functional_site = 0\n        with open(cif_path, 'r', errors='ignore') as f:\n            for line in f:\n                if any(line.strip().startswith(tag) for tag in SITE_ANNOTATION_TAGS):\n                    has_functional_site = 1\n                    break\n        parser = MMCIFParser(QUIET=True)\n        target_id = os.path.basename(cif_path).replace('.cif', '')\n        structure = parser.get_structure(target_id, cif_path)\n        res_map = {\"A\": \"A\", \"U\": \"U\", \"G\": \"G\", \"C\": \"C\"}\n        sequence = \"\"\n        for model in structure:\n            for chain in model:\n                for residue in chain:\n                    res_id, res_name = residue.get_id(), residue.get_resname().strip()\n                    if res_id[0] == ' ' and res_name in res_map:\n                        sequence += res_name\n            break\n        if not sequence: return None\n        return {'target_id': target_id, 'sequence': sequence, 'has_functional_site': has_functional_site}\n    except Exception:\n        return None\n\ndef generate_metadata_if_needed(cif_dir, output_dir):\n    full_metadata_path = os.path.join(output_dir, 'full_metadata.csv')\n    if os.path.exists(full_metadata_path):\n        print(\"✅ Metadata file already exists. Skipping generation.\")\n        df = pd.read_csv(full_metadata_path)\n        positive_count = df['has_functional_site'].sum()\n        print(f\"📄 Found {int(positive_count)} positive functional site samples in the existing metadata.\")\n        return df\n    cif_files = glob.glob(os.path.join(cif_dir, '*.cif'))\n    if not cif_files: raise FileNotFoundError(f\"No .cif files found in {cif_dir}\")\n    print(f\"⏳ Generating metadata for {len(cif_files)} CIF files...\")\n    data = [extract_info_from_cif(path) for path in tqdm(cif_files, desc=\"Scanning CIFs\")]\n    df = pd.DataFrame([d for d in data if d])\n    df.to_csv(full_metadata_path, index=False)\n    print(f\"✅ Metadata generated and saved to {full_metadata_path}\")\n    return df\n\nos.makedirs(BASE_DIR, exist_ok=True)\nos.makedirs(CIF_DIR, exist_ok=True)\nos.makedirs(CACHE_DIR, exist_ok=True)\nfull_df = generate_metadata_if_needed(CIF_DIR, BASE_DIR)\nprint(\"\\nMetadata generation complete.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-21T06:00:17.268219Z","iopub.execute_input":"2025-11-21T06:00:17.268571Z","iopub.status.idle":"2025-11-21T06:00:17.536974Z","shell.execute_reply.started":"2025-11-21T06:00:17.268551Z","shell.execute_reply":"2025-11-21T06:00:17.536214Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"--- Exploratory Data Analysis ---\")\nfull_df['seq_length'] = full_df['sequence'].str.len()\n\n# Graph 1: Distribution of Sequence Lengths\nplt.figure(figsize=(12, 5))\nplt.subplot(1, 2, 1)\nsns.histplot(full_df['seq_length'], bins=50, kde=True)\nplt.title('Distribution of RNA Sequence Lengths')\nplt.xlabel('Sequence Length'); plt.ylabel('Count')\n\n# Graph 2: Class Distribution\nplt.subplot(1, 2, 2)\nfull_df['has_functional_site'].value_counts().plot(kind='pie', autopct='%1.1f%%', colors=['skyblue', 'salmon'], labels=['Non-Functional', 'Functional'])\nplt.title('Class Distribution: Functional vs. Non-Functional')\nplt.ylabel('')\nplt.tight_layout(); plt.show()\n\n# Graph 3: Nucleotide Composition\nnucleotide_counts = pd.Series(list(''.join(full_df['sequence']))).value_counts()\nplt.figure(figsize=(12, 5))\nplt.subplot(1, 2, 1)\nsns.barplot(x=nucleotide_counts.index, y=nucleotide_counts.values)\nplt.title('Overall Nucleotide Composition')\nplt.xlabel('Nucleotide'); plt.ylabel('Total Count')\n\n# Graph 4: Sequence Length by Class\nplt.subplot(1, 2, 2)\nsns.boxplot(x='has_functional_site', y='seq_length', data=full_df)\nplt.title('Sequence Length vs. Functional Site Presence')\nplt.xlabel('Has Functional Site'); plt.ylabel('Sequence Length')\nplt.xticks([0, 1], ['No', 'Yes'])\nplt.tight_layout(); plt.show()\n\n# Graph 5: GC Content Distribution\nfull_df['gc_content'] = full_df['sequence'].apply(lambda x: (x.count('G') + x.count('C')) / len(x))\nplt.figure(figsize=(8, 5))\nsns.histplot(data=full_df, x='gc_content', hue='has_functional_site', multiple='stack', bins=30, kde=True)\nplt.title('GC Content Distribution by Class')\nplt.xlabel('GC Content'); plt.ylabel('Count')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-21T06:00:17.537760Z","iopub.execute_input":"2025-11-21T06:00:17.537997Z","iopub.status.idle":"2025-11-21T06:00:19.942398Z","shell.execute_reply.started":"2025-11-21T06:00:17.537978Z","shell.execute_reply":"2025-11-21T06:00:19.941680Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def calculate_dihedral(p0, p1, p2, p3):\n    b0 = -1.0 * (p1 - p0); b1 = p2 - p1; b2 = p3 - p2\n    b1_norm = np.linalg.norm(b1, axis=-1, keepdims=True)\n    b1_safe = np.divide(b1, b1_norm, out=np.zeros_like(b1), where=b1_norm!=0)\n    v = b0 - np.sum(b0 * b1_safe, axis=-1, keepdims=True) * b1_safe\n    w = b2 - np.sum(b2 * b1_safe, axis=-1, keepdims=True) * b1_safe\n    x = np.sum(v * w, axis=-1); y = np.sum(np.cross(b1_safe, v) * w, axis=-1)\n    return np.arctan2(y, x)\n\ndef process_cif_file(args):\n    cif_path, sequence, atom_map, num_atoms = args\n    site_residues = set()\n    try:\n        with open(cif_path, 'r', errors='ignore') as f: lines = f.readlines()\n        in_site_gen_loop, header_map = False, {}\n        for line in lines:\n            s_line = line.strip()\n            if s_line.startswith('loop_'): in_site_gen_loop, header_map = False, {}\n            elif s_line.startswith('_struct_site_gen.'):\n                in_site_gen_loop = True\n                header_map[s_line] = len(header_map)\n            elif in_site_gen_loop and not s_line.startswith('#') and s_line:\n                parts = s_line.split()\n                id_col, seq_col = header_map.get('_struct_site_gen.auth_asym_id'), header_map.get('_struct_site_gen.auth_seq_id')\n                if id_col is not None and seq_col is not None and len(parts) > max(id_col, seq_col):\n                    try: site_residues.add((parts[id_col], int(parts[seq_col])))\n                    except (ValueError, IndexError): continue\n        parser = MMCIFParser(QUIET=True)\n        structure = parser.get_structure(\"RNA\", cif_path)\n        model, seq_len = structure[0], len(sequence)\n        coords, fx_sites = np.full((seq_len, num_atoms, 3), np.nan, dtype=np.float32), np.zeros(seq_len, dtype=np.float32)\n        res_idx = 0\n        for chain in model:\n            if res_idx >= seq_len: break\n            for residue in chain:\n                if res_idx >= seq_len: break\n                res_id_tuple = residue.get_id()\n                if res_id_tuple[0] == ' ' and residue.get_resname().strip() in ['A', 'U', 'G', 'C']:\n                    if (chain.id, res_id_tuple[1]) in site_residues: fx_sites[res_idx] = 1.0\n                    for atom in residue:\n                        atom_name = atom.get_name().replace(\"*\", \"'\")\n                        if atom_name in atom_map: coords[res_idx, atom_map[atom_name], :] = atom.get_coord()\n                    res_idx += 1\n        atom_mask = ~np.isnan(coords).any(axis=-1)\n        coords[np.isnan(coords)] = 0\n        return {'target_id': os.path.basename(cif_path).replace('.cif', ''), 'sequence': sequence, 'coords': coords, 'atom_mask': atom_mask, 'fx_sites': fx_sites}\n    except Exception: return None\n\ndef rna_collate_fn_fx(batch):\n    batch = [item for item in batch if item is not None]\n    if not batch: return None\n    keys = batch[0].keys(); collated = {k: [d[k] for d in batch] for k in keys}\n    padded_sequences = pad_sequence([torch.tensor(s) for s in collated['sequence']], batch_first=True, padding_value=4).long()\n    coord_mask = (padded_sequences != 4).float()\n    return (padded_sequences, pad_sequence([torch.from_numpy(c) for c in collated['coords']], batch_first=True),\n            pad_sequence([torch.from_numpy(m) for m in collated['atom_mask']], batch_first=True), \n            pad_sequence([torch.from_numpy(f) for f in collated['fx_sites']], batch_first=True),\n            coord_mask, \n            pad_sequence([torch.from_numpy(t) for t in collated['torsionals']], batch_first=True),\n            pad_sequence([torch.from_numpy(m) for m in collated['torsionals_mask']], batch_first=True), \n            collated['target_id'])\n\nclass LengthBasedBatchSampler(Sampler):\n    def __init__(self, dataset, batch_size, drop_last):\n        self.dataset, self.batch_size, self.drop_last = dataset, batch_size, drop_last\n        self.groups = {}\n        for idx in range(len(dataset)):\n            length = len(dataset.metadata_df.iloc[idx]['sequence'])\n            if length not in self.groups: self.groups[length] = []\n            self.groups[length].append(idx)\n        self.batches = self._generate_batches()\n    def _generate_batches(self):\n        batches = []\n        for group in self.groups.values():\n            random.shuffle(group)\n            for i in range(0, len(group), self.batch_size):\n                batch = group[i:i+self.batch_size]\n                if len(batch) == self.batch_size or not self.drop_last: batches.append(batch)\n        random.shuffle(batches)\n        return batches\n    def __iter__(self): return iter(self.batches)\n    def __len__(self): return len(self.batches)\n\nclass RNADataset(Dataset):\n    def __init__(self, metadata_df, cif_dir, cache_dir, max_len=350, coord_mean=None, coord_std=None):\n        self.metadata_df = metadata_df.copy()\n        if max_len: self.metadata_df = self.metadata_df[self.metadata_df['sequence'].str.len() <= max_len].reset_index(drop=True)\n        self.cif_dir, self.cache_dir = cif_dir, cache_dir\n        os.makedirs(self.cache_dir, exist_ok=True)\n        self.atom_order=[\"P\",\"OP1\",\"OP2\",\"O5'\",\"C5'\",\"C4'\",\"O4'\",\"C3'\",\"O3'\",\"C2'\",\"O2'\",\"C1'\",\"N1\",\"C2\",\"N3\",\"C4\",\"C5\",\"C6\",\"N7\",\"C8\",\"N9\"]\n        self.atom_map = {name: i for i, name in enumerate(self.atom_order)}\n        self.num_atoms = len(self.atom_order)\n        self.nuc_map = {'A': 0, 'U': 1, 'G': 2, 'C': 3, 'N': 4}\n        self.data_cache = [None] * len(self.metadata_df)\n        self.coord_mean, self.coord_std = coord_mean, coord_std\n        if self.coord_mean is None: self._calculate_normalization_stats()\n    def _calculate_normalization_stats(self):\n        print(\"Calculating normalization stats (parsing files if not cached)...\")\n        all_coords_list = []\n        for idx in tqdm(range(len(self)), desc=\"Scanning for Stats\"):\n            data = self._get_item_data(idx)\n            if data and data['atom_mask'].any(): all_coords_list.append(data['coords'][data['atom_mask']])\n        if not all_coords_list: self.coord_mean, self.coord_std = 0.0, 1.0; return\n        all_coords_np = np.concatenate(all_coords_list)\n        self.coord_mean, self.coord_std = np.mean(all_coords_np), np.std(all_coords_np)\n        print(f\"Coord Mean: {self.coord_mean:.4f}, Std: {self.coord_std:.4f}\")\n    def _get_item_data(self, idx):\n        if idx < len(self.data_cache) and self.data_cache[idx] is not None: return self.data_cache[idx]\n        row = self.metadata_df.iloc[idx]\n        target_id, sequence = row['target_id'], row['sequence']\n        cache_path = os.path.join(self.cache_dir, f\"{target_id}.pt\")\n        if os.path.exists(cache_path): \n            data = torch.load(cache_path, weights_only=False)\n            if idx < len(self.data_cache): self.data_cache[idx] = data\n            return data\n        data = process_cif_file((os.path.join(self.cif_dir, f\"{target_id}.cif\"), sequence, self.atom_map, self.num_atoms))\n        if data: \n            torch.save(data, cache_path)\n            if idx < len(self.data_cache): self.data_cache[idx] = data\n        return data\n    def _get_torsional_angles(self, coords, atom_mask):\n        seq_len,_,_ = coords.shape; torsionals=np.zeros((seq_len,7,2),dtype=np.float32); torsionals_mask=np.zeros((seq_len,7),dtype=np.float32)\n        indices = self.atom_map\n        p,o5p,c5p,c4p,c3p,o3p = (indices.get(n, -1) for n in [\"P\",\"O5'\",\"C5'\",\"C4'\",\"C3'\",\"O3'\"])\n        for i in range(seq_len):\n            if i>0 and all(atom_mask[i-1,idx] for idx in [c4p,c3p,o3p] if idx!=-1) and atom_mask[i,p]:\n                angle=calculate_dihedral(coords[i-1,c4p],coords[i-1,c3p],coords[i-1,o3p],coords[i,p]); torsionals[i,0,:]=[np.sin(angle),np.cos(angle)]; torsionals_mask[i,0]=1.0\n            if i>0 and all(atom_mask[i-1,o3p] and atom_mask[i,idx] for idx in [p,o5p,c5p] if idx!=-1):\n                angle=calculate_dihedral(coords[i-1,o3p],coords[i,p],coords[i,o5p],coords[i,c5p]); torsionals[i,1,:]=[np.sin(angle),np.cos(angle)]; torsionals_mask[i,1]=1.0\n        return torsionals,torsionals_mask\n    def __len__(self): return len(self.metadata_df)\n    def __getitem__(self, idx):\n        data = self._get_item_data(idx)\n        if data is None: return self.__getitem__(np.random.randint(0, len(self)))\n        seq_tokens = [self.nuc_map.get(n, 4) for n in data['sequence']]\n        torsionals, torsionals_mask = self._get_torsional_angles(data['coords'], data['atom_mask'])\n        coords_for_norm = np.copy(data['coords']); coords_for_norm[~data['atom_mask']] = self.coord_mean\n        normalized_coords = (coords_for_norm - self.coord_mean) / (self.coord_std + 1e-8)\n        normalized_coords[~data['atom_mask']] = 0\n        return {\"sequence\":seq_tokens, \"coords\":normalized_coords, \"atom_mask\":data['atom_mask'], \"fx_sites\":data['fx_sites'], \"torsionals\":torsionals, \"torsionals_mask\":torsionals_mask, \"target_id\":data['target_id']}\n\nclass RefinementEGNNLayer(nn.Module):\n    def __init__(self, d_model):\n        super().__init__()\n        self.message_mlp = nn.Sequential(nn.Linear(d_model*2 + 1, d_model), nn.SiLU(), nn.Linear(d_model, d_model))\n        self.update_mlp = nn.Sequential(nn.Linear(d_model*2, d_model), nn.SiLU(), nn.Linear(d_model, d_model))\n        self.coord_update_mlp = nn.Sequential(nn.Linear(d_model, d_model), nn.SiLU(), nn.Linear(d_model, 1, bias=False))\n    def forward(self, s, coords, edge_index):\n        row, col = edge_index\n        rel_coords = coords[row] - coords[col]\n        dist = torch.norm(rel_coords, p=2, dim=-1, keepdim=True)\n        edge_features = torch.cat([s[row], s[col], dist], dim=-1)\n        messages = self.message_mlp(edge_features)\n        coord_update_scalar = self.coord_update_mlp(messages)\n        coord_shifts = (rel_coords / (dist + 1e-8)) * coord_update_scalar\n        agg_messages = torch.zeros_like(s).index_add_(0, col, messages.float())\n        agg_coord_shifts = torch.zeros_like(coords).index_add_(0, col, coord_shifts)\n        update_input = torch.cat([s, agg_messages], dim=-1)\n        s_out = s + self.update_mlp(update_input)\n        coords_out = coords + agg_coord_shifts\n        return s_out, coords_out\n\nclass SpatialRefinementModule(nn.Module):\n    def __init__(self, d_model, n_layers=2):\n        super().__init__()\n        self.layers = nn.ModuleList([RefinementEGNNLayer(d_model) for _ in range(n_layers)])\n        self.norm = nn.LayerNorm(d_model)\n    def forward(self, s, coords, coord_mask):\n        B, L, _ = s.shape\n        dist_matrix = torch.cdist(coords, coords)\n        mask = coord_mask.unsqueeze(1) * coord_mask.unsqueeze(2)\n        dist_matrix.masked_fill_(mask == 0, float('inf'))\n        k = min(16, L)\n        _, edge_index_col = torch.topk(dist_matrix, k=k, dim=-1, largest=False)\n        base_row = torch.arange(L, device=s.device).unsqueeze(-1).expand(-1, k)\n        rows, cols = [], []\n        for i in range(B):\n            offset = i * L\n            rows.append(base_row + offset)\n            cols.append(edge_index_col[i] + offset)\n        row_tensor, col_tensor = torch.cat(rows).flatten(), torch.cat(cols).flatten()\n        batch_edge_index = torch.stack([row_tensor, col_tensor])\n        s_flat, coords_flat = s.view(B * L, -1), coords.view(B * L, -1)\n        refined_s_flat = s_flat\n        for layer in self.layers:\n            refined_s_flat, _ = layer(refined_s_flat, coords_flat, batch_edge_index)\n        refined_s = refined_s_flat.view(B, L, -1)\n        return self.norm(s + refined_s)\n\nclass InvariantPointAttention(nn.Module):\n    def __init__(self, d_model, d_point, n_heads, d_head_point, d_head_scalar):\n        super().__init__()\n        self.n_heads,self.d_head_point=n_heads,d_head_point\n        self.q_scalar,self.k_scalar,self.v_scalar = [nn.Linear(d_model, n_heads * d_head_scalar) for _ in range(3)]\n        self.q_point,self.k_point,self.v_point = [nn.Linear(d_point * 3, n_heads * d_head_point * 3) for _ in range(3)]\n        self.trainable_point_weights = nn.Parameter(torch.randn(n_heads))\n        self.attn_out = nn.Linear(n_heads * (d_head_scalar + d_head_point * 3), d_model)\n        self.gamma = 1/math.sqrt(d_head_scalar)\n        self.register_buffer('virtual_points', torch.randn(d_point, 3))\n    def forward(self, s, z, rotations, translations, coord_mask):\n        B,L,_=s.shape\n        points=torch.einsum('blij,pj->blpi',rotations,self.virtual_points)+translations.unsqueeze(-2)\n        points_flat=points.reshape(B,L,-1)\n        q_s,k_s,v_s=self.q_scalar(s),self.k_scalar(s),self.v_scalar(s)\n        q_p,k_p,v_p=self.q_point(points_flat),self.k_point(points_flat),self.v_point(points_flat)\n        q_s,k_s,v_s=[x.reshape(B,L,self.n_heads,-1) for x in [q_s,k_s,v_s]]\n        q_p,k_p,v_p=[x.reshape(B,L,self.n_heads,-1,3) for x in [q_p,k_p,v_p]]\n        attn_logits=self.gamma*torch.einsum('bihd,bjhd->bijh',q_s,k_s)-0.5*torch.sum((q_p.unsqueeze(2)-k_p.unsqueeze(1))**2,dim=(-1,-2))*self.trainable_point_weights\n        attn_logits+=z.unsqueeze(0).permute(0,2,3,1)\n        mask=coord_mask.unsqueeze(1).unsqueeze(-1)*coord_mask.unsqueeze(2).unsqueeze(-1)\n        attn_logits=attn_logits.masked_fill(mask==0,-1e9)\n        attn=F.softmax(attn_logits,dim=2)\n        result_s=torch.einsum('bijh,bjhd->bihd',attn,v_s)\n        result_p=torch.einsum('bijh,bjhdp->bihdp',attn,v_p)\n        output=self.attn_out(torch.cat([result_s.reshape(B,L,-1),result_p.reshape(B,L,-1)],dim=-1))\n        return output\n\nclass EGNNLayer(nn.Module):\n    def __init__(self, d_model):\n        super().__init__()\n        self.message_mlp = nn.Sequential(nn.Linear(d_model*2 + 1, d_model), nn.SiLU(), nn.Linear(d_model, d_model))\n        self.update_mlp = nn.Sequential(nn.Linear(d_model*2, d_model), nn.SiLU(), nn.Linear(d_model, d_model))\n        self.coord_update_mlp = nn.Sequential(nn.Linear(d_model, d_model), nn.SiLU(), nn.Linear(d_model, 1, bias=False))\n    def forward(self, s, coords, edge_index):\n        row, col = edge_index\n        rel_coords = coords[:, row] - coords[:, col]\n        dist = torch.norm(rel_coords, p=2, dim=-1, keepdim=True)\n        edge_features = torch.cat([s[:, row], s[:, col], dist], dim=-1)\n        messages = self.message_mlp(edge_features)\n        coord_update_scalar = self.coord_update_mlp(messages)\n        coord_shifts = (rel_coords / (dist + 1e-8)) * coord_update_scalar\n        agg_messages = torch.zeros_like(s).index_add_(1, col, messages.float())\n        agg_coord_shifts = torch.zeros_like(coords).index_add_(1, col, coord_shifts)\n        update_input = torch.cat([s, agg_messages], dim=-1)\n        s_out = s + self.update_mlp(update_input)\n        coords_out = coords + agg_coord_shifts\n        return s_out, coords_out\n\nclass GeoformerIPABlock(nn.Module):\n    def __init__(self, d_model, d_point, n_heads, d_head_point, d_head_scalar, ff_dim, dropout=0.1):\n        super().__init__()\n        self.ipa = InvariantPointAttention(d_model, d_point, n_heads, d_head_point, d_head_scalar)\n        self.ipa_norm = nn.LayerNorm(d_model)\n        self.egnn = EGNNLayer(d_model)\n        self.egnn_norm = nn.LayerNorm(d_model)\n        self.ffn = nn.Sequential(nn.Linear(d_model, ff_dim), nn.ReLU(), nn.Dropout(dropout), nn.Linear(ff_dim, d_model), nn.Dropout(dropout))\n        self.ffn_norm = nn.LayerNorm(d_model)\n    def _forward_impl(self, s, z, rotations, translations, edge_index, coord_mask):\n        s = s + self.ipa(self.ipa_norm(s), z, rotations, translations, coord_mask)\n        s_norm, translations_norm = self.egnn_norm(s), translations\n        s, translations = self.egnn(s_norm, translations_norm, edge_index)\n        s = s + self.ffn(self.ffn_norm(s))\n        return s, translations\n    def forward(self, args):\n        s, z, rotations, translations, edge_index, coord_mask = args\n        return self._forward_impl(s, z, rotations, translations, edge_index, coord_mask)\n\nclass StructureModule(nn.Module):\n    def __init__(self, d_model, num_atoms): super().__init__(); self.num_atoms,self.atom_predictor=num_atoms,nn.Linear(d_model,num_atoms*3)\n    def forward(self, s, rotations, translations):\n        local_displacements = self.atom_predictor(s).view(*s.shape[:-1], self.num_atoms, 3)\n        return torch.einsum('blij,blaj->blai', rotations, local_displacements) + translations.unsqueeze(-2)\n\nclass SupervisedContrastiveLoss(nn.Module):\n    def __init__(self, temperature=0.07, max_samples=1024):\n        super(SupervisedContrastiveLoss, self).__init__()\n        self.temperature = temperature\n        self.max_samples = max_samples\n    def forward(self, features, labels):\n        if features.shape[0] == 0: return torch.tensor(0.0, device=features.device)\n        pos_indices = torch.where(labels == 1)[0]\n        neg_indices = torch.where(labels == 0)[0]\n        num_pos, num_neg = len(pos_indices), len(neg_indices)\n        if num_pos < 2 or num_neg < 2: return torch.tensor(0.0, device=features.device)\n        num_each = min(self.max_samples // 2, num_pos, num_neg)\n        sampled_indices = torch.cat([pos_indices[torch.randperm(num_pos)[:num_each]], neg_indices[torch.randperm(num_neg)[:num_each]]])\n        features, labels = features[sampled_indices], labels[sampled_indices]\n        features = F.normalize(features, p=2, dim=1)\n        labels = labels.contiguous().view(-1, 1)\n        mask = torch.eq(labels, labels.T).float().to(features.device)\n        anchor_dot_contrast = torch.div(torch.matmul(features, features.T), self.temperature)\n        logits_mask = torch.scatter(torch.ones_like(mask), 1, torch.arange(features.shape[0]).view(-1, 1).to(features.device), 0)\n        mask = mask * logits_mask\n        exp_logits = torch.exp(anchor_dot_contrast) * logits_mask\n        log_prob = anchor_dot_contrast - torch.log(exp_logits.sum(1, keepdim=True))\n        mean_log_prob_pos = (mask * log_prob).sum(1) / (mask.sum(1) + 1e-8)\n        loss = -mean_log_prob_pos[mean_log_prob_pos != 0].mean()\n        return loss if not torch.isnan(loss) else torch.tensor(0.0, device=features.device)\n\nclass AttentionPooling(nn.Module):\n    def __init__(self, input_dim, hidden_dim=128):\n        super().__init__()\n        self.attention_net = nn.Sequential(nn.Linear(input_dim, hidden_dim), nn.Tanh(), nn.Linear(hidden_dim, 1))\n    def forward(self, x, mask):\n        attention_logits = self.attention_net(x).squeeze(-1)\n        mask_value = torch.finfo(attention_logits.dtype).min\n        attention_logits.masked_fill_(mask == 0, mask_value)\n        attention_weights = F.softmax(attention_logits, dim=1).unsqueeze(1)\n        pooled_features = torch.bmm(attention_weights, x).squeeze(1)\n        return pooled_features\n\nclass GeoformerRNA(nn.Module):\n    def __init__(self, n_blocks, d_model, d_point, n_heads, d_head_point, d_head_scalar, ff_dim, num_atoms, dropout, rel_pos_bins=32):\n        super().__init__()\n        self.embedding_s = nn.Embedding(5, d_model, padding_idx=4)\n        self.rel_pos_embedding = nn.Embedding(2 * rel_pos_bins + 1, n_heads)\n        self.rel_pos_bins = rel_pos_bins\n        self.blocks = nn.ModuleList([GeoformerIPABlock(d_model, d_point, n_heads, d_head_point, d_head_scalar, ff_dim, dropout) for _ in range(n_blocks)])\n        self.structure_module = StructureModule(d_model, num_atoms)\n        self.to_s_point = nn.Linear(d_model, d_point * 3)\n        self.spatial_refiner = SpatialRefinementModule(d_model)\n        self.early_feature_proj = nn.Linear(d_model, d_model // 2)\n        fused_dim = d_model + d_model // 2\n        self.gate = nn.Sequential(nn.Linear(fused_dim, fused_dim), nn.Sigmoid())\n        self.functional_head = nn.Sequential(nn.LayerNorm(fused_dim), nn.Linear(fused_dim, 64), nn.ReLU(), nn.Dropout(0.25), nn.Linear(64, 1))\n        self.attention_pool = AttentionPooling(fused_dim)\n        self.projection_head = nn.Sequential(nn.Linear(fused_dim, fused_dim), nn.ReLU(), nn.Linear(fused_dim, 128))\n        self.torsional_head = nn.Sequential(nn.LayerNorm(d_model), nn.Linear(d_model, 64), nn.ReLU(), nn.Linear(64, 14))\n        self.secondary_structure_head = nn.Sequential(nn.LayerNorm(d_model), nn.Linear(d_model, 64), nn.ReLU())\n    def forward(self, seq, coord_mask):\n        B,L=seq.shape\n        s=self.embedding_s(seq)\n        pos=torch.arange(L,device=seq.device); rel_pos=torch.clamp(pos[None,:]-pos[:,None]+self.rel_pos_bins,0,2*self.rel_pos_bins)\n        z=self.rel_pos_embedding(rel_pos).permute(2,0,1)\n        rotations=torch.eye(3,device=seq.device,dtype=s.dtype).unsqueeze(0).unsqueeze(0).expand(B,L,-1,-1)\n        translations=torch.zeros(B,L,3,device=seq.device,dtype=s.dtype)\n        edge_index=torch.stack([torch.arange(L-1,device=seq.device),torch.arange(1,L,device=seq.device)],dim=0)\n        s_early = None\n        for i, block in enumerate(self.blocks):\n            s,translations=block((s,z,rotations,translations,edge_index,coord_mask))\n            s_point=self.to_s_point(s).view(B,L,-1,3)\n            translations=translations+torch.einsum('blij,blaj->blai',rotations,s_point).mean(dim=-2)\n            if i == 0: s_early = s\n        final_coords=self.structure_module(s,rotations,translations)\n        refined_s = self.spatial_refiner(s, translations, coord_mask)\n        s_early_proj = self.early_feature_proj(s_early)\n        raw_fused_features = torch.cat([s_early_proj, refined_s], dim=-1)\n        gate_values = self.gate(raw_fused_features)\n        fused_features = raw_fused_features * gate_values\n        functional_logits=self.functional_head(fused_features).squeeze(-1)\n        prototype = self.attention_pool(fused_features, coord_mask)\n        triplet_features = self.projection_head(prototype)\n        torsional_preds=F.normalize(self.torsional_head(s).view(B,L,7,2),dim=-1)\n        s_ss=self.secondary_structure_head(s)\n        ss_logits=torch.einsum('bid,bjd->bij',s_ss,s_ss)\n        return final_coords,rotations,translations,functional_logits,torsional_preds,ss_logits,triplet_features\n\nclass WarmupCosineScheduler(_LRScheduler):\n    def __init__(self,optimizer,warmup_steps,total_steps,last_epoch=-1): self.warmup_steps,self.total_steps=warmup_steps,total_steps; super().__init__(optimizer,last_epoch)\n    def get_lr(self):\n        if self.warmup_steps>0 and self.last_epoch<self.warmup_steps: return [base_lr*(self.last_epoch+1)/self.warmup_steps for base_lr in self.base_lrs]\n        progress=(self.last_epoch-self.warmup_steps)/max(1,self.total_steps-self.warmup_steps)\n        return [base_lr*0.5*(1.0+math.cos(math.pi*progress)) for base_lr in self.base_lrs]\n\nclass FocalLoss(nn.Module):\n    def __init__(self, alpha=0.1, gamma=2.0, pos_weight=None):\n        super(FocalLoss, self).__init__()\n        self.alpha, self.gamma, self.pos_weight = alpha, gamma, pos_weight\n    def forward(self, inputs, targets):\n        BCE_loss = F.binary_cross_entropy_with_logits(inputs, targets, reduction='none')\n        if self.pos_weight is not None:\n            weight_tensor = torch.ones_like(targets); weight_tensor[targets == 1] = self.pos_weight.item()\n            BCE_loss = BCE_loss * weight_tensor\n        pt = torch.exp(-BCE_loss)\n        F_loss = self.alpha * (1 - pt)**self.gamma * BCE_loss\n        return F_loss.mean()\n\ndef FAPE_loss(pred_coords,true_coords,rotations,translations,atom_mask,coord_mask,clamp_dist=10.0):\n    rotations,translations=rotations.float(),translations.float(); pred_coords,true_coords=pred_coords.float(),true_coords.float()\n    inv_rots=rotations.transpose(-1,-2)\n    relative_true_coords=true_coords.unsqueeze(1)-translations.unsqueeze(2).unsqueeze(3)\n    local_true_coords=torch.einsum('bilk,bijak->bijal',inv_rots,relative_true_coords)\n    error=torch.sqrt(torch.sum((local_true_coords-pred_coords.unsqueeze(2))**2,dim=-1)+1e-8)\n    mask=(coord_mask.unsqueeze(2)*coord_mask.unsqueeze(1)).unsqueeze(-1)*atom_mask.unsqueeze(1)\n    return (torch.clamp(error,max=clamp_dist)*mask).sum()/(mask.sum()+1e-8)\n\ndef torsional_loss(pred_torsionals,true_torsionals,mask): return ((1-torch.sum(pred_torsionals*true_torsionals,dim=-1))*mask).sum()/(mask.sum()+1e-8)\n\ndef steric_clash_loss(pred_coords,atom_mask,coord_mask,c4_idx=5,clash_threshold=1.5):\n    B,L,_,_=pred_coords.shape\n    if L<2: return torch.tensor(0.0,device=pred_coords.device)\n    coords,mask=pred_coords[:,:,c4_idx,:],(atom_mask[:,:,c4_idx]*coord_mask).bool()\n    total_clash_loss=torch.tensor(0.0,device=pred_coords.device)\n    for b in range(B):\n        valid_coords=coords[b,mask[b]]\n        if valid_coords.shape[0]<2: continue\n        violations=F.relu(clash_threshold-torch.pdist(valid_coords))\n        total_clash_loss+=violations.mean() if violations.numel()>0 else 0.0\n    return total_clash_loss/B\n\ndef secondary_structure_geometric_loss(pred_coords,ss_logits,seq_map,atom_mask,atom_map):\n    with torch.no_grad():\n        c1_idx=atom_map.get(\"C1'\")\n        if c1_idx is None: return torch.tensor(0.0,device=pred_coords.device)\n        dist_c1=torch.cdist(pred_coords[:,:,c1_idx,:],pred_coords[:,:,c1_idx,:])\n        au_mask=((seq_map==0)|(seq_map==1)); gc_mask=((seq_map==2)|(seq_map==3))\n        pseudo_labels=((dist_c1<10.0)&au_mask.unsqueeze(1)&au_mask.unsqueeze(2))|((dist_c1<9.0)&gc_mask.unsqueeze(1)&gc_mask.unsqueeze(2))\n        pseudo_labels=pseudo_labels.float()\n        pseudo_labels.diagonal(dim1=-2,dim2=-1).zero_()\n        pseudo_labels*=atom_mask[:,:,c1_idx].unsqueeze(1)&atom_mask[:,:,c1_idx].unsqueeze(2)\n    ss_logits_symm=(ss_logits+ss_logits.transpose(-1,-2))/2\n    return F.binary_cross_entropy_with_logits(ss_logits_symm,pseudo_labels,reduction='none').sum()/(pseudo_labels.sum()+1e-8)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-21T06:00:19.944311Z","iopub.execute_input":"2025-11-21T06:00:19.944588Z","iopub.status.idle":"2025-11-21T06:00:20.016268Z","shell.execute_reply.started":"2025-11-21T06:00:19.944567Z","shell.execute_reply":"2025-11-21T06:00:20.015529Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def calculate_rmsd(pred_coords,true_coords,atom_masks):\n    sq_dist=torch.sum((pred_coords-true_coords)**2,dim=-1)\n    num_valid=atom_masks.sum()\n    return torch.sqrt((sq_dist*atom_masks).sum()/(num_valid+1e-8)).item() if num_valid>0 else 0.0\n\ndef train_epoch(model, loader, optimizer, scheduler, scaler, accumulation_steps, bce_pos_weight, atom_map, current_epoch, total_epochs):\n    model.train(); losses = {k: 0 for k in ['total', 'fape', 'torsion', 'clash', 'fx', 'ss', 'triplet']}\n    progress = current_epoch / max(1, total_epochs - 1)\n    current_fx_weight = INITIAL_FX_WEIGHT + progress * (FINAL_FX_WEIGHT - INITIAL_FX_WEIGHT)\n    focal_loss_fn = FocalLoss(alpha=0.1, gamma=2.0, pos_weight=bce_pos_weight).to(DEVICE)\n    triplet_loss_fn = nn.TripletMarginLoss(margin=0.5).to(DEVICE)\n    optimizer.zero_grad(set_to_none=True)\n    for i, batch in enumerate(tqdm(loader, desc=f\"Training Epoch {current_epoch+1}/{total_epochs}\", leave=False)):\n        if batch is None: continue\n        seqs, true_coords, atom_masks, fx_sites, coord_masks, true_torsionals, torsionals_masks = [t.to(DEVICE) for t in batch[:-1]]\n        with torch.amp.autocast(device_type='cuda', dtype=torch.float16, enabled=(DEVICE.type == 'cuda')):\n            pred_coords, rotations, translations, functional_logits, pred_torsionals, ss_logits, triplet_features = model(seqs, coord_masks)\n            loss_fape = FAPE_WEIGHT * FAPE_loss(pred_coords, true_coords, rotations, translations, atom_masks, coord_masks, clamp_dist=FAPE_CLAMP_DIST)\n            loss_torsion = TORSIONAL_WEIGHT * torsional_loss(pred_torsionals, true_torsionals, torsionals_masks)\n            loss_clash = STERIC_CLASH_WEIGHT * steric_clash_loss(pred_coords, atom_masks, coord_masks)\n            loss_ss = SECONDARY_STRUCTURE_WEIGHT * secondary_structure_geometric_loss(pred_coords, ss_logits, seqs, atom_masks, atom_map)\n            loss_fx = current_fx_weight * (focal_loss_fn(functional_logits, fx_sites) * coord_masks).sum() / (coord_masks.sum() + 1e-8)\n            batch_labels = (fx_sites.sum(dim=1) > 0).long()\n            pos_indices = torch.where(batch_labels == 1)[0]\n            neg_indices = torch.where(batch_labels == 0)[0]\n            loss_triplet = torch.tensor(0.0, device=DEVICE)\n            if len(pos_indices) >= 2 and len(neg_indices) >= 1:\n                anchor_idx = pos_indices[torch.randint(len(pos_indices), (1,))]\n                positive_idx = pos_indices[torch.randint(len(pos_indices), (1,))]\n                negative_idx = neg_indices[torch.randint(len(neg_indices), (1,))]\n                anchor = triplet_features[anchor_idx]\n                positive = triplet_features[positive_idx]\n                negative = triplet_features[negative_idx]\n                loss_triplet = TRIPLET_WEIGHT * current_fx_weight * triplet_loss_fn(anchor, positive, negative)\n            loss = loss_fape + loss_torsion + loss_clash + loss_ss + loss_fx + loss_triplet\n        if torch.isnan(loss) or torch.isinf(loss):\n            print(f\"WARNING: NaN or Inf loss at batch {i}. Skipping.\"); optimizer.zero_grad(set_to_none=True); continue\n        scaler.scale(loss / accumulation_steps).backward()\n        if (i + 1) % accumulation_steps == 0 or (i + 1) == len(loader):\n            scaler.unscale_(optimizer)\n            torch.nn.utils.clip_grad_norm_(model.parameters(), CLIP_GRAD_NORM)\n            scaler.step(optimizer)\n            scaler.update()\n            scheduler.step()\n            optimizer.zero_grad(set_to_none=True)\n        for k, v in zip(losses.keys(), [loss, loss_fape, loss_torsion, loss_clash, loss_fx, loss_ss, loss_triplet]):\n            losses[k] += v.item() if torch.is_tensor(v) else v\n    return {k: v/len(loader) for k, v in losses.items()}\n\ndef validate_epoch(model, loader, nuc_map):\n    model.eval(); all_results = []\n    c4_idx = 5\n    rev_nuc_map = {v: k for k, v in nuc_map.items()}\n    with torch.no_grad():\n        for batch in tqdm(loader, desc=\"Validating\", leave=False):\n            if batch is None: continue\n            tensors_to_move, target_ids = batch[:-1], batch[-1]\n            seqs, true_coords, atom_masks, fx_sites, coord_masks, _, _ = [t.to(DEVICE) for t in tensors_to_move]\n            with torch.amp.autocast(device_type='cuda', dtype=torch.float16, enabled=(DEVICE.type == 'cuda')):\n                 pred_coords, _, _, functional_logits, _, _, _ = model(seqs, coord_masks)\n            for i in range(seqs.shape[0]):\n                mask = coord_masks[i].bool()\n                seq_len = mask.sum().item()\n                tm_score = 0.0\n                try:\n                    true_c4 = true_coords[i, :seq_len, c4_idx].cpu().numpy()\n                    pred_c4 = pred_coords[i, :seq_len, c4_idx].cpu().numpy()\n                    sequence_str = \"\".join([rev_nuc_map.get(token.item(), 'A') for token in seqs[i, :seq_len]])\n                    if true_c4.shape[0] > 0 and pred_c4.shape[0] > 0:\n                        res = tmtools.tm_align(true_c4, pred_c4, sequence_str, sequence_str)\n                        tm_score = res.tm_norm_chain1\n                except:\n                    pass\n                all_results.append({\n                    'target_id': target_ids[i], \n                    'rmsd': calculate_rmsd(pred_coords[i], true_coords[i], atom_masks[i]),\n                    'tm_score': tm_score,\n                    'fx_true': fx_sites[i][mask].cpu().numpy().flatten().astype(int),\n                    'fx_preds': (torch.sigmoid(functional_logits[i][mask]).cpu().numpy() > 0.5).astype(int)\n                })\n    return all_results","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-21T06:00:20.017037Z","iopub.execute_input":"2025-11-21T06:00:20.017260Z","iopub.status.idle":"2025-11-21T06:00:20.036784Z","shell.execute_reply.started":"2025-11-21T06:00:20.017235Z","shell.execute_reply":"2025-11-21T06:00:20.036039Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_metrics_dashboard(results_df):\n    if results_df.empty: return\n    plt.figure(figsize=(14, 6)); sns.set_style(\"whitegrid\")\n    \n    # Plot losses\n    plt.subplot(1, 2, 1)\n    for col in [c for c in results_df.columns if 'loss' in c]: \n        plt.plot(results_df.index, results_df[col], label=col.replace('loss_', ''))\n    plt.xlabel('Epochs'); plt.ylabel('Loss'); plt.title('Training Loss Components'); plt.legend(); plt.yscale('log')\n    \n    # Plot validation metrics (RMSD and TM-Score)\n    plt.subplot(1, 2, 2)\n    ax1 = plt.gca()\n    ax1.plot(results_df.index, results_df['val_rmsd'], label='Validation RMSD (Å)', color='green', marker='o')\n    ax1.set_xlabel('Epochs'); ax1.set_ylabel('RMSD (Å)', color='green')\n    ax1.tick_params(axis='y', labelcolor='green')\n    ax1.legend(loc='upper left')\n\n    ax2 = ax1.twinx()\n    ax2.plot(results_df.index, results_df['val_tm_score'], label='Validation TM-Score', color='red', marker='.')\n    ax2.set_ylabel('TM-Score', color='red')\n    ax2.tick_params(axis='y', labelcolor='red')\n    ax2.legend(loc='upper right')\n    \n    plt.title('Validation Metrics'); plt.tight_layout(); plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-21T06:00:20.037590Z","iopub.execute_input":"2025-11-21T06:00:20.037860Z","iopub.status.idle":"2025-11-21T06:00:20.052192Z","shell.execute_reply.started":"2025-11-21T06:00:20.037845Z","shell.execute_reply":"2025-11-21T06:00:20.051668Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"y_stratify = full_df['has_functional_site']\nif y_stratify.sum() < 2:\n    y_stratify = None\n    print(\"⚠️ WARNING: Not enough positive samples for stratification.\")\n\ntrain_df, val_df = train_test_split(full_df, test_size=0.1, random_state=42, stratify=y_stratify)\nprint(f\"🔬 Training on {len(train_df)} samples, validating on {len(val_df)} samples.\")\n\ntrain_dataset = RNADataset(train_df, CIF_DIR, CACHE_DIR, max_len=MAX_LEN)\nval_dataset = RNADataset(val_df, CIF_DIR, CACHE_DIR, max_len=MAX_LEN, coord_mean=train_dataset.coord_mean, coord_std=train_dataset.coord_std)\n\ntrain_batch_sampler = LengthBasedBatchSampler(train_dataset, batch_size=BATCH_SIZE, drop_last=True)\ntrain_loader = DataLoader(train_dataset, batch_sampler=train_batch_sampler, collate_fn=rna_collate_fn_fx, num_workers=2, pin_memory=True)\nval_loader = DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=False, collate_fn=rna_collate_fn_fx, num_workers=2, pin_memory=True)\n\nnum_pos_residues = sum(d['fx_sites'].sum() for d in train_dataset.data_cache if d is not None and 'fx_sites' in d)\nnum_total_residues = sum(len(d['fx_sites']) for d in train_dataset.data_cache if d is not None and 'fx_sites' in d)\n\nif num_pos_residues == 0:\n    print(\"❌ CRITICAL ERROR: 0 positive residues found. Cannot train.\")\n    bce_pos_weight = torch.tensor([1.0], device=DEVICE)\nelse:\n    num_neg_residues = num_total_residues - num_pos_residues\n    bce_pos_weight = torch.tensor([num_neg_residues / (num_pos_residues + 1e-8)], device=DEVICE)\n    print(f\"⚖️ BCE Positive Weight: {bce_pos_weight.item():.2f} ({int(num_pos_residues)} positive residues found)\")\n\nmodel = GeoformerRNA(N_BLOCKS,D_MODEL,D_POINT,N_HEADS,D_HEAD_POINT,D_HEAD_SCALAR,FF_DIM,num_atoms=train_dataset.num_atoms,dropout=DROPOUT_RATE).to(DEVICE)\noptimizer = torch.optim.AdamW(model.parameters(), lr=LEARNING_RATE, weight_decay=0.01) # Recommended addition\nscaler = GradScaler(enabled=(DEVICE.type == 'cuda')) # This line is correct\nscheduler = WarmupCosineScheduler(optimizer, WARMUP_STEPS, len(train_loader) * EPOCHS)\n\nhistory = []\nstart_time = time.time()\nbest_val_tm_score = -1.0\n\nfor epoch in range(EPOCHS):\n    print(f\"\\n{'='*20} EPOCH {epoch + 1}/{EPOCHS} {'='*20}\")\n    losses = train_epoch(model, train_loader, optimizer, scheduler, scaler, GRAD_ACCUMULATION_STEPS, bce_pos_weight, train_dataset.atom_map, current_epoch=epoch, total_epochs=EPOCHS)\n    results = validate_epoch(model, val_loader, train_dataset.nuc_map)\n\n    val_rmsd = np.mean([r['rmsd'] for r in results]) if results else float('inf')\n    val_tm_score = np.mean([r['tm_score'] for r in results if r['tm_score'] > 0]) if results else 0.0\n\n    print(f\"  Epoch {epoch+1:02d}/{EPOCHS} -> Train Loss: {losses['total']:.4f} | Val RMSD: {val_rmsd:.4f} | Val TM-Score: {val_tm_score:.4f}\")\n\n    if val_tm_score > best_val_tm_score:\n        best_val_tm_score = val_tm_score\n        MODEL_SAVE_PATH = \"best_geoformer_model.pt\"\n\n        torch.save(model.state_dict(), MODEL_SAVE_PATH)\n        print(f\"  🏅 New best model saved to {MODEL_SAVE_PATH} (TM-Score: {best_val_tm_score:.4f})\")\n\n    epoch_data = {**{f'loss_{k}': v for k, v in losses.items()}, 'val_rmsd': val_rmsd, 'val_tm_score': val_tm_score}\n    history.append(epoch_data)\n\nprint(f\"\\n✅ Training finished in {(time.time() - start_time)/60:.2f} minutes.\")\nhistory_df = pd.DataFrame(history)\nplot_metrics_dashboard(history_df)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-21T06:00:20.052982Z","iopub.execute_input":"2025-11-21T06:00:20.053236Z","iopub.status.idle":"2025-11-21T06:06:10.561141Z","shell.execute_reply.started":"2025-11-21T06:00:20.053213Z","shell.execute_reply":"2025-11-21T06:06:10.560308Z"}},"outputs":[],"execution_count":null}]}