{"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 V16 Training\n\n**Stanford RNA 3D Folding Part 2 - Kaggle Competition**\n\n## V16 Key Changes from V15\n\n1. **Much shorter warmup**: 1% instead of 10% (~1862 steps vs 18620)\n2. **Higher LR**: 5e-4 instead of 1e-4 (5x faster learning)\n3. **Bond distance loss**: Encourages proper backbone geometry (~5.9Å between residues)\n4. **Phi monitoring**: Re-enabled for better tracking\n5. **Gradient clipping**: Still 1.0 for stability\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 V16 - Faster Learning')\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    \"\"\"V16 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    # V16: Faster learning\n    epochs: int = 100\n    batch_size: int = 2\n    lr: float = 5e-4           # V16: 5x higher than V15\n    min_lr: float = 1e-6\n    weight_decay: float = 0.01\n    warmup_pct: float = 0.01   # V16: MUCH shorter warmup (1% = ~1862 steps)\n    grad_clip: float = 1.0\n    \n    # V16: Loss weights with bond distance\n    fape_weight: float = 5.0\n    bond_weight: float = 2.0   # V16: NEW - encourage proper backbone\n    nv_weight: float = 0.0     # Keep disabled for now\n    \n    # Geometry priors (RNA backbone)\n    mu_bond: float = 5.9       # C1'-C1' typical distance\n    sig_bond: float = 1.5      # Allow some variation\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 V16 (FASTER LEARNING) ===')\nprint(f'LR: {cfg.lr} (5x higher than V15)')\nprint(f'Warmup: {cfg.warmup_pct*100:.0f}% (~{int(1862 * cfg.warmup_pct * 100)} steps)')\nprint(f'Gradient Clipping: {cfg.grad_clip}')\nprint(f'FAPE weight: {cfg.fape_weight}')\nprint(f'Bond weight: {cfg.bond_weight} (NEW)')"},{"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_v16(coords, mask=None, eps=1e-6):\n    \"\"\"V16: 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_v16(normalized, center, scale):\n    \"\"\"V16: 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 V16 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 EGNNLayerV16(nn.Module):\n    \"\"\"V16: 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 V16 ready')"},{"cell_type":"code","metadata":{},"execution_count":null,"outputs":[],"source":"# ==============================================================================\n# FLOW BACKBONE"},{"cell_type":"code","metadata":{},"execution_count":null,"outputs":[],"source":"# ==============================================================================\n\nclass FlowBackboneV16(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            EGNNLayerV16(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 V16 ready')"},{"cell_type":"code","metadata":{},"execution_count":null,"outputs":[],"source":"# ==============================================================================\n# MAIN MODEL"},{"cell_type":"code","metadata":{},"execution_count":null,"outputs":[],"source":"# ==============================================================================\n\nclass NullivaRNAFlowV16(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 = FlowBackboneV16(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 V16 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_v16(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 V16 ready')"},{"cell_type":"code","metadata":{},"execution_count":null,"outputs":[],"source":"# ==============================================================================\n# DATASET"},{"cell_type":"code","metadata":{},"execution_count":null,"outputs":[],"source":"# ==============================================================================\n\nclass RNAFlowDatasetV16(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_v16(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 V16 ===')\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_v16(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 V16 ready')"},{"cell_type":"code","metadata":{},"execution_count":null,"outputs":[],"source":"# ==============================================================================\n# LOSS FUNCTIONS V16"},{"cell_type":"code","metadata":{},"execution_count":null,"outputs":[],"source":"# ==============================================================================\n\ndef compute_fape_loss_v16(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_v16(pred_coords, mask, mu_bond=5.9, sig_bond=1.5, scale=None):\n    \"\"\"V16: Bond distance loss to encourage proper backbone geometry.\n    \n    In normalized space, we need to account for scale.\n    \"\"\"\n    # Consecutive residue distances\n    bonds = pred_coords[:, 1:] - pred_coords[:, :-1]  # [B, N-1, 3]\n    bond_dist = bonds.norm(dim=-1)  # [B, N-1]\n    \n    # If we have scale, convert target to normalized space\n    if scale is not None:\n        # Average scale across axes and batch\n        avg_scale = scale.mean(dim=-1, keepdim=True).mean(dim=0)  # scalar-ish\n        target_dist = mu_bond / avg_scale.clamp(min=1.0)\n        sigma = sig_bond / avg_scale.clamp(min=1.0)\n    else:\n        target_dist = mu_bond\n        sigma = sig_bond\n    \n    # Gaussian penalty\n    bond_loss = ((bond_dist - target_dist) ** 2) / (2 * sigma ** 2)\n    \n    # Mask for valid consecutive pairs\n    if mask is not None:\n        # Both residues must be valid\n        pair_mask = mask[:, 1:] * mask[:, :-1]\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_phi_v16(coords, mask, eps=1e-6):\n    \"\"\"V16: Compute normalized variance (Phi) for monitoring.\"\"\"\n    # Per-sequence std\n    if mask is not None:\n        # Compute std only for valid residues\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 LossV16(nn.Module):\n    \"\"\"V16: Loss function with bond distance term.\"\"\"\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):\n        eps = 1e-6\n        \n        # 1. Flow loss\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                'total': float('nan'), 'Phi': float('nan'), '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\n        L_fape = compute_fape_loss_v16(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. V16: Bond distance loss\n        L_bond = compute_bond_loss_v16(\n            pred_coords_norm, coord_mask, \n            mu_bond=self.cfg.mu_bond, \n            sig_bond=self.cfg.sig_bond,\n            scale=scales\n        )\n        \n        if torch.isnan(L_bond):\n            L_bond = torch.tensor(0.0, device=v_pred.device, requires_grad=True)\n        \n        # Combined loss\n        total = L_flow + self.cfg.fape_weight * L_fape + self.cfg.bond_weight * L_bond\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_v16(pred_coords_norm, coord_mask)\n        \n        # Mean bond distance in Angstroms\n        coords_angstrom = denormalize_coordinates_v16(pred_coords_norm, centers, scales)\n        bonds = (coords_angstrom[:, 1:] - coords_angstrom[:, :-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            '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('V16 Loss with bond distance ready')\nprint(f'  FAPE weight: {cfg.fape_weight}')\nprint(f'  Bond weight: {cfg.bond_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_v16(model, batch, loss_fn, device):\n    \"\"\"V16: 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    \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)\n    \n    if torch.isnan(loss):\n        return None, {'skipped': True}\n    \n    return loss, metrics\n\n\n@torch.no_grad()\ndef validate_v16(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        \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)\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_v16(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_v16(pred_normalized, centers, scales)\n        \n        for b in range(len(target_ids)):\n            tm, debug = compute_tm_score_v16(\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 V16 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 V16 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 = RNAFlowDatasetV16(\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 = RNAFlowDatasetV16(\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_v16, 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_v16, 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 = NullivaRNAFlowV16(\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 V16: {total_params:,} params')\n    \n    loss_fn = LossV16(cfg)\n    optimizer = torch.optim.AdamW(model.parameters(), lr=cfg.lr, weight_decay=cfg.weight_decay)\n    \n    # V16: Much shorter warmup\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 V16')\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_v16(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_v16(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}')\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_v16(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_v16(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()"}]}