{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":118765,"databundleVersionId":15231210,"isSourceIdPinned":false},{"sourceType":"datasetVersion","sourceId":15138794,"datasetId":9694599,"databundleVersionId":16028046}],"dockerImageVersionId":31286,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"#Environment Setup\n\n# Core libraries\nimport os\nimport gc\nimport math\nimport random\nimport numpy as np\nimport pandas as pd\nfrom pathlib import Path\n\n# Visualization (for later analysis)\nimport matplotlib.pyplot as plt\n\n# Deep Learning\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\n\n# Progress bars\nfrom tqdm.auto import tqdm\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T08:33:06.265335Z","iopub.execute_input":"2026-03-16T08:33:06.265862Z","iopub.status.idle":"2026-03-16T08:33:10.726104Z","shell.execute_reply.started":"2026-03-16T08:33:06.26583Z","shell.execute_reply":"2026-03-16T08:33:10.725335Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Global Configuration\n\nclass CFG:\n    COMPETITION_DIR = Path(\"/kaggle/input/competitions/stanford-rna-3d-folding-2\")\n    TRAIN_SEQ_PATH = COMPETITION_DIR / \"train_sequences.csv\"\n    VAL_SEQ_PATH   = COMPETITION_DIR / \"validation_sequences.csv\"\n    TEST_SEQ_PATH  = COMPETITION_DIR / \"test_sequences.csv\"\n    TRAIN_LABELS_PATH = COMPETITION_DIR / \"train_labels.csv\"\n    VAL_LABELS_PATH   = COMPETITION_DIR / \"validation_labels.csv\"\n    MSA_DIR = COMPETITION_DIR / \"MSA\"\n    PDB_DIR = COMPETITION_DIR / \"PDB_RNA\"\n    SEED = 42\n    DEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n    MAX_SEQ_LEN = 4640\n    NUM_WORKERS = 1\n    BATCH_SIZE = 4\n    MIXED_PRECISION = False\nprint(\"CFG paths verified.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T08:33:10.727641Z","iopub.execute_input":"2026-03-16T08:33:10.728151Z","iopub.status.idle":"2026-03-16T08:33:10.971436Z","shell.execute_reply.started":"2026-03-16T08:33:10.728124Z","shell.execute_reply":"2026-03-16T08:33:10.970801Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#import os\n#base = \"/kaggle/input/competitions\"\n#print(\"Level 1:\", os.listdir(base))\n#for folder in os.listdir(base):\n#    path2 = os.path.join(base, folder)\n#    print(f\"\\nInside {folder}:\")\n#    print(os.listdir(path2)[:20])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T08:33:10.972153Z","iopub.execute_input":"2026-03-16T08:33:10.972348Z","iopub.status.idle":"2026-03-16T08:33:10.983358Z","shell.execute_reply.started":"2026-03-16T08:33:10.972328Z","shell.execute_reply":"2026-03-16T08:33:10.982791Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Reproducibility\ndef set_seed(seed=42):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n\nset_seed(CFG.SEED)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T08:33:10.984293Z","iopub.execute_input":"2026-03-16T08:33:10.984525Z","iopub.status.idle":"2026-03-16T08:33:10.999617Z","shell.execute_reply.started":"2026-03-16T08:33:10.984504Z","shell.execute_reply":"2026-03-16T08:33:10.999049Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# GPU Check\nprint(\"Device:\", CFG.DEVICE)\nif CFG.DEVICE == \"cuda\":\n    print(\"GPU:\", torch.cuda.get_device_name(0))\n    print(\"VRAM:\", round(torch.cuda.get_device_properties(0).total_memory / 1e9, 2), \"GB\")\n\n# Memory Cleanup Helper\ndef cleanup():\n    gc.collect()\n    torch.cuda.empty_cache()\nprint(\"Environment setup complete.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T08:33:11.001431Z","iopub.execute_input":"2026-03-16T08:33:11.001639Z","iopub.status.idle":"2026-03-16T08:33:11.044809Z","shell.execute_reply.started":"2026-03-16T08:33:11.00162Z","shell.execute_reply":"2026-03-16T08:33:11.044298Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#RNA Tokenizer\n\n# RNA vocabulary\nRNA_VOCAB = {\"A\": 0,\"C\": 1,\"G\": 2,\"U\": 3,\"PAD\": 4}\nIDX2RNA = {v: k for k, v in RNA_VOCAB.items()}\nVOCAB_SIZE = len(RNA_VOCAB)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T08:33:11.045546Z","iopub.execute_input":"2026-03-16T08:33:11.045803Z","iopub.status.idle":"2026-03-16T08:33:11.049877Z","shell.execute_reply.started":"2026-03-16T08:33:11.045776Z","shell.execute_reply":"2026-03-16T08:33:11.049177Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Tokenizer\ndef tokenize_sequence(seq: str):\n    \"\"\"\n    Convert RNA sequence string to list of token IDs.\n    Unknown characters are ignored.\n    \"\"\"\n    tokens = [RNA_VOCAB.get(nt, RNA_VOCAB[\"PAD\"]) for nt in seq]\n    return tokens","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T08:33:11.050808Z","iopub.execute_input":"2026-03-16T08:33:11.051067Z","iopub.status.idle":"2026-03-16T08:33:11.061629Z","shell.execute_reply.started":"2026-03-16T08:33:11.05104Z","shell.execute_reply":"2026-03-16T08:33:11.06097Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Padding Utility\ndef pad_sequence(tokens, max_len):\n    \"\"\"\n    Pad sequence to max_len and create attention mask.\n    \"\"\"\n    length = len(tokens)\n    if length > max_len:\n        tokens = tokens[:max_len]\n        mask = [1] * max_len\n    else:\n        pad_len = max_len - length\n        tokens = tokens + [RNA_VOCAB[\"PAD\"]] * pad_len\n        mask = [1] * length + [0] * pad_len\n    return tokens, mask","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T08:33:11.062483Z","iopub.execute_input":"2026-03-16T08:33:11.062829Z","iopub.status.idle":"2026-03-16T08:33:11.074311Z","shell.execute_reply.started":"2026-03-16T08:33:11.062799Z","shell.execute_reply":"2026-03-16T08:33:11.073799Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Tensor Converter\ndef encode_sequence(seq, max_len):\n    tokens = tokenize_sequence(seq)\n    tokens, mask = pad_sequence(tokens, max_len)\n    tokens = torch.tensor(tokens, dtype=torch.long)\n    mask = torch.tensor(mask, dtype=torch.float32)\n    return tokens, mask","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T08:33:11.075121Z","iopub.execute_input":"2026-03-16T08:33:11.07544Z","iopub.status.idle":"2026-03-16T08:33:11.08457Z","shell.execute_reply.started":"2026-03-16T08:33:11.07541Z","shell.execute_reply":"2026-03-16T08:33:11.083918Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Test Example\ntest_seq = \"ACGUACGUAGCU\"\ntokens, mask = encode_sequence(test_seq, max_len=20)\nprint(\"Sequence:\", test_seq)\nprint(\"Tokens:\", tokens)\nprint(\"Mask:\", mask)\nprint(\"Vocab size:\", VOCAB_SIZE)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T08:33:11.085399Z","iopub.execute_input":"2026-03-16T08:33:11.085878Z","iopub.status.idle":"2026-03-16T08:33:11.183175Z","shell.execute_reply.started":"2026-03-16T08:33:11.085857Z","shell.execute_reply":"2026-03-16T08:33:11.182474Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Sequence CSV Parsing\ndef load_sequence_csv(path):\n    df = pd.read_csv(path)\n    print(f\"Loaded {path.name}: {len(df)} entries\")\n    return df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T08:33:11.184013Z","iopub.execute_input":"2026-03-16T08:33:11.18428Z","iopub.status.idle":"2026-03-16T08:33:11.187586Z","shell.execute_reply.started":"2026-03-16T08:33:11.18426Z","shell.execute_reply":"2026-03-16T08:33:11.186969Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nbase = \"/kaggle/input/competitions\"\nprint(\"Level 1:\", os.listdir(base))\n\nfor folder in os.listdir(base):\n    path2 = os.path.join(base, folder)\n    print(f\"\\nInside {folder}:\")\n    print(os.listdir(path2)[:20])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T08:33:11.188336Z","iopub.execute_input":"2026-03-16T08:33:11.188597Z","iopub.status.idle":"2026-03-16T08:33:11.199575Z","shell.execute_reply.started":"2026-03-16T08:33:11.188576Z","shell.execute_reply":"2026-03-16T08:33:11.199016Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load datasets\ntrain_seq_df = load_sequence_csv(CFG.TRAIN_SEQ_PATH)\nval_seq_df   = load_sequence_csv(CFG.VAL_SEQ_PATH)\ntest_seq_df  = load_sequence_csv(CFG.TEST_SEQ_PATH)\nprint(\"\\nColumns:\")\ndisplay(train_seq_df.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T08:33:11.200445Z","iopub.execute_input":"2026-03-16T08:33:11.200818Z","iopub.status.idle":"2026-03-16T08:33:11.813601Z","shell.execute_reply.started":"2026-03-16T08:33:11.200797Z","shell.execute_reply":"2026-03-16T08:33:11.812962Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Basic sequence statistics\ndef sequence_stats(df, name):\n    lengths = df[\"sequence\"].apply(len)\n    print(f\"\\n{name} Sequence Stats\")\n    print(f\"Count: {len(lengths)}\")\n    print(f\"Min length: {lengths.min()}\")\n    print(f\"Max length: {lengths.max()}\")\n    print(f\"Mean length: {int(lengths.mean())}\")\nsequence_stats(train_seq_df, \"Train\")\nsequence_stats(val_seq_df, \"Validation\")\nsequence_stats(test_seq_df, \"Test\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T08:33:11.816215Z","iopub.execute_input":"2026-03-16T08:33:11.816869Z","iopub.status.idle":"2026-03-16T08:33:11.824817Z","shell.execute_reply.started":"2026-03-16T08:33:11.816842Z","shell.execute_reply":"2026-03-16T08:33:11.824055Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Example tokenization preview\nexample_seq = train_seq_df.iloc[0][\"sequence\"]\ntokens, mask = encode_sequence(example_seq, max_len=min(len(example_seq), CFG.MAX_SEQ_LEN))\nprint(\"\\nExample target_id:\", train_seq_df.iloc[0][\"target_id\"])\nprint(\"Sequence snippet:\", example_seq[:60], \"...\")\nprint(\"Tokenized shape:\", tokens.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T08:33:11.825665Z","iopub.execute_input":"2026-03-16T08:33:11.825893Z","iopub.status.idle":"2026-03-16T08:33:11.840664Z","shell.execute_reply.started":"2026-03-16T08:33:11.825872Z","shell.execute_reply":"2026-03-16T08:33:11.840059Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Label Parsing\ndef load_labels_csv(path):\n    df = pd.read_csv(path)\n    print(f\"Loaded {path.name}: {len(df)} rows\")\n    return df\n    \n# Load label datasets\ntrain_labels_df = load_labels_csv(CFG.TRAIN_LABELS_PATH)\nval_labels_df   = load_labels_csv(CFG.VAL_LABELS_PATH)\ndisplay(train_labels_df.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T08:33:11.841508Z","iopub.execute_input":"2026-03-16T08:33:11.841795Z","iopub.status.idle":"2026-03-16T08:33:20.610491Z","shell.execute_reply.started":"2026-03-16T08:33:11.841763Z","shell.execute_reply":"2026-03-16T08:33:20.609873Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Identify coordinate columns\ncoord_cols = [c for c in train_labels_df.columns if c.startswith((\"x_\", \"y_\", \"z_\"))]\nprint(f\"\\nTotal coordinate columns: {len(coord_cols)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T08:33:20.611358Z","iopub.execute_input":"2026-03-16T08:33:20.611613Z","iopub.status.idle":"2026-03-16T08:33:20.616118Z","shell.execute_reply.started":"2026-03-16T08:33:20.611592Z","shell.execute_reply":"2026-03-16T08:33:20.615522Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def normalize_coords(coords):\n    return coords / 50.0","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T08:33:20.61701Z","iopub.execute_input":"2026-03-16T08:33:20.617273Z","iopub.status.idle":"2026-03-16T08:33:20.628304Z","shell.execute_reply.started":"2026-03-16T08:33:20.617253Z","shell.execute_reply":"2026-03-16T08:33:20.627632Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Parse target_id from ID column\ndef extract_target_id(full_id):\n    return full_id.split(\"_\")[0]\ntrain_labels_df[\"target_id\"] = train_labels_df[\"ID\"].apply(extract_target_id)\nval_labels_df[\"target_id\"]   = val_labels_df[\"ID\"].apply(extract_target_id)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T08:33:20.629062Z","iopub.execute_input":"2026-03-16T08:33:20.629335Z","iopub.status.idle":"2026-03-16T08:33:22.931877Z","shell.execute_reply.started":"2026-03-16T08:33:20.629305Z","shell.execute_reply":"2026-03-16T08:33:22.931294Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Build structured coordinate arrays\ndef build_structure_dict(df):\n    structures = {}\n    for target_id, group in tqdm(df.groupby(\"target_id\")):\n        coords = group[coord_cols].values.astype(np.float32)\n        is_nan      = np.isnan(coords).any(axis=-1)\n        is_sentinel = (np.abs(coords) > 1e10).any(axis=-1)\n        invalid     = is_nan | is_sentinel\n\n        valid_mask = ~invalid \n        coords[invalid] = 0.0\n        coords = np.nan_to_num(coords, nan=0.0)\n        structures[target_id] = { \"coords\": coords,\"valid_mask\": valid_mask }\n    return structures\n\n# Rebuild both\nprint(\"Rebuilding training structures...\")\ntrain_structures = build_structure_dict(train_labels_df)\nprint(\"Rebuilding validation structures...\")\nval_structures = build_structure_dict(val_labels_df)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T08:33:22.932636Z","iopub.execute_input":"2026-03-16T08:33:22.932834Z","iopub.status.idle":"2026-03-16T08:33:26.97363Z","shell.execute_reply.started":"2026-03-16T08:33:22.932815Z","shell.execute_reply":"2026-03-16T08:33:26.972738Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Bug 1 Diagnostic — run this after building val_structures\nprint(\"val_structures keys sample:\", list(val_structures.keys())[:5])\nprint(\"val_seq_df target_ids sample:\", val_seq_df[\"target_id\"].values[:5])\nmissing = []\nfor tid in val_seq_df[\"target_id\"].values:\n    if tid not in val_structures:\n        missing.append(tid)\nprint(f\"\\nMissing val structures: {len(missing)} / {len(val_seq_df)}\")\nif missing:\n    print(\"Missing IDs:\", missing)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T08:33:26.974797Z","iopub.execute_input":"2026-03-16T08:33:26.975181Z","iopub.status.idle":"2026-03-16T08:33:26.980584Z","shell.execute_reply.started":"2026-03-16T08:33:26.975147Z","shell.execute_reply":"2026-03-16T08:33:26.979801Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Example structure\nexample_id = list(train_structures.keys())[0]\nprint(f\"\\nExample target_id: {example_id}\")\nprint(\"Coords shape:     \", train_structures[example_id][\"coords\"].shape)\nprint(\"Valid mask shape: \", train_structures[example_id][\"valid_mask\"].shape)\nprint(\"Valid residues:   \", train_structures[example_id][\"valid_mask\"].sum(), \n      \"/\", train_structures[example_id][\"valid_mask\"].shape[0])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T08:33:26.981474Z","iopub.execute_input":"2026-03-16T08:33:26.981743Z","iopub.status.idle":"2026-03-16T08:33:26.995358Z","shell.execute_reply.started":"2026-03-16T08:33:26.98172Z","shell.execute_reply":"2026-03-16T08:33:26.994799Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#MSA Parsing\n# RNA mapping extended for gaps\nMSA_VOCAB = {\"A\": 0,\"C\": 1,\"G\": 2,\"U\": 3,\"-\": 4,\"PAD\": 5}\nMSA_PAD_IDX = MSA_VOCAB[\"PAD\"]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T08:33:26.996047Z","iopub.execute_input":"2026-03-16T08:33:26.996447Z","iopub.status.idle":"2026-03-16T08:33:27.005904Z","shell.execute_reply.started":"2026-03-16T08:33:26.996426Z","shell.execute_reply":"2026-03-16T08:33:27.005239Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Parse single MSA file\n#def parse_msa_file(msa_path, max_alignments=128):\n#    \"\"\"\n#    Reads MSA FASTA and returns tensor:\n#    (num_alignments, seq_len)\n#    \"\"\"\n#    sequences = []\n#    for record in SeqIO.parse(msa_path, \"fasta\"):\n#        seq = str(record.seq)\n#        tokens = [MSA_VOCAB.get(nt, MSA_PAD_IDX) for nt in seq]\n#        sequences.append(tokens)        \n#        if len(sequences) >= max_alignments:\n#            break   \n#    if not sequences:\n#        return None\n#    msa = torch.tensor(sequences, dtype=torch.long)\n#    return msa\n\ndef parse_msa_file(msa_path, max_alignments=128):\n    sequences = []\n    current_seq = []\n    with open(msa_path, \"r\") as f:\n        for line in f:\n            line = line.strip()\n            if not line:\n                continue\n            if line.startswith(\">\"):\n                if current_seq:\n                    seq = \"\".join(current_seq)\n                    tokens = [MSA_VOCAB.get(nt, MSA_PAD_IDX) for nt in seq]\n                    sequences.append(tokens)\n                    current_seq = []\n                if len(sequences) >= max_alignments:\n                    break\n            else:\n                current_seq.append(line)\n\n    # Save last sequence\n    if current_seq and len(sequences) < max_alignments:\n        seq = \"\".join(current_seq)\n        tokens = [MSA_VOCAB.get(nt, MSA_PAD_IDX) for nt in seq]\n        sequences.append(tokens)\n    if not sequences:\n        return None\n    msa = torch.tensor(sequences, dtype=torch.long)\n    return msa","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T08:33:27.006758Z","iopub.execute_input":"2026-03-16T08:33:27.007099Z","iopub.status.idle":"2026-03-16T08:33:27.016362Z","shell.execute_reply.started":"2026-03-16T08:33:27.00707Z","shell.execute_reply":"2026-03-16T08:33:27.015822Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load MSA for one target\ndef load_msa(target_id):\n    msa_file = CFG.MSA_DIR / f\"{target_id}.MSA.fasta\"    \n    if not msa_file.exists():\n        return None    \n    msa_tensor = parse_msa_file(msa_file)\n    return msa_tensor","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T08:33:27.017036Z","iopub.execute_input":"2026-03-16T08:33:27.017274Z","iopub.status.idle":"2026-03-16T08:33:27.025426Z","shell.execute_reply.started":"2026-03-16T08:33:27.017247Z","shell.execute_reply":"2026-03-16T08:33:27.024784Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Test MSA loading\nexample_target = train_seq_df.iloc[0][\"target_id\"]\nmsa_tensor = load_msa(example_target)\nif msa_tensor is not None:\n    print(\"Example target:\", example_target)\n    print(\"MSA shape:\", msa_tensor.shape)\nelse:\n    print(\"No MSA found for:\", example_target)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T08:33:27.026384Z","iopub.execute_input":"2026-03-16T08:33:27.026623Z","iopub.status.idle":"2026-03-16T08:33:27.044357Z","shell.execute_reply.started":"2026-03-16T08:33:27.026594Z","shell.execute_reply":"2026-03-16T08:33:27.043692Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Feature Builder\ndef build_sample(target_id, seq_df, structure_dict):\n    row = seq_df[seq_df[\"target_id\"] == target_id].iloc[0]\n    seq = row[\"sequence\"]\n    tokens, mask = encode_sequence(seq, CFG.MAX_SEQ_LEN)\n    msa_tensor = load_msa(target_id)\n    MAX_MSA = 128\n    if msa_tensor is not None:\n        msa_len = msa_tensor.shape[1]\n        if msa_len > CFG.MAX_SEQ_LEN:\n            msa_tensor = msa_tensor[:, :CFG.MAX_SEQ_LEN]\n        else:\n            pad_cols = CFG.MAX_SEQ_LEN - msa_len\n            pad_tensor = torch.full((msa_tensor.shape[0], pad_cols), MSA_PAD_IDX, dtype=torch.long)\n            msa_tensor = torch.cat([msa_tensor, pad_tensor], dim=1)\n        msa_rows = msa_tensor.shape[0]\n        if msa_rows > MAX_MSA:\n            msa_tensor = msa_tensor[:MAX_MSA]\n        else:\n            pad_rows = MAX_MSA - msa_rows\n            pad_tensor = torch.full((pad_rows, CFG.MAX_SEQ_LEN), MSA_PAD_IDX, dtype=torch.long)\n            msa_tensor = torch.cat([msa_tensor, pad_tensor], dim=0)\n    else:\n        msa_tensor = torch.full((MAX_MSA, CFG.MAX_SEQ_LEN), MSA_PAD_IDX, dtype=torch.long)\n    entry = structure_dict.get(target_id)\n    if entry is not None:\n        coords     = torch.tensor(entry[\"coords\"],     dtype=torch.float32)\n        coord_mask = torch.tensor(entry[\"valid_mask\"], dtype=torch.float32)  # ✅ per-residue valid\n        coords     = normalize_coords(coords)\n\n        if coords.shape[0] > CFG.MAX_SEQ_LEN:\n            coords     = coords[:CFG.MAX_SEQ_LEN]\n            coord_mask = coord_mask[:CFG.MAX_SEQ_LEN]\n        else:\n            pad_rows = CFG.MAX_SEQ_LEN - coords.shape[0]\n            coords     = torch.cat([coords, torch.zeros((pad_rows, 3))], dim=0)\n            coord_mask = torch.cat([coord_mask, torch.zeros(pad_rows)],  dim=0)\n    else:\n        coords     = torch.zeros((CFG.MAX_SEQ_LEN, 3),  dtype=torch.float32)\n        coord_mask = torch.zeros(CFG.MAX_SEQ_LEN,       dtype=torch.float32)\n\n    return { \"seq_tokens\": tokens, \"seq_mask\":   mask, \"msa\":        msa_tensor, \"coords\":     coords, \"coord_mask\": coord_mask}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T08:33:27.045174Z","iopub.execute_input":"2026-03-16T08:33:27.045466Z","iopub.status.idle":"2026-03-16T08:33:27.053417Z","shell.execute_reply.started":"2026-03-16T08:33:27.045446Z","shell.execute_reply":"2026-03-16T08:33:27.052824Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Test feature building\ntest_id = train_seq_df.iloc[0][\"target_id\"]\nsample = build_sample(test_id, train_seq_df, train_structures)\nprint(\"Sample target:\", test_id)\nprint(\"Sequence tokens:\", sample[\"seq_tokens\"].shape)\nprint(\"Sequence mask:\", sample[\"seq_mask\"].shape)\nprint(\"MSA tensor:\", sample[\"msa\"].shape)\nprint(\"Coordinates:\", sample[\"coords\"].shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T08:33:27.054297Z","iopub.execute_input":"2026-03-16T08:33:27.05463Z","iopub.status.idle":"2026-03-16T08:33:27.083544Z","shell.execute_reply.started":"2026-03-16T08:33:27.054601Z","shell.execute_reply":"2026-03-16T08:33:27.083035Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Pytorch Dataset & DataLoader\nclass RNADataset(Dataset):\n    def __init__(self, seq_df, structure_dict):\n        self.seq_df = seq_df\n        self.structure_dict = structure_dict\n        self.target_ids = seq_df[\"target_id\"].values    \n    def __len__(self):\n        return len(self.target_ids)    \n    def __getitem__(self, idx):\n        target_id = self.target_ids[idx]\n        sample = build_sample(target_id, self.seq_df, self.structure_dict)\n        return sample","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T08:33:27.084955Z","iopub.execute_input":"2026-03-16T08:33:27.085322Z","iopub.status.idle":"2026-03-16T08:33:27.089747Z","shell.execute_reply.started":"2026-03-16T08:33:27.0853Z","shell.execute_reply":"2026-03-16T08:33:27.089126Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Collate function for batching\ndef collate_fn(batch):\n    MAX_MSA = 128\n    MAX_LEN = CFG.MAX_SEQ_LEN\n    seq_tokens = torch.stack([b[\"seq_tokens\"] for b in batch])\n    seq_mask   = torch.stack([b[\"seq_mask\"] for b in batch])\n    coords     = torch.stack([b[\"coords\"] for b in batch])    \n    coord_mask = torch.stack([b[\"coord_mask\"] for b in batch])\n    msa_list = []\n    for b in batch:\n        msa = b[\"msa\"]        \n        rows, cols = msa.shape\n        # pad rows\n        if rows < MAX_MSA:\n            pad = torch.full((MAX_MSA - rows, cols), MSA_PAD_IDX, dtype=torch.long)\n            msa = torch.cat([msa, pad], dim=0)\n        elif rows > MAX_MSA:\n            msa = msa[:MAX_MSA]\n        # pad cols\n        if cols < MAX_LEN:\n            pad = torch.full((msa.shape[0], MAX_LEN - cols), MSA_PAD_IDX, dtype=torch.long)\n            msa = torch.cat([msa, pad], dim=1)\n        elif cols > MAX_LEN:\n            msa = msa[:, :MAX_LEN]\n        msa_list.append(msa)    \n    msa = torch.stack(msa_list)    \n    return {\"seq_tokens\": seq_tokens,\"seq_mask\": seq_mask,\"msa\": msa,\"coords\": coords,\"coord_mask\": coord_mask}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T08:33:27.090747Z","iopub.execute_input":"2026-03-16T08:33:27.090967Z","iopub.status.idle":"2026-03-16T08:33:27.101433Z","shell.execute_reply.started":"2026-03-16T08:33:27.090937Z","shell.execute_reply":"2026-03-16T08:33:27.100824Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Create datasets\ntrain_dataset = RNADataset(train_seq_df, train_structures)\nval_dataset   = RNADataset(val_seq_df, val_structures)\ntrain_loader = DataLoader(train_dataset, batch_size=CFG.BATCH_SIZE,shuffle=True,  num_workers=CFG.NUM_WORKERS, collate_fn=collate_fn)\nval_loader   = DataLoader(val_dataset,   batch_size=CFG.BATCH_SIZE,shuffle=False, num_workers=CFG.NUM_WORKERS, collate_fn=collate_fn)\n\n# Verify all batches are clean\nfor i, val_batch in enumerate(val_loader):\n    coords = val_batch[\"coords\"]\n    print(f\"Batch {i} — max coord: {coords.abs().max().item():.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T08:33:27.102338Z","iopub.execute_input":"2026-03-16T08:33:27.102599Z","iopub.status.idle":"2026-03-16T08:33:27.696548Z","shell.execute_reply.started":"2026-03-16T08:33:27.102571Z","shell.execute_reply":"2026-03-16T08:33:27.69573Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# DataLoaders\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=CFG.BATCH_SIZE,\n    shuffle=True,\n    num_workers=CFG.NUM_WORKERS,\n    collate_fn=collate_fn)\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=CFG.BATCH_SIZE,\n    shuffle=False,\n    num_workers=CFG.NUM_WORKERS,\n    collate_fn=collate_fn)\n\n# Test batch\nbatch = next(iter(train_loader))\nprint(\"Batch seq_tokens:\", batch[\"seq_tokens\"].shape)\nprint(\"Batch seq_mask:\", batch[\"seq_mask\"].shape)\nprint(\"Batch msa:\", batch[\"msa\"].shape)\nprint(\"Batch coords:\", batch[\"coords\"].shape)\nprint(\"Batch coord_mask:\", batch[\"coord_mask\"].shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T08:33:27.698422Z","iopub.execute_input":"2026-03-16T08:33:27.698676Z","iopub.status.idle":"2026-03-16T08:33:28.174332Z","shell.execute_reply.started":"2026-03-16T08:33:27.69865Z","shell.execute_reply":"2026-03-16T08:33:28.173513Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Baseline RNA Structure Prediction Model","metadata":{}},{"cell_type":"code","source":"#RNA Structure Prediction Model\nclass PositionalEncoding(nn.Module):\n    def __init__(self, d_model, max_len=1024):\n        super().__init__()\n        pe = torch.zeros(max_len, d_model)\n        position = torch.arange(0, max_len).unsqueeze(1)\n        div_term = torch.exp(torch.arange(0, d_model, 2) * (-math.log(10000.0) / d_model))\n        pe[:, 0::2] = torch.sin(position * div_term)\n        pe[:, 1::2] = torch.cos(position * div_term)\n        pe = pe.unsqueeze(0)\n        self.register_buffer('pe', pe)\n\n    def forward(self, x):\n        return x + self.pe[:, :x.size(1)]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T08:33:28.175659Z","iopub.execute_input":"2026-03-16T08:33:28.176345Z","iopub.status.idle":"2026-03-16T08:33:28.181717Z","shell.execute_reply.started":"2026-03-16T08:33:28.176315Z","shell.execute_reply":"2026-03-16T08:33:28.180936Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class RNACoordinateModel(nn.Module):\n    def __init__(self, vocab_size, embed_dim=128):\n        super().__init__()\n\n        self.embedding = nn.Embedding(vocab_size, embed_dim)\n        encoder_layer = nn.TransformerEncoderLayer(\n            d_model=embed_dim,\n            nhead=4,\n            dim_feedforward=embed_dim, \n            dropout=0.0,               \n            activation=\"gelu\",\n            batch_first=True,\n            norm_first=True)\n\n        self.transformer = nn.TransformerEncoder(encoder_layer, num_layers=2)\n        self.head = nn.Linear(embed_dim, 3)\n    def forward(self, seq_tokens, seq_mask):\n        x = self.embedding(seq_tokens)\n        x = torch.clamp(x, -5.0, 5.0)\n        pad_mask = (seq_mask == 0)\n        x = self.transformer(x, src_key_padding_mask=pad_mask)\n        coords = self.head(x)\n        coords = torch.tanh(coords) * 20.0\n        coords = torch.nan_to_num(coords, nan=0.0, posinf=10.0, neginf=-10.0)\n        return coords\n        \n# Initialize model\nmodel = RNACoordinateModel(vocab_size=VOCAB_SIZE).to(CFG.DEVICE)\nprint(model)\n\ndef init_weights(m):\n    if isinstance(m, nn.Linear):\n        nn.init.xavier_uniform_(m.weight)\n        if m.bias is not None:\n            nn.init.zeros_(m.bias)\n    elif isinstance(m, nn.Embedding):\n        nn.init.normal_(m.weight, mean=0.0, std=0.02)\nmodel.apply(init_weights)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T08:33:28.182652Z","iopub.execute_input":"2026-03-16T08:33:28.18331Z","iopub.status.idle":"2026-03-16T08:33:28.501466Z","shell.execute_reply.started":"2026-03-16T08:33:28.183288Z","shell.execute_reply":"2026-03-16T08:33:28.500904Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Structural Loss Function\n\n# Pairwise Distance Matrix\ndef pairwise_distances(coords):\n    diff = coords.unsqueeze(2) - coords.unsqueeze(1)\n    dist_sq = (diff ** 2).sum(-1)\n    return dist_sq\n# Coordinate Loss\ndef coordinate_loss(pred, true, mask):\n    mask = mask.unsqueeze(-1)\n    return F.mse_loss(pred * mask, true * mask)\n\n# Distance Matrix Loss\ndef distance_loss(pred, true, mask):\n    pred_dist = pairwise_distances(pred)\n    true_dist = pairwise_distances(true)   \n    mask2d = mask.unsqueeze(1) * mask.unsqueeze(2)\n    return F.mse_loss(pred_dist * mask2d, true_dist * mask2d)\n\n# Bond Length Regularization\ndef bond_regularization(coords, mask):\n    diffs = coords[:, 1:] - coords[:, :-1]\n    bond_lengths = torch.sqrt((diffs**2).sum(-1) + 1e-8)   \n    mask = mask[:, 1:] * mask[:, :-1]\n    ideal = 3.8 \n    return F.mse_loss(bond_lengths * mask, torch.full_like(bond_lengths, ideal) * mask)\n\n# Total Structural Loss\ndef structural_loss(pred, true, coord_mask):\n    mask3d  = coord_mask.unsqueeze(-1).expand_as(pred)\n    sq_err  = (pred - true) ** 2 * mask3d\n    n_valid = mask3d.sum().clamp(min=1.0)\n    return sq_err.sum() / n_valid, None, None","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T08:33:28.502306Z","iopub.execute_input":"2026-03-16T08:33:28.502791Z","iopub.status.idle":"2026-03-16T08:33:28.509419Z","shell.execute_reply.started":"2026-03-16T08:33:28.502768Z","shell.execute_reply":"2026-03-16T08:33:28.508726Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Training Engine\nfrom torch.cuda.amp import autocast, GradScaler\nEPOCHS = 8 \nbest_val_loss = float(\"inf\")\noptimizer  = torch.optim.AdamW(model.parameters(), lr=1e-4, weight_decay=1e-2)\nscheduler  = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=EPOCHS, eta_min=1e-6)\nscaler     = torch.amp.GradScaler('cuda', enabled=CFG.MIXED_PRECISION)\n\ndef train_one_epoch():\n    model.train()\n    total_loss = 0\n    for batch in train_loader:\n        seq_tokens = batch[\"seq_tokens\"].to(CFG.DEVICE)\n        seq_mask   = batch[\"seq_mask\"].to(CFG.DEVICE)\n        coords     = batch[\"coords\"].to(CFG.DEVICE)\n        coord_mask = batch[\"coord_mask\"].to(CFG.DEVICE)\n        optimizer.zero_grad()\n        pred_coords = model(seq_tokens, seq_mask)\n        loss = structural_loss(pred_coords, coords, coord_mask)[0]\n        loss.backward()\n        torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n        optimizer.step()\n        total_loss += loss.item()\n    return total_loss / len(train_loader)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T08:33:28.510332Z","iopub.execute_input":"2026-03-16T08:33:28.510694Z","iopub.status.idle":"2026-03-16T08:33:31.349534Z","shell.execute_reply.started":"2026-03-16T08:33:28.510663Z","shell.execute_reply":"2026-03-16T08:33:31.348926Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Sanity check — run ONE batch manually before full training\nmodel.train()\nbatch = next(iter(train_loader))\nseq_tokens = batch[\"seq_tokens\"].to(CFG.DEVICE)\nseq_mask   = batch[\"seq_mask\"].to(CFG.DEVICE)\ncoords     = batch[\"coords\"].to(CFG.DEVICE)\npred = model(seq_tokens, seq_mask)\nloss, _, _ = structural_loss(pred, coords, seq_mask)\nprint(\"Sanity loss:\", loss.item())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T08:33:31.350407Z","iopub.execute_input":"2026-03-16T08:33:31.350835Z","iopub.status.idle":"2026-03-16T08:33:32.285586Z","shell.execute_reply.started":"2026-03-16T08:33:31.350806Z","shell.execute_reply":"2026-03-16T08:33:32.28481Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Validation Function\ndef validate():\n    model.eval()\n    total_loss = 0\n    with torch.no_grad():\n        for batch in val_loader:\n            seq_tokens = batch[\"seq_tokens\"].to(CFG.DEVICE)\n            seq_mask   = batch[\"seq_mask\"].to(CFG.DEVICE)\n            coords     = batch[\"coords\"].to(CFG.DEVICE)\n            coord_mask = batch[\"coord_mask\"].to(CFG.DEVICE)\n\n            pred_coords = model(seq_tokens, seq_mask)\n            loss = structural_loss(pred_coords, coords, coord_mask)[0]\n            total_loss += loss.item()\n    return total_loss / len(val_loader)\n\n# Training Loop\nfor epoch in range(EPOCHS):\n    train_loss = train_one_epoch()\n    val_loss   = validate()\n    scheduler.step()\n    print(f\"Epoch {epoch+1}/{EPOCHS}\")\n    print(f\"  Train Loss: {train_loss:.4f}\")\n    print(f\"  Val Loss:   {val_loss:.4f}\")\n    if val_loss < best_val_loss:\n        best_val_loss = val_loss\n        torch.save(model.state_dict(), \"best_rna_model.pt\")\n        print(\"Saved best model\")\nprint(f\"\\nBest Val Loss: {best_val_loss:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T08:33:32.286762Z","iopub.execute_input":"2026-03-16T08:33:32.287041Z","iopub.status.idle":"2026-03-16T09:16:56.920957Z","shell.execute_reply.started":"2026-03-16T08:33:32.287004Z","shell.execute_reply":"2026-03-16T09:16:56.920123Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Inference & Submission Pipeline\n\n#Load Best Model\nmodel.load_state_dict(torch.load(\"best_rna_model.pt\", map_location=CFG.DEVICE))\nmodel.eval()\nprint(\"Best model loaded for inference\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T09:16:56.922249Z","iopub.execute_input":"2026-03-16T09:16:56.922494Z","iopub.status.idle":"2026-03-16T09:16:56.935797Z","shell.execute_reply.started":"2026-03-16T09:16:56.922469Z","shell.execute_reply":"2026-03-16T09:16:56.935218Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Predict 3D structure for one RNA\ndef predict_structure(sample):\n    seq_tokens = sample[\"seq_tokens\"].unsqueeze(0).to(CFG.DEVICE)\n    seq_mask   = sample[\"seq_mask\"].unsqueeze(0).to(CFG.DEVICE)\n    with torch.no_grad():\n        coords = model(seq_tokens, seq_mask)[0]\n    coords = coords.cpu()\n    return coords","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T09:16:56.936765Z","iopub.execute_input":"2026-03-16T09:16:56.937042Z","iopub.status.idle":"2026-03-16T09:16:56.947739Z","shell.execute_reply.started":"2026-03-16T09:16:56.937013Z","shell.execute_reply":"2026-03-16T09:16:56.947037Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Generate 5 conformation\ndef predict_5_structures(sample):\n    base_coords = predict_structure(sample)\n    preds = []\n    for i in range(5):\n        noise = torch.randn_like(base_coords) * 0.05\n        coords = base_coords + noise\n        preds.append(coords)\n    return preds","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T09:16:56.95215Z","iopub.execute_input":"2026-03-16T09:16:56.952447Z","iopub.status.idle":"2026-03-16T09:16:56.963445Z","shell.execute_reply.started":"2026-03-16T09:16:56.952427Z","shell.execute_reply":"2026-03-16T09:16:56.962904Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Predict all test RNAs\ntest_predictions = {}\nfor idx in range(len(test_seq_df)):\n    target_id = test_seq_df.iloc[idx][\"target_id\"]\n    try:\n        sample = build_sample(target_id, test_seq_df, structure_dict={})\n        preds = predict_5_structures(sample)\n        test_predictions[target_id] = preds\n    except Exception as e:\n        print(f\" Failed {target_id}: {e}\")\n        zero = torch.zeros((CFG.MAX_SEQ_LEN, 3))\n        test_predictions[target_id] = [zero] * 5\nprint(\"Test predictions complete:\", len(test_predictions))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T09:16:56.964375Z","iopub.execute_input":"2026-03-16T09:16:56.964641Z","iopub.status.idle":"2026-03-16T09:16:58.044385Z","shell.execute_reply.started":"2026-03-16T09:16:56.964612Z","shell.execute_reply":"2026-03-16T09:16:58.043673Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Manually enter your epoch results\ntrain_loss = [5.9917, 5.3263, 4.9313, 4.7988, 4.4331, 4.0749, 3.8130, 3.7052, 3.6064, 3.5785]\nval_loss   = [4.6094, 4.2001, 3.7275, 4.1520, 3.5114, 3.6634,  4.2789, 4.4184, 4.4168, 4.3003]\n\n#Training & Validation Loss Curves\nepochs = range(1, len(train_loss) + 1)\nplt.figure(figsize=(6, 4), dpi=80)\nplt.plot(epochs, train_loss)\nplt.plot(epochs, val_loss)\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Loss\")\nplt.title(\"Training vs Validation Loss\")\nplt.legend([\"Train\", \"Validation\"])\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T09:16:58.045373Z","iopub.execute_input":"2026-03-16T09:16:58.045659Z","iopub.status.idle":"2026-03-16T09:16:58.231734Z","shell.execute_reply.started":"2026-03-16T09:16:58.045634Z","shell.execute_reply":"2026-03-16T09:16:58.231156Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#RNA Length Distribution\nimport matplotlib.pyplot as plt\nlengths = train_seq_df[\"sequence\"].apply(len)\nplt.figure(figsize=(6, 4), dpi=80)\nplt.hist(lengths, bins=50)\nplt.xlabel(\"RNA Length\")\nplt.ylabel(\"Count\")\nplt.title(\"RNA Length Distribution\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T09:16:58.232551Z","iopub.execute_input":"2026-03-16T09:16:58.232828Z","iopub.status.idle":"2026-03-16T09:16:58.384724Z","shell.execute_reply.started":"2026-03-16T09:16:58.232797Z","shell.execute_reply":"2026-03-16T09:16:58.384035Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Visualize all 28 test RNA structure projections\nfig, axes = plt.subplots(4, 7, figsize=(20, 12), dpi=80)\naxes = axes.flatten()\nfor idx, row in test_seq_df.iterrows():\n    target_id = row[\"target_id\"]\n    seq_len   = min(len(row[\"sequence\"]), CFG.MAX_SEQ_LEN)\n\n    # Get prediction\n    sample = build_sample(target_id, test_seq_df, structure_dict={})\n    preds  = predict_5_structures(sample)\n\n    # First conformation, denormalize, crop padding\n    coords = preds[0].cpu().numpy() * 50.0\n    coords = coords[:seq_len]\n    ax = axes[idx]\n    ax.scatter(coords[:, 0], coords[:, 1], s=1, alpha=0.5)\n    ax.set_title(target_id, fontsize=7)\n    ax.set_xlabel(\"X\", fontsize=6)\n    ax.set_ylabel(\"Y\", fontsize=6)\n    ax.tick_params(labelsize=5)\nplt.suptitle(\"RNA Structure Projections (XY) — All Test Targets\", fontsize=12)\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T09:16:58.38577Z","iopub.execute_input":"2026-03-16T09:16:58.38614Z","iopub.status.idle":"2026-03-16T09:17:02.56645Z","shell.execute_reply.started":"2026-03-16T09:16:58.386109Z","shell.execute_reply":"2026-03-16T09:17:02.565692Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Compare True vs Predicted Structure\nsample = next(iter(val_loader))\nseq_tokens = sample[\"seq_tokens\"].to(CFG.DEVICE)\nseq_mask   = sample[\"seq_mask\"].to(CFG.DEVICE)\nwith torch.no_grad():\n    preds = model(seq_tokens, seq_mask).cpu()\ntrue = sample[\"coords\"]\nplt.figure(figsize=(6, 4), dpi=80)\nplt.scatter(true[0,:,0], true[0,:,1], s=6)\nplt.scatter(preds[0,:,0], preds[0,:,1], s=6)\nplt.legend([\"True\", \"Predicted\"])\nplt.title(\"True vs Predicted RNA Structure (XY)\")\nplt.xlabel(\"X\")\nplt.ylabel(\"Y\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T09:17:02.567417Z","iopub.execute_input":"2026-03-16T09:17:02.567648Z","iopub.status.idle":"2026-03-16T09:17:03.232602Z","shell.execute_reply.started":"2026-03-16T09:17:02.567626Z","shell.execute_reply":"2026-03-16T09:17:03.231721Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Build Kaggle submission\nsample_sub = pd.read_csv(CFG.COMPETITION_DIR / \"sample_submission.csv\")\nprint(\"Sample sub shape:\", sample_sub.shape)\nprint(\"Sample sub columns:\", sample_sub.columns.tolist())\n\n# Build lookup from test sequences for resname\ntest_resname_lookup = {}\nfor _, row in test_seq_df.iterrows():\n    tid = row[\"target_id\"]\n    seq = row[\"sequence\"]\n    for i, nt in enumerate(seq):\n        test_resname_lookup[f\"{tid}_{i+1}\"] = nt\n\n# Build submission matching sample format exactly\nsubmission_rows = []\nfor target_id, preds in test_predictions.items():\n    seq = test_seq_df[test_seq_df[\"target_id\"] == target_id].iloc[0][\"sequence\"]\n    seq_len = len(seq)\n    effective_len = min(seq_len, CFG.MAX_SEQ_LEN)\n    for resid in range(1, effective_len + 1):\n        try:\n            row_id  = f\"{target_id}_{resid}\"\n            resname = seq[resid - 1]\n            row_dict = { \"ID\":      row_id, \"resname\": resname, \"resid\":   resid, }\n            for i, coords_tensor in enumerate(preds):\n                if resid - 1 < coords_tensor.shape[0]:\n                    xyz = coords_tensor[resid - 1].cpu().numpy()\n                else:\n                    xyz = np.zeros(3) \n                row_dict[f\"x_{i+1}\"] = float(xyz[0])\n                row_dict[f\"y_{i+1}\"] = float(xyz[1])\n                row_dict[f\"z_{i+1}\"] = float(xyz[2])\n            submission_rows.append(row_dict)\n        except Exception as e:\n            print(f\"⚠️ Skipped {target_id}_{resid}: {e}\")\nsubmission_df = pd.DataFrame(submission_rows)\ncol_order = [\"ID\", \"resname\", \"resid\",\n             \"x_1\",\"y_1\",\"z_1\",\n             \"x_2\",\"y_2\",\"z_2\",\n             \"x_3\",\"y_3\",\"z_3\",\n             \"x_4\",\"y_4\",\"z_4\",\n             \"x_5\",\"y_5\",\"z_5\"]\nsubmission_df = submission_df[col_order]\nprint(\"Your submission shape:\", submission_df.shape)\nprint(\"Expected shape:       \", sample_sub.shape)\nprint(\"Nulls:\", submission_df.isnull().sum().sum())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T09:17:03.233908Z","iopub.execute_input":"2026-03-16T09:17:03.234262Z","iopub.status.idle":"2026-03-16T09:17:03.627523Z","shell.execute_reply.started":"2026-03-16T09:17:03.234232Z","shell.execute_reply":"2026-03-16T09:17:03.626737Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Fill missing IDs with zeros and finalize submission\nmissing_ids = set(sample_sub[\"ID\"]) - set(submission_df[\"ID\"])\nprint(\"Missing IDs:\", len(missing_ids))\n\n# Get missing rows from sample submission\nmissing_rows = sample_sub[sample_sub[\"ID\"].isin(missing_ids)].copy()\n\n# Zero out all coordinate columns\ncoord_cols_sub = [\"x_1\",\"y_1\",\"z_1\",\"x_2\",\"y_2\",\"z_2\",\n                  \"x_3\",\"y_3\",\"z_3\",\"x_4\",\"y_4\",\"z_4\",\n                  \"x_5\",\"y_5\",\"z_5\"]\nmissing_rows[coord_cols_sub] = 0.0\nsubmission_final = pd.concat([submission_df, missing_rows], ignore_index=True)\n\n# Sort to match sample submission order exactly\nsubmission_final = submission_final.merge(\n    sample_sub[[\"ID\"]].reset_index().rename(columns={\"index\": \"sort_order\"}),\n    on=\"ID\"\n).sort_values(\"sort_order\").drop(\"sort_order\", axis=1).reset_index(drop=True)\n\n# Verify\nprint(\"Final shape:   \", submission_final.shape)\nprint(\"Expected shape:\", sample_sub.shape)\nprint(\"Nulls:\", submission_final.isnull().sum().sum())\n#print(submission_final.head())\n\n# Save final\nsubmission_final.to_csv(\"submission.csv\", index=False)\nprint(\"\\nsubmission.csv saved\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T09:17:03.628453Z","iopub.execute_input":"2026-03-16T09:17:03.628791Z","iopub.status.idle":"2026-03-16T09:17:03.88462Z","shell.execute_reply.started":"2026-03-16T09:17:03.628768Z","shell.execute_reply":"2026-03-16T09:17:03.884015Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(submission_final.shape)\n#print(submission_final.head())\n#print(submission_final.isnull().sum())\n\n#Download\nfrom IPython.display import FileLink\nFileLink(\"submission.csv\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-16T09:17:03.885548Z","iopub.execute_input":"2026-03-16T09:17:03.885852Z","iopub.status.idle":"2026-03-16T09:17:03.891485Z","shell.execute_reply.started":"2026-03-16T09:17:03.885816Z","shell.execute_reply":"2026-03-16T09:17:03.890721Z"}},"outputs":[],"execution_count":null}]}