{"nbformat":4,"nbformat_minor":4,"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.10.0"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":92965,"databundleVersionId":13093583,"sourceType":"competition"},{"sourceType":"datasetVersion","sourceId":11198765,"datasetId":6566778}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"cells":[{"cell_type":"markdown","metadata":{},"source":"# NullivaRNA-Flow V17 Training\n\n**Stanford RNA 3D Folding Part 2 - Kaggle Competition**\n\n## V17 Key Fixes from V16\n\n1. **Bond loss in RAW SPACE**: Now computed in Angstroms, not normalized space\n2. **Lower bond_weight**: 0.5 instead of 2.0 (was dominating loss)\n3. **Pairwise distance loss**: Encourages correct global structure\n4. **Better geometry supervision**: Match true backbone + be reasonable\n\n**Target**: TM-score >= 0.5\n"},{"cell_type":"code","metadata":{},"execution_count":null,"outputs":[],"source":"import os\nimport sys\nimport time\nimport gc\nimport math\nimport glob\nimport warnings\nfrom pathlib import Path\nfrom dataclasses import dataclass\nfrom typing import Dict, Optional, Tuple, List\n\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\n\nprint('Imports OK')"},{"cell_type":"code","metadata":{},"execution_count":null,"outputs":[],"source":"# ==============================================================================\n# CONFIGURATION"},{"cell_type":"code","metadata":{},"execution_count":null,"outputs":[],"source":"# ==============================================================================\n\ntorch.backends.cuda.matmul.allow_tf32 = False\ntorch.backends.cudnn.allow_tf32 = False\n\nKAGGLE_MODE = 'KAGGLE_KERNEL_RUN_TYPE' in os.environ\nDEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\nprint('='*60)\nprint('NullivaRNA-Flow V17 - Correct Bond Loss')\nprint('='*60)\nprint(f'PyTorch: {torch.__version__}')\nprint(f'Device: {DEVICE}')\nprint(f'Kaggle mode: {KAGGLE_MODE}')\n\nif torch.cuda.is_available():\n    print(f'GPU: {torch.cuda.get_device_name(0)}')\n    print(f'Memory: {torch.cuda.get_device_properties(0).total_memory / 1e9:.1f} GB')\n\nSTART_TIME = time.time()\nMAX_HOURS = 11.5\n\ndef time_remaining():\n    return MAX_HOURS - (time.time() - START_TIME) / 3600\n\ndef should_stop():\n    return time_remaining() < 0.25\n\n\ndef find_data_path():\n    possible_paths = [\n        '/kaggle/input/stanford-rna-3d-folding-2',\n        '/kaggle/input/stanford-rna-3d-folding-part-2',\n    ]\n    for pattern in ['/kaggle/input/*/train_sequences.csv', '/kaggle/input/*/*/train_sequences.csv']:\n        matches = glob.glob(pattern)\n        if matches:\n            return os.path.dirname(matches[0])\n    for p in possible_paths:\n        if os.path.exists(p):\n            return p\n    return None\n\n\nDATA_PATH = find_data_path()\nif DATA_PATH is None:\n    print('WARNING: Could not find competition data!')\n    DATA_PATH = '/kaggle/input/stanford-rna-3d-folding-2'\nelse:\n    print(f'Found data at: {DATA_PATH}')\n\n\n@dataclass\nclass TrainConfig:\n    \"\"\"V17 Training configuration.\"\"\"\n    train_seq_path: str = f'{DATA_PATH}/train_sequences.csv'\n    train_labels_path: str = f'{DATA_PATH}/train_labels.csv'\n    val_seq_path: str = f'{DATA_PATH}/validation_sequences.csv'\n    val_labels_path: str = f'{DATA_PATH}/validation_labels.csv'\n    \n    # Model\n    embed_dim: int = 256\n    n_layers: int = 4\n    max_seq_len: int = 512\n    dropout: float = 0.1\n    \n    # Training\n    epochs: int = 100\n    batch_size: int = 2\n    lr: float = 3e-4           # V17: Slightly lower for stability\n    min_lr: float = 1e-6\n    weight_decay: float = 0.01\n    warmup_pct: float = 0.02   # V17: 2% warmup\n    grad_clip: float = 1.0\n    \n    # V17: Corrected loss weights\n    fape_weight: float = 5.0\n    bond_weight: float = 0.5   # V17: MUCH lower - was dominating\n    dist_weight: float = 0.1   # V17: NEW - pairwise distance supervision\n    \n    # Geometry priors (RNA backbone in ANGSTROMS)\n    mu_bond: float = 5.9       # C1'-C1' typical distance\n    sig_bond: float = 1.5\n    \n    # Flow\n    num_flow_steps: int = 50\n    \n    # System\n    use_amp: bool = False\n    checkpoint_dir: str = '/kaggle/working/checkpoints'\n    log_every: int = 100\n    eval_every: int = 1\n    patience: int = 30\n    max_train_samples: int = 5000\n    max_val_samples: int = 50\n    min_valid_residues: int = 50\n\n\ncfg = TrainConfig()\n\nprint(f'\\n=== CONFIG V17 (CORRECT BOND LOSS) ===')\nprint(f'LR: {cfg.lr}')\nprint(f'Warmup: {cfg.warmup_pct*100:.0f}%')\nprint(f'Gradient Clipping: {cfg.grad_clip}')\nprint(f'FAPE weight: {cfg.fape_weight}')\nprint(f'Bond weight: {cfg.bond_weight} (in Angstroms)')\nprint(f'Dist weight: {cfg.dist_weight}')"},{"cell_type":"code","metadata":{},"execution_count":null,"outputs":[],"source":"# ==============================================================================\n# VOCABULARY"},{"cell_type":"code","metadata":{},"execution_count":null,"outputs":[],"source":"# ==============================================================================\n\nVOCAB = {'<PAD>': 0, '<UNK>': 1, 'A': 2, 'C': 3, 'G': 4, 'U': 5, 'N': 6}\nVOCAB_SIZE = len(VOCAB)\n\ndef tokenize_sequence(seq: str) -> torch.Tensor:\n    return torch.tensor([VOCAB.get(c.upper(), VOCAB['<UNK>']) for c in seq], dtype=torch.long)"},{"cell_type":"code","metadata":{},"execution_count":null,"outputs":[],"source":"# ==============================================================================\n# NORMALIZATION"},{"cell_type":"code","metadata":{},"execution_count":null,"outputs":[],"source":"# ==============================================================================\n\ndef normalize_coordinates_v17(coords, mask=None, eps=1e-6):\n    \"\"\"V17: Per-axis normalization.\"\"\"\n    if mask is not None:\n        valid_mask = mask.bool()\n        valid_coords = coords[valid_mask]\n    else:\n        valid_coords = coords\n    \n    n_valid = len(valid_coords)\n    if n_valid < 3:\n        return coords.clone(), torch.zeros(3, device=coords.device), torch.ones(3, device=coords.device)\n    \n    center = valid_coords.mean(dim=0)\n    centered = coords - center\n    \n    if mask is not None:\n        valid_centered = centered[valid_mask]\n    else:\n        valid_centered = centered\n    \n    scale = valid_centered.std(dim=0).clamp(min=eps)\n    scale = scale.clamp(min=1.0, max=100.0)\n    \n    normalized = centered / scale\n    normalized = normalized.clamp(-10.0, 10.0)\n    \n    return normalized, center, scale\n\n\ndef denormalize_coordinates_v17(normalized, center, scale):\n    \"\"\"V17: Convert back to Angstroms.\"\"\"\n    normalized = normalized.clamp(-10.0, 10.0)\n    \n    if normalized.dim() == 3 and center.dim() == 1:\n        return normalized * scale.view(1, 1, 3) + center.view(1, 1, 3)\n    elif normalized.dim() == 3 and center.dim() == 2:\n        return normalized * scale.unsqueeze(1) + center.unsqueeze(1)\n    else:\n        return normalized * scale + center\n\nprint('Normalization V17 ready')"},{"cell_type":"code","metadata":{},"execution_count":null,"outputs":[],"source":"# ==============================================================================\n# TIME EMBEDDING"},{"cell_type":"code","metadata":{},"execution_count":null,"outputs":[],"source":"# ==============================================================================\n\nclass SinusoidalTimeEmbedding(nn.Module):\n    def __init__(self, dim):\n        super().__init__()\n        self.dim = dim\n        self.mlp = nn.Sequential(\n            nn.Linear(dim, dim * 4),\n            nn.GELU(),\n            nn.Linear(dim * 4, dim),\n        )\n        for m in self.mlp.modules():\n            if isinstance(m, nn.Linear):\n                nn.init.xavier_uniform_(m.weight, gain=0.1)\n                nn.init.zeros_(m.bias)\n    \n    def forward(self, t):\n        half = self.dim // 2\n        freqs = torch.exp(-math.log(10000) * torch.arange(half, device=t.device, dtype=t.dtype) / half)\n        args = t.unsqueeze(-1) * freqs.unsqueeze(0)\n        emb = torch.cat([torch.sin(args), torch.cos(args)], dim=-1)\n        return self.mlp(emb)"},{"cell_type":"code","metadata":{},"execution_count":null,"outputs":[],"source":"# ==============================================================================\n# EGNN LAYER"},{"cell_type":"code","metadata":{},"execution_count":null,"outputs":[],"source":"# ==============================================================================\n\nclass EGNNLayerV17(nn.Module):\n    \"\"\"V17: EGNN with stability.\"\"\"\n    \n    def __init__(self, hidden_dim, dropout=0.1):\n        super().__init__()\n        self.hidden_dim = hidden_dim\n        \n        self.edge_mlp = nn.Sequential(\n            nn.Linear(hidden_dim * 2 + 1, hidden_dim),\n            nn.SiLU(),\n            nn.Dropout(dropout),\n            nn.Linear(hidden_dim, hidden_dim),\n            nn.SiLU(),\n        )\n        \n        self.node_mlp = nn.Sequential(\n            nn.Linear(hidden_dim * 2, hidden_dim),\n            nn.SiLU(),\n            nn.Dropout(dropout),\n            nn.Linear(hidden_dim, hidden_dim),\n        )\n        \n        self.coord_mlp = nn.Sequential(\n            nn.Linear(hidden_dim, hidden_dim),\n            nn.SiLU(),\n            nn.Dropout(dropout),\n            nn.Linear(hidden_dim, 1),\n        )\n        \n        self._init_weights()\n    \n    def _init_weights(self):\n        for m in self.modules():\n            if isinstance(m, nn.Linear):\n                nn.init.xavier_uniform_(m.weight, gain=0.5)\n                if m.bias is not None:\n                    nn.init.zeros_(m.bias)\n    \n    def forward(self, h, x, edge_index, mask=None):\n        row, col = edge_index\n        \n        x = x.clamp(-50.0, 50.0)\n        \n        diff = x[row] - x[col]\n        dist = (diff ** 2).sum(dim=-1, keepdim=True).clamp(min=1e-6).sqrt()\n        dist = dist.clamp(max=100.0)\n        \n        edge_input = torch.cat([h[row], h[col], dist], dim=-1)\n        edge_feat = self.edge_mlp(edge_input)\n        \n        agg = torch.zeros(h.shape, dtype=edge_feat.dtype, device=h.device)\n        agg.index_add_(0, row, edge_feat)\n        \n        node_input = torch.cat([h, agg], dim=-1)\n        h_out = h + self.node_mlp(node_input)\n        \n        coord_weights = self.coord_mlp(edge_feat).clamp(-1.0, 1.0)\n        weighted_diff = diff * coord_weights\n        coord_update = torch.zeros(x.shape, dtype=weighted_diff.dtype, device=x.device)\n        coord_update.index_add_(0, row, weighted_diff)\n        coord_update = coord_update.clamp(-5.0, 5.0)\n        x_out = x + coord_update\n        \n        if mask is not None:\n            mask_exp = mask.unsqueeze(-1).to(h_out.dtype)\n            h_out = h_out * mask_exp\n            x_out = x_out * mask_exp\n        \n        return h_out, x_out\n\nprint('EGNN V17 ready')"},{"cell_type":"code","metadata":{},"execution_count":null,"outputs":[],"source":"# ==============================================================================\n# FLOW BACKBONE"},{"cell_type":"code","metadata":{},"execution_count":null,"outputs":[],"source":"# ==============================================================================\n\nclass FlowBackboneV17(nn.Module):\n    def __init__(self, hidden_dim, n_layers=4, dropout=0.1, k_neighbors=16):\n        super().__init__()\n        self.hidden_dim = hidden_dim\n        self.n_layers = n_layers\n        self.k_neighbors = k_neighbors\n        \n        self.time_embed = SinusoidalTimeEmbedding(hidden_dim)\n        self.time_proj = nn.Linear(hidden_dim, hidden_dim)\n        \n        self.layers = nn.ModuleList([\n            EGNNLayerV17(hidden_dim, dropout=dropout) for _ in range(n_layers)\n        ])\n        \n        self.vel_head = nn.Sequential(\n            nn.Linear(hidden_dim, hidden_dim),\n            nn.SiLU(),\n            nn.Linear(hidden_dim, 3)\n        )\n        \n        nn.init.zeros_(self.vel_head[-1].weight)\n        nn.init.zeros_(self.vel_head[-1].bias)\n    \n    def _build_edges(self, N, device):\n        idx = torch.arange(N, device=device)\n        edges_src, edges_dst = [], []\n        \n        for offset in range(1, min(self.k_neighbors + 1, N)):\n            src = idx[:-offset]\n            dst = idx[offset:]\n            edges_src.extend([src, dst])\n            edges_dst.extend([dst, src])\n        \n        if not edges_src:\n            return torch.zeros(2, 0, dtype=torch.long, device=device)\n        \n        src = torch.cat(edges_src)\n        dst = torch.cat(edges_dst)\n        return torch.stack([src, dst], dim=0)\n    \n    def forward(self, h, x, t, mask=None):\n        N = h.shape[0]\n        device = h.device\n        \n        t_emb = self.time_embed(t.view(1))\n        t_proj = self.time_proj(t_emb)\n        h = h + t_proj.expand(N, -1)\n        \n        edge_index = self._build_edges(N, device)\n        x = x.clamp(-50.0, 50.0)\n        \n        for layer in self.layers:\n            h, x = layer(h, x, edge_index, mask)\n            h = h.clamp(-100.0, 100.0)\n            x = x.clamp(-50.0, 50.0)\n        \n        v = self.vel_head(h)\n        v = v.clamp(-10.0, 10.0)\n        \n        return v, h\n\nprint('Flow Backbone V17 ready')"},{"cell_type":"code","metadata":{},"execution_count":null,"outputs":[],"source":"# ==============================================================================\n# MAIN MODEL"},{"cell_type":"code","metadata":{},"execution_count":null,"outputs":[],"source":"# ==============================================================================\n\nclass NullivaRNAFlowV17(nn.Module):\n    def __init__(self, vocab_size=7, embed_dim=256, n_layers=4, max_seq_len=512, dropout=0.1):\n        super().__init__()\n        \n        self.embed = nn.Embedding(vocab_size, embed_dim, padding_idx=0)\n        self.pos_embed = nn.Embedding(max_seq_len, embed_dim)\n        \n        self.flow = FlowBackboneV17(embed_dim, n_layers=n_layers, dropout=dropout)\n        \n        nn.init.normal_(self.embed.weight, mean=0, std=0.02)\n        nn.init.normal_(self.pos_embed.weight, mean=0, std=0.02)\n    \n    def encode(self, tokens, mask=None):\n        B, N = tokens.shape\n        positions = torch.arange(N, device=tokens.device).unsqueeze(0).expand(B, -1)\n        h = self.embed(tokens) + self.pos_embed(positions)\n        if mask is not None:\n            h = h * mask.unsqueeze(-1)\n        return h\n    \n    def forward(self, x, t, tokens, mask=None):\n        B, N, _ = x.shape\n        \n        v_list = []\n        h_list = []\n        \n        for b in range(B):\n            h_b = self.encode(tokens[b:b+1], mask[b:b+1] if mask is not None else None)[0]\n            x_b = x[b]\n            t_b = t[b]\n            mask_b = mask[b] if mask is not None else None\n            \n            v_b, h_out_b = self.flow(h_b, x_b, t_b, mask_b)\n            v_list.append(v_b)\n            h_list.append(h_out_b)\n        \n        v = torch.stack(v_list, dim=0)\n        \n        return v, {'h': torch.stack(h_list, dim=0)}\n    \n    @torch.no_grad()\n    def sample(self, tokens, mask=None, num_steps=50):\n        B, N = tokens.shape\n        device = tokens.device\n        \n        x = torch.randn(B, N, 3, device=device).clamp(-3.0, 3.0)\n        \n        dt = 1.0 / num_steps\n        \n        for step in range(num_steps):\n            t = torch.full((B,), step * dt, device=device)\n            v, _ = self.forward(x, t, tokens, mask)\n            v = v.clamp(-10.0, 10.0)\n            x = x + dt * v\n            x = x.clamp(-20.0, 20.0)\n        \n        return x\n\nprint('Model V17 ready')"},{"cell_type":"code","metadata":{},"execution_count":null,"outputs":[],"source":"# ==============================================================================\n# TM-SCORE"},{"cell_type":"code","metadata":{},"execution_count":null,"outputs":[],"source":"# ==============================================================================\n\ndef compute_tm_score_v17(pred, true, mask=None, min_valid=50):\n    if mask is not None:\n        valid = mask.bool()\n        pred = pred[valid]\n        true = true[valid]\n    \n    n = len(pred)\n    if n < min_valid:\n        return 0.0, {'skipped': True, 'reason': f'n={n} < min_valid={min_valid}'}\n    \n    if torch.isnan(pred).any() or torch.isinf(pred).any():\n        return 0.0, {'skipped': True, 'reason': 'pred contains NaN/Inf'}\n    if torch.isnan(true).any() or torch.isinf(true).any():\n        return 0.0, {'skipped': True, 'reason': 'true contains NaN/Inf'}\n    \n    pred_centered = pred - pred.mean(dim=0)\n    true_centered = true - true.mean(dim=0)\n    \n    H = pred_centered.T @ true_centered\n    U, S, Vt = torch.linalg.svd(H)\n    \n    d = torch.det(Vt.T @ U.T)\n    sign_matrix = torch.diag(torch.tensor([1.0, 1.0, d.sign()], device=pred.device))\n    R = Vt.T @ sign_matrix @ U.T\n    \n    pred_aligned = pred_centered @ R\n    \n    dist = (pred_aligned - true_centered).norm(dim=-1)\n    dist = dist.clamp(max=1000.0)\n    \n    d0 = 1.24 * (n - 15) ** (1/3) - 1.8\n    d0 = max(d0, 0.5)\n    \n    tm = (1 / (1 + (dist / d0) ** 2)).mean().item()\n    \n    return tm, {'n_valid': n, 'mean_dist': dist.mean().item(), 'd0': d0}\n\nprint('TM-score V17 ready')"},{"cell_type":"code","metadata":{},"execution_count":null,"outputs":[],"source":"# ==============================================================================\n# DATASET"},{"cell_type":"code","metadata":{},"execution_count":null,"outputs":[],"source":"# ==============================================================================\n\nclass RNAFlowDatasetV17(Dataset):\n    def __init__(self, seq_df, labels_df, max_seq_len=512, max_samples=None, \n                 name='DATASET', normalize=True, min_valid_ratio=0.5):\n        self.max_seq_len = max_seq_len\n        self.normalize = normalize\n        self.min_valid_ratio = min_valid_ratio\n        self.name = name\n        self.samples = []\n        \n        self.diag = {\n            'total_sequences': len(seq_df),\n            'matched_labels': 0,\n            'dropped_len_gt_max': 0,\n            'dropped_low_valid_ratio': 0,\n            'valid_residues': 0,\n            'total_residues': 0,\n            'final_samples': 0,\n            'norm_stats': {'check_std': [], 'scales': []},\n        }\n        \n        print(f'  Processing labels ({len(labels_df)} rows)...')\n        labels_df = labels_df.copy()\n        labels_df['target'] = labels_df['ID'].str.rsplit('_', n=1).str[0]\n        labels_df['resid'] = labels_df['ID'].str.rsplit('_', n=1).str[1].astype(int)\n        grouped = {k: v for k, v in labels_df.groupby('target')}\n        print(f'  Unique targets in labels: {len(grouped)}')\n        \n        for idx, row in seq_df.iterrows():\n            if max_samples and len(self.samples) >= max_samples:\n                break\n            \n            target_id = row.get('target_id', row.get('sequence_id', str(idx)))\n            sequence = row.get('sequence', '')\n            \n            if len(sequence) > max_seq_len:\n                self.diag['dropped_len_gt_max'] += 1\n                continue\n            \n            if target_id not in grouped:\n                continue\n            \n            self.diag['matched_labels'] += 1\n            target_labels = grouped[target_id].sort_values('resid')\n            \n            try:\n                result = self._parse_and_normalize_coords(target_labels, len(sequence))\n                if result is None:\n                    continue\n                \n                coords, coord_mask, center, scale, raw_coords = result\n                \n                n_valid = coord_mask.sum().item()\n                self.diag['valid_residues'] += n_valid\n                self.diag['total_residues'] += len(sequence)\n                \n                if n_valid / len(sequence) < self.min_valid_ratio:\n                    self.diag['dropped_low_valid_ratio'] += 1\n                    continue\n                \n                if torch.isnan(coords).any() or torch.isnan(raw_coords).any():\n                    continue\n                \n                if self.normalize:\n                    check_std = coords[coord_mask.bool()].std(dim=0).mean().item()\n                    self.diag['norm_stats']['check_std'].append(check_std)\n                    self.diag['norm_stats']['scales'].append(scale.mean().item())\n                \n                self.samples.append({\n                    'target_id': target_id,\n                    'sequence': sequence,\n                    'coords': coords,\n                    'coord_mask': coord_mask,\n                    'center': center,\n                    'scale': scale,\n                    'raw_coords': raw_coords,\n                })\n            except Exception as e:\n                continue\n        \n        self.diag['final_samples'] = len(self.samples)\n        self._print_diagnostic()\n    \n    def _parse_and_normalize_coords(self, labels_df, seq_len):\n        if len(labels_df) != seq_len:\n            return None\n        \n        coords = []\n        mask = []\n        \n        for _, row in labels_df.iterrows():\n            x = row.get('x_1', row.get('x', None))\n            y = row.get('y_1', row.get('y', None))\n            z = row.get('z_1', row.get('z', None))\n            \n            if pd.isna(x) or pd.isna(y) or pd.isna(z):\n                coords.append([0.0, 0.0, 0.0])\n                mask.append(0)\n            elif abs(float(x)) > 1e6 or abs(float(y)) > 1e6 or abs(float(z)) > 1e6:\n                coords.append([0.0, 0.0, 0.0])\n                mask.append(0)\n            else:\n                coords.append([float(x), float(y), float(z)])\n                mask.append(1)\n        \n        raw_coords = torch.tensor(coords, dtype=torch.float32)\n        coord_mask = torch.tensor(mask, dtype=torch.float32)\n        \n        if self.normalize:\n            normalized, center, scale = normalize_coordinates_v17(raw_coords, coord_mask)\n            return (normalized, coord_mask, center, scale, raw_coords)\n        else:\n            center = raw_coords[coord_mask.bool()].mean(dim=0) if coord_mask.sum() > 0 else torch.zeros(3)\n            scale = torch.ones(3)\n            return (raw_coords, coord_mask, center, scale, raw_coords)\n    \n    def _print_diagnostic(self):\n        d = self.diag\n        print(f'\\n=== [{self.name}] Dataset Diagnostic V17 ===')\n        print(f'  Total sequences:       {d[\"total_sequences\"]}')\n        print(f'  Matched with labels:   {d[\"matched_labels\"]}')\n        print(f'  Dropped (len > max):   {d[\"dropped_len_gt_max\"]}')\n        print(f'  Dropped (low valid):   {d[\"dropped_low_valid_ratio\"]}')\n        if d['total_residues'] > 0:\n            valid_pct = d['valid_residues']/d['total_residues']*100\n            print(f'  Valid residues:        {d[\"valid_residues\"]}/{d[\"total_residues\"]} ({valid_pct:.1f}%)')\n        if self.diag['norm_stats']['check_std']:\n            mean_std = np.mean(self.diag['norm_stats']['check_std'])\n            mean_scale = np.mean(self.diag['norm_stats']['scales'])\n            print(f'  Norm check: std={mean_std:.3f}, scale={mean_scale:.2f}Å')\n        print(f'  FINAL SAMPLES: {d[\"final_samples\"]}')\n    \n    def __len__(self):\n        return len(self.samples)\n    \n    def __getitem__(self, idx):\n        sample = self.samples[idx]\n        tokens = tokenize_sequence(sample['sequence'])\n        \n        return {\n            'target_id': sample['target_id'],\n            'tokens': tokens,\n            'coords': sample['coords'],\n            'coord_mask': sample['coord_mask'],\n            'seq_mask': torch.ones(len(tokens)),\n            'center': sample['center'],\n            'scale': sample['scale'],\n            'raw_coords': sample['raw_coords'],\n        }\n\n\ndef collate_fn_v17(batch):\n    max_len = max(len(item['tokens']) for item in batch)\n    \n    tokens = torch.zeros(len(batch), max_len, dtype=torch.long)\n    coords = torch.zeros(len(batch), max_len, 3)\n    raw_coords = torch.zeros(len(batch), max_len, 3)\n    coord_mask = torch.zeros(len(batch), max_len)\n    seq_mask = torch.zeros(len(batch), max_len)\n    centers = torch.zeros(len(batch), 3)\n    scales = torch.zeros(len(batch), 3)\n    target_ids = []\n    \n    for i, item in enumerate(batch):\n        L = len(item['tokens'])\n        tokens[i, :L] = item['tokens']\n        coords[i, :L] = item['coords']\n        raw_coords[i, :L] = item['raw_coords']\n        coord_mask[i, :L] = item['coord_mask']\n        seq_mask[i, :L] = item['seq_mask']\n        centers[i] = item['center']\n        scales[i] = item['scale']\n        target_ids.append(item['target_id'])\n    \n    return {\n        'target_ids': target_ids,\n        'tokens': tokens,\n        'coords': coords,\n        'raw_coords': raw_coords,\n        'coord_mask': coord_mask,\n        'seq_mask': seq_mask,\n        'centers': centers,\n        'scales': scales,\n    }\n\nprint('Dataset V17 ready')"},{"cell_type":"code","metadata":{},"execution_count":null,"outputs":[],"source":"# ==============================================================================\n# LOSS FUNCTIONS V17"},{"cell_type":"code","metadata":{},"execution_count":null,"outputs":[],"source":"# ==============================================================================\n\ndef compute_fape_loss_v17(pred_coords, true_coords, mask, eps=1e-6):\n    \"\"\"FAPE loss with NaN protection.\"\"\"\n    diff = pred_coords - true_coords\n    dist = (diff ** 2).sum(dim=-1).clamp(min=eps).sqrt()\n    \n    d_clamp = 10.0\n    fape = torch.clamp(dist, max=d_clamp)\n    \n    if torch.isnan(fape).any():\n        return torch.tensor(0.0, device=pred_coords.device, requires_grad=True)\n    \n    if mask is not None:\n        fape = (fape * mask).sum() / (mask.sum() + eps)\n    else:\n        fape = fape.mean()\n    \n    return fape\n\n\ndef compute_bond_loss_v17(pred_coords_raw, true_coords_raw, mask, mu_bond=5.9, sig_bond=1.5):\n    \"\"\"V17: Bond distance loss in RAW SPACE (Angstroms).\n    \n    This computes how well the predicted backbone matches the true backbone\n    in terms of consecutive residue distances.\n    \n    Args:\n        pred_coords_raw: [B, N, 3] predicted coords in Angstroms\n        true_coords_raw: [B, N, 3] true coords in Angstroms\n        mask: [B, N] validity mask\n        mu_bond: target C1'-C1' distance (~5.9Å for RNA)\n        sig_bond: allowed variation\n    \n    Returns:\n        L_bond: scalar loss\n    \"\"\"\n    # Consecutive distances for predicted\n    pred_bonds = pred_coords_raw[:, 1:] - pred_coords_raw[:, :-1]  # [B, N-1, 3]\n    pred_bond_dist = pred_bonds.norm(dim=-1)  # [B, N-1]\n    \n    # Consecutive distances for true (as reference)\n    true_bonds = true_coords_raw[:, 1:] - true_coords_raw[:, :-1]\n    true_bond_dist = true_bonds.norm(dim=-1)\n    \n    # Mask for valid consecutive pairs\n    if mask is not None:\n        pair_mask = mask[:, 1:] * mask[:, :-1]\n    else:\n        pair_mask = torch.ones_like(pred_bond_dist)\n    \n    # V17: Loss = how much predicted differs from true backbone geometry\n    # Plus penalty if too far from typical RNA backbone (~5.9Å)\n    bond_diff = (pred_bond_dist - true_bond_dist).abs()\n    \n    # Also penalize if predicted is too far from typical RNA backbone\n    deviation_from_typical = ((pred_bond_dist - mu_bond) / sig_bond).clamp(-5, 5) ** 2\n    \n    # Combined: match true + be reasonable\n    bond_loss = bond_diff + 0.1 * deviation_from_typical\n    \n    if mask is not None:\n        bond_loss = (bond_loss * pair_mask).sum() / (pair_mask.sum() + 1e-6)\n    else:\n        bond_loss = bond_loss.mean()\n    \n    return bond_loss\n\n\ndef compute_distance_matrix_loss_v17(pred_coords_raw, true_coords_raw, mask, max_pairs=1000):\n    \"\"\"V17: Pairwise distance matrix loss.\n    \n    Encourages the overall shape to match by comparing pairwise distances.\n    Uses sampling to avoid O(N^2) computation.\n    \"\"\"\n    B, N, _ = pred_coords_raw.shape\n    \n    total_loss = 0.0\n    count = 0\n    \n    for b in range(B):\n        if mask is not None:\n            valid_idx = mask[b].bool().nonzero(as_tuple=True)[0]\n        else:\n            valid_idx = torch.arange(N, device=pred_coords_raw.device)\n        \n        n_valid = len(valid_idx)\n        if n_valid < 10:\n            continue\n        \n        # Sample pairs if too many\n        if n_valid * (n_valid - 1) // 2 > max_pairs:\n            n_sample = int(math.sqrt(2 * max_pairs))\n            sample_idx = valid_idx[torch.randperm(n_valid, device=valid_idx.device)[:n_sample]]\n        else:\n            sample_idx = valid_idx\n        \n        n_s = len(sample_idx)\n        \n        # Compute pairwise distances\n        pred_sub = pred_coords_raw[b, sample_idx]  # [n_s, 3]\n        true_sub = true_coords_raw[b, sample_idx]\n        \n        # [n_s, n_s]\n        pred_dist = torch.cdist(pred_sub, pred_sub)\n        true_dist = torch.cdist(true_sub, true_sub)\n        \n        # Upper triangle only (exclude diagonal)\n        triu_idx = torch.triu_indices(n_s, n_s, offset=1, device=pred_dist.device)\n        pred_flat = pred_dist[triu_idx[0], triu_idx[1]]\n        true_flat = true_dist[triu_idx[0], triu_idx[1]]\n        \n        # L1 loss on distances\n        dist_loss = (pred_flat - true_flat).abs().mean()\n        \n        total_loss += dist_loss\n        count += 1\n    \n    if count == 0:\n        return torch.tensor(0.0, device=pred_coords_raw.device, requires_grad=True)\n    \n    return total_loss / count\n\n\ndef compute_phi_v17(coords, mask, eps=1e-6):\n    \"\"\"V17: Compute normalized variance (Phi) for monitoring.\"\"\"\n    if mask is not None:\n        B = coords.shape[0]\n        stds = []\n        for b in range(B):\n            valid = mask[b].bool()\n            if valid.sum() >= 2:\n                valid_coords = coords[b][valid]\n                std = valid_coords.std(dim=0).mean().item()\n                stds.append(std)\n        \n        if stds:\n            return np.mean(stds)\n        return 0.0\n    else:\n        return coords.std(dim=1).mean(dim=-1).mean().item()\n\n\nclass LossV17(nn.Module):\n    \"\"\"V17: Loss function with correct bond loss in Angstroms.\"\"\"\n    \n    def __init__(self, cfg):\n        super().__init__()\n        self.cfg = cfg\n    \n    def forward(self, v_pred, v_true, pred_coords_norm, true_coords_norm, \n                coord_mask, centers, scales, raw_coords):\n        eps = 1e-6\n        \n        # 1. Flow loss (in normalized space)\n        flow_diff = (v_pred - v_true) ** 2\n        flow_diff = flow_diff.sum(dim=-1)\n        \n        if torch.isnan(flow_diff).any():\n            return torch.tensor(0.0, device=v_pred.device, requires_grad=True), {\n                'L_flow': float('nan'), 'L_fape': float('nan'), 'L_bond': float('nan'),\n                'L_dist': float('nan'), 'total': float('nan'), 'Phi': float('nan'),\n                'mean_bond_dist': float('nan'),\n            }\n        \n        if coord_mask is not None:\n            L_flow = (flow_diff * coord_mask).sum() / (coord_mask.sum() + eps)\n        else:\n            L_flow = flow_diff.mean()\n        \n        # 2. FAPE loss (in normalized space)\n        L_fape = compute_fape_loss_v17(pred_coords_norm, true_coords_norm, coord_mask)\n        \n        if torch.isnan(L_fape):\n            L_fape = torch.tensor(0.0, device=v_pred.device, requires_grad=True)\n        \n        # 3. V17: Denormalize to get Angstrom coords for geometry losses\n        pred_coords_raw = denormalize_coordinates_v17(pred_coords_norm, centers, scales)\n        \n        # 4. V17: Bond loss in Angstroms\n        L_bond = compute_bond_loss_v17(\n            pred_coords_raw, raw_coords, coord_mask,\n            mu_bond=self.cfg.mu_bond,\n            sig_bond=self.cfg.sig_bond\n        )\n        \n        if torch.isnan(L_bond):\n            L_bond = torch.tensor(0.0, device=v_pred.device, requires_grad=True)\n        \n        # 5. V17: Distance matrix loss\n        L_dist = compute_distance_matrix_loss_v17(pred_coords_raw, raw_coords, coord_mask)\n        \n        if torch.isnan(L_dist):\n            L_dist = torch.tensor(0.0, device=v_pred.device, requires_grad=True)\n        \n        # Combined loss\n        total = (L_flow + \n                 self.cfg.fape_weight * L_fape + \n                 self.cfg.bond_weight * L_bond +\n                 self.cfg.dist_weight * L_dist)\n        \n        if torch.isnan(total):\n            total = torch.tensor(0.0, device=v_pred.device, requires_grad=True)\n        \n        # Compute Phi for monitoring\n        phi = compute_phi_v17(pred_coords_norm, coord_mask)\n        \n        # Mean bond distance in Angstroms\n        bonds = (pred_coords_raw[:, 1:] - pred_coords_raw[:, :-1]).norm(dim=-1)\n        mean_bond = bonds.mean().item() if not torch.isnan(bonds).any() else 0.0\n        \n        metrics = {\n            'L_flow': L_flow.item() if not torch.isnan(L_flow) else float('nan'),\n            'L_fape': L_fape.item() if not torch.isnan(L_fape) else float('nan'),\n            'L_bond': L_bond.item() if not torch.isnan(L_bond) else float('nan'),\n            'L_dist': L_dist.item() if not torch.isnan(L_dist) else float('nan'),\n            'total': total.item() if not torch.isnan(total) else float('nan'),\n            'Phi': phi,\n            'mean_bond_dist': mean_bond,\n        }\n        \n        return total, metrics\n\nprint('V17 Loss with correct bond loss in Angstroms')\nprint(f'  FAPE weight: {cfg.fape_weight}')\nprint(f'  Bond weight: {cfg.bond_weight}')\nprint(f'  Dist weight: {cfg.dist_weight}')"},{"cell_type":"code","metadata":{},"execution_count":null,"outputs":[],"source":"# ==============================================================================\n# TRAINING FUNCTIONS"},{"cell_type":"code","metadata":{},"execution_count":null,"outputs":[],"source":"# ==============================================================================\n\ndef train_step_v17(model, batch, loss_fn, device):\n    \"\"\"V17: Training step.\"\"\"\n    tokens = batch['tokens'].to(device)\n    coords = batch['coords'].to(device)\n    coord_mask = batch['coord_mask'].to(device)\n    centers = batch['centers'].to(device)\n    scales = batch['scales'].to(device)\n    raw_coords = batch['raw_coords'].to(device)\n    \n    B, N, _ = coords.shape\n    \n    if torch.isnan(coords).any():\n        return None, {'skipped': True}\n    \n    x_0 = torch.randn(B, N, 3, device=device).clamp(-3.0, 3.0)\n    x_1 = coords\n    \n    t = torch.rand(B, device=device)\n    t_exp = t.view(B, 1, 1)\n    \n    x_t = (1 - t_exp) * x_0 + t_exp * x_1\n    v_true = x_1 - x_0\n    \n    v_pred, _ = model(x_t, t, tokens, coord_mask)\n    \n    if torch.isnan(v_pred).any():\n        return None, {'skipped': True}\n    \n    pred_x1 = x_t + (1 - t_exp) * v_pred\n    \n    loss, metrics = loss_fn(v_pred, v_true, pred_x1, x_1, coord_mask, centers, scales, raw_coords)\n    \n    if torch.isnan(loss):\n        return None, {'skipped': True}\n    \n    return loss, metrics\n\n\n@torch.no_grad()\ndef validate_v17(model, val_loader, loss_fn, device):\n    model.eval()\n    total_loss = 0\n    total_metrics = {}\n    n_batches = 0\n    \n    for batch in val_loader:\n        tokens = batch['tokens'].to(device)\n        coords = batch['coords'].to(device)\n        coord_mask = batch['coord_mask'].to(device)\n        centers = batch['centers'].to(device)\n        scales = batch['scales'].to(device)\n        raw_coords = batch['raw_coords'].to(device)\n        \n        B, N, _ = coords.shape\n        \n        x_0 = torch.randn(B, N, 3, device=device).clamp(-3.0, 3.0)\n        x_1 = coords\n        \n        t = torch.rand(B, device=device)\n        t_exp = t.view(B, 1, 1)\n        x_t = (1 - t_exp) * x_0 + t_exp * x_1\n        v_true = x_1 - x_0\n        \n        v_pred, _ = model(x_t, t, tokens, coord_mask)\n        pred_x1 = x_t + (1 - t_exp) * v_pred\n        \n        loss, metrics = loss_fn(v_pred, v_true, pred_x1, x_1, coord_mask, centers, scales, raw_coords)\n        \n        if not torch.isnan(loss):\n            total_loss += loss.item()\n            for k, v in metrics.items():\n                if not (isinstance(v, float) and math.isnan(v)):\n                    total_metrics[k] = total_metrics.get(k, 0) + v\n            n_batches += 1\n    \n    model.train()\n    \n    if n_batches == 0:\n        return float('inf'), {}\n    \n    return total_loss / n_batches, {k: v / n_batches for k, v in total_metrics.items()}\n\n\n@torch.no_grad()\ndef evaluate_tm_score_v17(model, val_loader, device, num_samples=10, num_flow_steps=50, min_valid=50):\n    model.eval()\n    tm_scores = []\n    \n    for i, batch in enumerate(val_loader):\n        if i >= num_samples:\n            break\n        \n        tokens = batch['tokens'].to(device)\n        raw_coords = batch['raw_coords'].to(device)\n        seq_mask = batch['seq_mask'].to(device)\n        coord_mask = batch['coord_mask'].to(device)\n        centers = batch['centers'].to(device)\n        scales = batch['scales'].to(device)\n        target_ids = batch['target_ids']\n        \n        pred_normalized = model.sample(tokens, seq_mask, num_steps=num_flow_steps)\n        pred_raw = denormalize_coordinates_v17(pred_normalized, centers, scales)\n        \n        for b in range(len(target_ids)):\n            tm, debug = compute_tm_score_v17(\n                pred_raw[b], raw_coords[b], coord_mask[b], min_valid=min_valid\n            )\n            \n            if debug.get('skipped', False):\n                print(f'    {target_ids[b]}: SKIPPED ({debug.get(\"reason\", \"\")})')\n            else:\n                print(f'    {target_ids[b]}: TM={tm:.4f} | valid={debug[\"n_valid\"]} | dist={debug[\"mean_dist\"]:.2f}Å')\n                tm_scores.append(tm)\n    \n    model.train()\n    \n    if not tm_scores:\n        return 0.0\n    \n    mean_tm = np.mean(tm_scores)\n    std_tm = np.std(tm_scores)\n    print(f'  [TM] Mean TM-score: {mean_tm:.4f} ± {std_tm:.4f}')\n    \n    return mean_tm\n\nprint('Training functions V17 ready')"},{"cell_type":"code","metadata":{},"execution_count":null,"outputs":[],"source":"# ==============================================================================\n# MAIN"},{"cell_type":"code","metadata":{},"execution_count":null,"outputs":[],"source":"# ==============================================================================\n\ndef main():\n    print('\\n' + '='*60)\n    print('STARTING V17 TRAINING')\n    print('='*60)\n    \n    print('Loading data...')\n    \n    train_seq_df = pd.read_csv(cfg.train_seq_path)\n    train_labels_df = pd.read_csv(cfg.train_labels_path, low_memory=False)\n    print(f'Train: {len(train_seq_df)} sequences, {len(train_labels_df)} label rows')\n    \n    val_seq_df = pd.read_csv(cfg.val_seq_path)\n    val_labels_df = pd.read_csv(cfg.val_labels_path, low_memory=False)\n    print(f'Val: {len(val_seq_df)} sequences, {len(val_labels_df)} label rows')\n    \n    print('\\nBuilding datasets...')\n    train_dataset = RNAFlowDatasetV17(\n        train_seq_df, train_labels_df,\n        max_seq_len=cfg.max_seq_len,\n        max_samples=cfg.max_train_samples,\n        name='TRAIN',\n        normalize=True,\n    )\n    \n    val_dataset = RNAFlowDatasetV17(\n        val_seq_df, val_labels_df,\n        max_seq_len=cfg.max_seq_len,\n        max_samples=cfg.max_val_samples,\n        name='VAL',\n        normalize=True,\n    )\n    \n    train_loader = DataLoader(\n        train_dataset, batch_size=cfg.batch_size, shuffle=True,\n        collate_fn=collate_fn_v17, num_workers=0, pin_memory=True\n    )\n    \n    val_loader = DataLoader(\n        val_dataset, batch_size=cfg.batch_size, shuffle=False,\n        collate_fn=collate_fn_v17, num_workers=0, pin_memory=True\n    )\n    \n    print(f'\\nTrain batches: {len(train_loader)}')\n    print(f'Val batches: {len(val_loader)}')\n    \n    # Model\n    model = NullivaRNAFlowV17(\n        vocab_size=VOCAB_SIZE,\n        embed_dim=cfg.embed_dim,\n        n_layers=cfg.n_layers,\n        max_seq_len=cfg.max_seq_len,\n        dropout=cfg.dropout,\n    ).to(DEVICE)\n    \n    total_params = sum(p.numel() for p in model.parameters())\n    print(f'\\nModel V17: {total_params:,} params')\n    \n    loss_fn = LossV17(cfg)\n    optimizer = torch.optim.AdamW(model.parameters(), lr=cfg.lr, weight_decay=cfg.weight_decay)\n    \n    total_steps = len(train_loader) * cfg.epochs\n    warmup_steps = int(total_steps * cfg.warmup_pct)\n    \n    print(f'\\nOptimizer: AdamW, LR={cfg.lr}')\n    print(f'Warmup: {warmup_steps} steps ({cfg.warmup_pct*100:.0f}%)')\n    print(f'Total steps: {total_steps}')\n    \n    def lr_lambda(step):\n        if step < warmup_steps:\n            return step / max(1, warmup_steps)\n        progress = (step - warmup_steps) / max(1, total_steps - warmup_steps)\n        return 0.5 * (1 + math.cos(math.pi * progress))\n    \n    scheduler = torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)\n    \n    Path(cfg.checkpoint_dir).mkdir(parents=True, exist_ok=True)\n    \n    # Training loop\n    print('\\n' + '='*60)\n    print('TRAINING V17')\n    print('='*60)\n    \n    best_val_loss = float('inf')\n    best_tm_score = 0.0\n    patience_counter = 0\n    global_step = 0\n    nan_count = 0\n    \n    for epoch in range(1, cfg.epochs + 1):\n        if should_stop():\n            print(f'\\n[!] Time limit. Stopping.')\n            break\n        \n        model.train()\n        epoch_losses = []\n        epoch_phi = []\n        epoch_bond = []\n        \n        for batch_idx, batch in enumerate(train_loader):\n            global_step += 1\n            \n            optimizer.zero_grad()\n            \n            loss, metrics = train_step_v17(model, batch, loss_fn, DEVICE)\n            \n            if loss is None:\n                nan_count += 1\n                continue\n            \n            loss.backward()\n            \n            torch.nn.utils.clip_grad_norm_(model.parameters(), cfg.grad_clip)\n            \n            optimizer.step()\n            scheduler.step()\n            \n            epoch_losses.append(loss.item())\n            if 'Phi' in metrics and not math.isnan(metrics['Phi']):\n                epoch_phi.append(metrics['Phi'])\n            if 'mean_bond_dist' in metrics and not math.isnan(metrics['mean_bond_dist']):\n                epoch_bond.append(metrics['mean_bond_dist'])\n            \n            if global_step % cfg.log_every == 0:\n                avg_loss = np.mean(epoch_losses[-100:]) if epoch_losses else 0\n                avg_phi = np.mean(epoch_phi[-100:]) if epoch_phi else 0\n                avg_bond = np.mean(epoch_bond[-100:]) if epoch_bond else 0\n                lr = scheduler.get_last_lr()[0]\n                print(f'[E{epoch}] Step {global_step} | Loss: {avg_loss:.2f} | Φ: {avg_phi:.3f} | Bond: {avg_bond:.1f}Å | LR: {lr:.2e} | NaN: {nan_count}')\n        \n        # Validation\n        val_loss, val_metrics = validate_v17(model, val_loader, loss_fn, DEVICE)\n        \n        avg_train_loss = np.mean(epoch_losses) if epoch_losses else 0\n        avg_phi = np.mean(epoch_phi) if epoch_phi else 0\n        \n        print(f'\\n[=] Epoch {epoch}/{cfg.epochs} | Train: {avg_train_loss:.2f} | Val: {val_loss:.2f} | Φ: {avg_phi:.3f}')\n        print(f'    L_flow: {val_metrics.get(\"L_flow\", 0):.2f} | L_fape: {val_metrics.get(\"L_fape\", 0):.2f} | L_bond: {val_metrics.get(\"L_bond\", 0):.2f} | L_dist: {val_metrics.get(\"L_dist\", 0):.2f}')\n        print(f'    Mean bond dist: {val_metrics.get(\"mean_bond_dist\", 0):.2f}Å (target: ~{cfg.mu_bond}Å)')\n        \n        if val_loss < best_val_loss:\n            best_val_loss = val_loss\n            patience_counter = 0\n            torch.save({\n                'epoch': epoch,\n                'model_state_dict': model.state_dict(),\n                'optimizer_state_dict': optimizer.state_dict(),\n                'val_loss': val_loss,\n            }, f'{cfg.checkpoint_dir}/best_model.pt')\n            print(f'    [*] New best! Val Loss: {val_loss:.4f}')\n        else:\n            patience_counter += 1\n            print(f'    No improvement. Patience: {patience_counter}/{cfg.patience}')\n        \n        # TM-score\n        if epoch % cfg.eval_every == 0:\n            print(f'  [TM] Evaluating...')\n            try:\n                checkpoint = torch.load(f'{cfg.checkpoint_dir}/best_model.pt', weights_only=True)\n                model.load_state_dict(checkpoint['model_state_dict'])\n            except:\n                pass\n            \n            tm_score = evaluate_tm_score_v17(model, val_loader, DEVICE, \n                                             num_samples=10, \n                                             num_flow_steps=cfg.num_flow_steps,\n                                             min_valid=cfg.min_valid_residues)\n            \n            if tm_score > best_tm_score:\n                best_tm_score = tm_score\n                print(f'  [TM] New best: {best_tm_score:.4f}')\n                torch.save({\n                    'epoch': epoch,\n                    'model_state_dict': model.state_dict(),\n                    'tm_score': tm_score,\n                }, f'{cfg.checkpoint_dir}/best_tm_model.pt')\n        \n        if patience_counter >= cfg.patience:\n            print(f'\\n[!] Early stopping at epoch {epoch}')\n            break\n    \n    print('\\n' + '='*60)\n    print('TRAINING COMPLETE')\n    print('='*60)\n    print(f'Best Val Loss: {best_val_loss:.4f}')\n    print(f'Best TM-score: {best_tm_score:.4f}')\n    print(f'Total NaN: {nan_count}')\n    \n    # Final eval\n    print('\\n=== FINAL EVALUATION ===')\n    \n    try:\n        checkpoint = torch.load(f'{cfg.checkpoint_dir}/best_tm_model.pt', weights_only=True)\n        model.load_state_dict(checkpoint['model_state_dict'])\n        print(f'Loaded best TM model')\n    except:\n        try:\n            checkpoint = torch.load(f'{cfg.checkpoint_dir}/best_model.pt', weights_only=True)\n            model.load_state_dict(checkpoint['model_state_dict'])\n            print(f'Loaded best loss model')\n        except:\n            print('Using current weights')\n    \n    tm_score = evaluate_tm_score_v17(model, val_loader, DEVICE, \n                                      num_samples=50, \n                                      num_flow_steps=cfg.num_flow_steps,\n                                      min_valid=cfg.min_valid_residues)\n    \n    print(f'\\n=== FINAL ===')\n    print(f'Mean TM-score: {tm_score:.4f}')\n    print(f'Target: >= 0.5')\n    \n    if tm_score >= 0.5:\n        print('\\n[SUCCESS] Target achieved!')\n    else:\n        print(f'\\n[PROGRESS] Need improvement')\n\n\nif __name__ == '__main__':\n    main()"}]}