{"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":[{"sourceId":118765,"databundleVersionId":15231210,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":3751217,"sourceType":"datasetVersion","datasetId":2241434},{"sourceId":6233380,"sourceType":"datasetVersion","datasetId":3580819},{"sourceId":10878276,"sourceType":"datasetVersion","datasetId":6758842},{"sourceId":10878463,"sourceType":"datasetVersion","datasetId":6759157},{"sourceId":10880374,"sourceType":"datasetVersion","datasetId":6760482},{"sourceId":10880419,"sourceType":"datasetVersion","datasetId":6760509},{"sourceId":11974777,"sourceType":"datasetVersion","datasetId":7530539},{"sourceId":14441699,"sourceType":"datasetVersion","datasetId":9224635},{"sourceId":14496440,"sourceType":"datasetVersion","datasetId":9259043}],"dockerImageVersionId":31234,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Stanford RNA 3D Folding — RhoFold + Template-Based Modeling with Robust Logging","metadata":{}},{"cell_type":"markdown","source":"# Testing Rhofold to see if it's working or not ","metadata":{}},{"cell_type":"code","source":"!pip install /kaggle/input/rhofold-dependencies-for-offline-use/ml_collections-1.1.0-py3-none-any.whl\n!pip install /kaggle/input/rhofold-dependencies-for-offline-use/simtk-0.1.0-py2.py3-none-any.whl\n!pip install /kaggle/input/biopython-cp312/biopython-1.86-cp312-cp312-manylinux_2_17_x86_64.manylinux2014_x86_64.whl\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-18T07:10:47.524254Z","iopub.execute_input":"2026-01-18T07:10:47.524834Z","iopub.status.idle":"2026-01-18T07:11:00.330135Z","shell.execute_reply.started":"2026-01-18T07:10:47.524803Z","shell.execute_reply":"2026-01-18T07:11:00.329424Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%cd /kaggle/input/rhofold-main/RhoFold-main\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-18T07:11:00.332003Z","iopub.execute_input":"2026-01-18T07:11:00.332318Z","iopub.status.idle":"2026-01-18T07:11:00.340692Z","shell.execute_reply.started":"2026-01-18T07:11:00.332291Z","shell.execute_reply":"2026-01-18T07:11:00.340103Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%ls\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-18T07:11:00.341568Z","iopub.execute_input":"2026-01-18T07:11:00.341914Z","iopub.status.idle":"2026-01-18T07:11:00.491284Z","shell.execute_reply.started":"2026-01-18T07:11:00.341886Z","shell.execute_reply":"2026-01-18T07:11:00.490673Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nos.environ[\"DISABLE_OPENMM\"] = \"1\"\nprint(\"OpenMM disabled\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-18T07:11:00.493472Z","iopub.execute_input":"2026-01-18T07:11:00.493768Z","iopub.status.idle":"2026-01-18T07:11:00.498096Z","shell.execute_reply.started":"2026-01-18T07:11:00.49374Z","shell.execute_reply":"2026-01-18T07:11:00.497404Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import sys\nsys.path.append(\"/kaggle/input/rhofold-main/RhoFold-main\")\n\nimport rhofold\nprint(\"✅ RhOFold import successful\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-18T07:11:00.498996Z","iopub.execute_input":"2026-01-18T07:11:00.499247Z","iopub.status.idle":"2026-01-18T07:11:07.499266Z","shell.execute_reply.started":"2026-01-18T07:11:00.499227Z","shell.execute_reply":"2026-01-18T07:11:07.498523Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!cp -r /kaggle/input/rhofold-main/RhoFold-main /kaggle/working/RhoFold-main\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-18T07:11:07.500284Z","iopub.execute_input":"2026-01-18T07:11:07.500745Z","iopub.status.idle":"2026-01-18T07:11:07.977191Z","shell.execute_reply.started":"2026-01-18T07:11:07.500719Z","shell.execute_reply":"2026-01-18T07:11:07.976212Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!ls /kaggle/working/RhoFold-main\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-18T07:11:07.97855Z","iopub.execute_input":"2026-01-18T07:11:07.978823Z","iopub.status.idle":"2026-01-18T07:11:08.101758Z","shell.execute_reply.started":"2026-01-18T07:11:07.978787Z","shell.execute_reply":"2026-01-18T07:11:08.101072Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from pathlib import Path\n\ninference_file = Path(\"/kaggle/working/RhoFold-main/inference.py\")\ntext = inference_file.read_text()\n\ntext = text.replace(\n    \"from rhofold.relax.relax import AmberRelaxation\",\n    \"# from rhofold.relax.relax import AmberRelaxation  # DISABLED FOR KAGGLE\"\n)\n\ninference_file.write_text(text)\nprint(\"✅ Patched inference.py\")\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-18T07:11:08.103601Z","iopub.execute_input":"2026-01-18T07:11:08.103908Z","iopub.status.idle":"2026-01-18T07:11:08.109859Z","shell.execute_reply.started":"2026-01-18T07:11:08.10388Z","shell.execute_reply":"2026-01-18T07:11:08.109024Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"relax_file = Path(\"/kaggle/working/RhoFold-main/rhofold/relax/relax.py\")\ntext = relax_file.read_text()\n\ntext = text.replace(\n    \"from simtk.openmm.app import *\",\n    \"# from simtk.openmm.app import *  # DISABLED FOR KAGGLE\"\n)\n\ntext = text.replace(\n    \"from simtk.openmm import *\",\n    \"# from simtk.openmm import *  # DISABLED FOR KAGGLE\"\n)\n\nrelax_file.write_text(text)\nprint(\"✅ Patched relax.py\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-18T07:11:08.110603Z","iopub.execute_input":"2026-01-18T07:11:08.110884Z","iopub.status.idle":"2026-01-18T07:11:08.126162Z","shell.execute_reply.started":"2026-01-18T07:11:08.110851Z","shell.execute_reply":"2026-01-18T07:11:08.125415Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nos.environ[\"DISABLE_OPENMM\"] = \"1\"\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-18T07:11:09.216791Z","iopub.execute_input":"2026-01-18T07:11:09.217098Z","iopub.status.idle":"2026-01-18T07:11:09.22083Z","shell.execute_reply.started":"2026-01-18T07:11:09.217072Z","shell.execute_reply":"2026-01-18T07:11:09.220244Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nfasta_path = \"/kaggle/working/test_rna.fasta\"\n\nwith open(fasta_path, \"w\") as f:\n    f.write(\">test_rna\\n\")\n    f.write(\"GGGAAAUCC\\n\")\n\nprint(\"✅ FASTA file created:\", fasta_path)\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-18T07:11:11.173638Z","iopub.execute_input":"2026-01-18T07:11:11.17428Z","iopub.status.idle":"2026-01-18T07:11:11.17902Z","shell.execute_reply.started":"2026-01-18T07:11:11.174253Z","shell.execute_reply":"2026-01-18T07:11:11.178247Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!cd /kaggle/working/RhoFold-main && \\\npython inference.py \\\n  --device cpu \\\n  --input_fas /kaggle/working/test_rna.fasta \\\n  --output_dir /kaggle/working/test_output \\\n  --single_seq_pred true \\\n  --relax_steps 0 \\\n  --ckpt /kaggle/input/rhofold-pretrained-weights/rhofold_pretrained.pt\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-18T07:11:13.204816Z","iopub.execute_input":"2026-01-18T07:11:13.205137Z","iopub.status.idle":"2026-01-18T07:12:02.214324Z","shell.execute_reply.started":"2026-01-18T07:11:13.205109Z","shell.execute_reply":"2026-01-18T07:12:02.213594Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"with open(\"/kaggle/working/test_output/unrelaxed_model.pdb\") as f:\n    for _ in range(5):\n        print(f.readline().strip())\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-18T07:12:02.215983Z","iopub.execute_input":"2026-01-18T07:12:02.216257Z","iopub.status.idle":"2026-01-18T07:12:02.221884Z","shell.execute_reply.started":"2026-01-18T07:12:02.216228Z","shell.execute_reply":"2026-01-18T07:12:02.221117Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport subprocess\nimport textwrap\n\nRHOFOLD_DIR = \"/kaggle/working/RhoFold-main\"\nWEIGHTS = \"/kaggle/input/rhofold-pretrained-weights/rhofold_pretrained.pt\"\n\n# Create test FASTA\nfasta = \"/kaggle/working/test.fasta\"\nwith open(fasta, \"w\") as f:\n    f.write(\">test\\nGGGGAAAACCCC\\n\")\n\ncmd = [\n    \"python\", \"inference.py\",\n    \"--input_fas\", fasta,\n    \"--output_dir\", \"/kaggle/working/test_out\",\n    \"--ckpt\", WEIGHTS,\n    \"--device\", \"cpu\",\n    \"--single_seq_pred\", \"1\",\n    \"--relax_steps\", \"0\"\n]\n\n\nprint(\"Running RhoFold...\")\nresult = subprocess.run(\n    cmd,\n    cwd=RHOFOLD_DIR,\n    capture_output=True,\n    text=True\n)\n\nprint(result.stdout)\nprint(result.stderr)\n\nassert os.path.exists(\"/kaggle/working/test_out/unrelaxed_model.pdb\")\nprint(\"✅ RhoFold WORKING\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-18T07:12:02.222903Z","iopub.execute_input":"2026-01-18T07:12:02.223509Z","iopub.status.idle":"2026-01-18T07:12:43.647635Z","shell.execute_reply.started":"2026-01-18T07:12:02.223488Z","shell.execute_reply":"2026-01-18T07:12:43.646872Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"confirming rhofold directories ","metadata":{}},{"cell_type":"code","source":"# ============================================================\n# SETUP RHOFOLD - Extract to working directory\n# ============================================================\n\nimport os\nimport shutil\nimport zipfile\n\nprint(\"=\"*60)\nprint(\"RHOFOLD SETUP\")\nprint(\"=\"*60)\n\n# Source path (where you added RHOfold dataset)\nsource_path = \"/kaggle/input/rhofold-main/RhoFold-main\"\n\n# Destination path (working directory)\ndest_path = \"/kaggle/working/RhoFold-main\"\n\nprint(f\"\\nSource: {source_path}\")\nprint(f\"Destination: {dest_path}\")\n\n# Check if source exists\nif not os.path.exists(source_path):\n    print(f\"\\n✗ Source not found!\")\n    print(\"\\nExpected paths:\")\n    print(\"  /kaggle/input/rhofold-main/RhoFold-main\")\n    print(\"  /kaggle/input/rhofold-main/RhoFold-main.zip\")\n    \n    # Try to find it\n    print(\"\\nSearching for RHOfold...\")\n    for root, dirs, files in os.walk(\"/kaggle/input\"):\n        for d in dirs:\n            if \"rhofold\" in d.lower():\n                print(f\"  Found: {os.path.join(root, d)}\")\n        for f in files:\n            if \"rhofold\" in f.lower():\n                print(f\"  Found: {os.path.join(root, f)}\")\n    exit(1)\n\n# Check if already exists\nif os.path.exists(dest_path):\n    print(f\"\\n✓ RHOfold already exists at {dest_path}\")\n    print(\"  Skipping extraction\")\nelse:\n    # Copy RHOfold to working directory\n    print(f\"\\nCopying RHOfold to working directory...\")\n    \n    try:\n        shutil.copytree(source_path, dest_path)\n        print(\"✓ Copy complete\")\n    except Exception as e:\n        print(f\"✗ Copy failed: {e}\")\n        exit(1)\n\n# Verify critical files\nprint(\"\\nVerifying installation...\")\ncritical_files = [\n    \"inference.py\",\n    \"rhofold/__init__.py\",\n    \"rhofold/model/__init__.py\",\n    \"rhofold/utils/__init__.py\"\n]\n\nall_good = True\nfor file in critical_files:\n    path = os.path.join(dest_path, file)\n    exists = os.path.exists(path)\n    status = \"✓\" if exists else \"✗\"\n    print(f\"  {status} {file}\")\n    if not exists:\n        all_good = False\n\n# Check weights\nweights_path = \"/kaggle/input/rhofold-pretrained-weights/rhofold_pretrained.pt\"\nweights_exists = os.path.exists(weights_path)\nprint(f\"\\n  {'✓' if weights_exists else '✗'} Weights: {weights_path}\")\n\nif not weights_exists:\n    print(\"\\n✗ Weights not found!\")\n    print(\"  Make sure to add 'rhofold-pretrained-weights' dataset\")\n    all_good = False\n\nif all_good:\n    print(\"\\n\" + \"=\"*60)\n    print(\"✅ SETUP COMPLETE\")\n    print(\"=\"*60)\n    print(\"\\nRHOfold is ready to use!\")\n    print(\"You can now run the main submission script.\")\nelse:\n    print(\"\\n\" + \"=\"*60)\n    print(\"❌ SETUP INCOMPLETE\")\n    print(\"=\"*60)\n    print(\"\\nSome files are missing. Check the errors above.\")\n\n# Show structure\nprint(\"\\nDirectory structure:\")\nif os.path.exists(dest_path):\n    for item in os.listdir(dest_path)[:10]:\n        print(f\"  - {item}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-18T07:12:43.649311Z","iopub.execute_input":"2026-01-18T07:12:43.649582Z","iopub.status.idle":"2026-01-18T07:12:43.662898Z","shell.execute_reply.started":"2026-01-18T07:12:43.64956Z","shell.execute_reply":"2026-01-18T07:12:43.662335Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"---------------------------------------------------------------------------------------------------","metadata":{}},{"cell_type":"markdown","source":"# This script implements a speed-optimized RNA 3D structure prediction pipeline for the Stanford RNA 3D Folding Kaggle competition.\n\nIt generates 5 alternative 3D structures per RNA sequence using a hybrid strategy:\n\nRhoFold deep-learning predictions (highest quality, when feasible)\n\nTemplate-based structure transfer from known RNA structures\n\nGeometric fallback models (RNA helix) to guarantee valid outputs\n\nThe primary goal is robust submission generation under strict runtime constraints.","metadata":{}},{"cell_type":"code","source":"import os\nimport shutil\nimport subprocess\nimport tempfile\nimport warnings\nimport numpy as np\nimport pandas as pd\nimport time\nfrom pathlib import Path\nfrom scipy.spatial.distance import cdist\nimport gc\n\nwarnings.filterwarnings(\"ignore\")\n\n# ============================================================\n# CONFIG - OPTIMIZED FOR SCORING\n# ============================================================\n\nDATA_PATH = \"/kaggle/input/stanford-rna-3d-folding-2\"\nRHOFOLD_DIR = \"/kaggle/working/RhoFold-main\"\nRHOFOLD_WEIGHTS = \"/kaggle/input/rhofold-pretrained-weights/rhofold_pretrained.pt\"\nPDB_DIR = f\"{DATA_PATH}/PDB_RNA\"\nOUTPUT_PATH = \"/kaggle/working/submission.csv\"\n\nMSA_DIR = f\"{DATA_PATH}/MSA\"\nN_PREDICTIONS = 5\nSEED = 42\nMAX_RHOFOLD_LEN = 500\nRHOFOLD_TIMEOUT = 400\n\n# Balanced memory/performance\nMAX_TRAIN_TEMPLATES = 600\nMAX_PDB_TEMPLATES = 300  # Add PDB templates!\nMAX_PDB_CACHE = 50\nCHUNK_SIZE = 5000\n\n# Enhanced structural constraints\nBOND_LENGTH = 5.9\nBOND_TOL = 0.8\nMIN_CLASH_DIST = 3.5\nMAX_ITER_REFINE = 150  # Increased for better quality\nLEARNING_RATE = 0.16  # Slower but more stable\n\nDEVICE = \"cuda\" if os.path.exists(\"/dev/nvidia0\") else \"cpu\"\nnp.random.seed(SEED)\n\n# ============================================================\n# CIF PARSER WITH QUALITY FILTERING\n# ============================================================\n\nclass CIFStructureCache:\n    \"\"\"Memory-efficient CIF cache with quality filtering\"\"\"\n    \n    def __init__(self, max_cache=MAX_PDB_CACHE):\n        self.pdb_files = {}\n        self.cache = {}\n        self.max_cache = max_cache\n        self.access_count = {}\n        \n        if os.path.exists(PDB_DIR):\n            print(\"Indexing CIF files... \", end=\"\", flush=True)\n            start = time.time()\n            for cif_file in Path(PDB_DIR).glob(\"*.cif\"):\n                pdb_id = cif_file.stem.split('.')[0]\n                self.pdb_files[pdb_id] = str(cif_file)\n            print(f\"{len(self.pdb_files)} indexed in {time.time()-start:.1f}s\")\n    \n    def parse_cif_with_sequence(self, filepath):\n        \"\"\"Parse CIF and extract both coords and sequence - ROBUST\"\"\"\n        coords = []\n        sequence = []\n        \n        res_map = {\n            'A': 'A', 'C': 'C', 'G': 'G', 'U': 'U',\n            'ADE': 'A', 'CYT': 'C', 'GUA': 'G', 'URA': 'U',\n            'DA': 'A', 'DC': 'C', 'DG': 'G', 'DT': 'U',\n            'T': 'U', 'THY': 'U'\n        }\n        \n        try:\n            with open(filepath, 'r') as f:\n                in_atom_section = False\n                prev_resid = None\n                \n                for line in f:\n                    # Detect atom section\n                    if 'ATOM' in line and 'type_symbol' in line:\n                        in_atom_section = True\n                        continue\n                    \n                    if not in_atom_section and not line.startswith('ATOM'):\n                        continue\n                    \n                    if line.startswith('ATOM') or (in_atom_section and len(line.strip()) > 0):\n                        parts = line.split()\n                        if len(parts) < 17:\n                            continue\n                        \n                        # Find C1' atom\n                        atom_name = parts[3] if len(parts) > 3 else ''\n                        if atom_name != \"C1'\" and \"C1'\" not in atom_name:\n                            continue\n                        \n                        try:\n                            # Parse coordinates\n                            x = float(parts[10])\n                            y = float(parts[11])\n                            z = float(parts[12])\n                            \n                            # Parse residue\n                            residue = parts[5] if len(parts) > 5 else parts[4]\n                            current_resid = parts[8] if len(parts) > 8 else parts[6]\n                            \n                            # Avoid duplicates from same residue\n                            if current_resid == prev_resid:\n                                continue\n                            prev_resid = current_resid\n                            \n                            # Map residue to base\n                            base = res_map.get(residue.upper(), None)\n                            if base:\n                                coords.append([x, y, z])\n                                sequence.append(base)\n                        except (ValueError, IndexError):\n                            continue\n            \n            if len(coords) >= 10:  # Minimum useful length\n                arr = np.array(coords, dtype=np.float32)\n                seq = ''.join(sequence)\n                # Allow a few unknowns\n                if seq.count('N') < len(seq) * 0.1:  # Less than 10% unknown\n                    return arr, seq.replace('N', 'A')  # Replace unknowns with A\n        except Exception as e:\n            pass\n        \n        return None, None\n    \n    def get_structure(self, pdb_id):\n        \"\"\"Get structure with LRU caching\"\"\"\n        if pdb_id in self.cache:\n            self.access_count[pdb_id] += 1\n            return self.cache[pdb_id]\n        \n        if pdb_id not in self.pdb_files:\n            return None, None\n        \n        coords, seq = self.parse_cif_with_sequence(self.pdb_files[pdb_id])\n        \n        if coords is not None:\n            if len(self.cache) >= self.max_cache:\n                lru_id = min(self.access_count.items(), key=lambda x: x[1])[0]\n                del self.cache[lru_id]\n                del self.access_count[lru_id]\n            \n            self.cache[pdb_id] = (coords, seq)\n            self.access_count[pdb_id] = 1\n        \n        return coords, seq\n\n# ============================================================\n# MSA CACHE\n# ============================================================\n\nclass MSACache:\n    def __init__(self):\n        self.cache = {}\n        if os.path.exists(MSA_DIR):\n            print(\"Indexing MSA files... \", end=\"\", flush=True)\n            start = time.time()\n            for msa_file in Path(MSA_DIR).glob(\"*.fasta\"):\n                target_id = msa_file.stem.split('.')[0]\n                self.cache[target_id] = str(msa_file)\n            print(f\"{len(self.cache)} indexed in {time.time()-start:.1f}s\")\n    \n    def get(self, target_id):\n        return self.cache.get(target_id)\n\n# ============================================================\n# QUALITY METRICS\n# ============================================================\n\ndef compute_structure_quality_fast(coords):\n    \"\"\"Fast quality check with bond geometry\"\"\"\n    if len(coords) < 2:\n        return 0.0\n    \n    bond_lengths = np.linalg.norm(coords[1:] - coords[:-1], axis=1)\n    \n    # Check bond length consistency (ideal ~5.9Å)\n    bond_dev = np.abs(bond_lengths - BOND_LENGTH)\n    bond_score = np.mean(bond_dev < 1.5)  # Fraction of good bonds\n    \n    # Check for severe clashes\n    if len(coords) > 3:\n        min_sep = np.min(bond_lengths)\n        if min_sep < 2.0:  # Severe clash\n            return 0.0\n    \n    return bond_score\n\n# ============================================================\n# TEMPLATE LOADING WITH PDB\n# ============================================================\n\ndef load_templates_hybrid(cif_cache, max_train=MAX_TRAIN_TEMPLATES, max_pdb=MAX_PDB_TEMPLATES):\n    \"\"\"Load templates from both train data and PDB\"\"\"\n    print(\"Loading templates... \", end=\"\", flush=True)\n    start = time.time()\n    \n    templates = []\n    \n    # 1. Load train templates\n    train_df = pd.read_csv(f\"{DATA_PATH}/train_sequences.csv\")\n    \n    chunk_iter = pd.read_csv(\n        f\"{DATA_PATH}/train_labels.csv\",\n        chunksize=CHUNK_SIZE,\n        dtype={'x_1': np.float32, 'y_1': np.float32, 'z_1': np.float32}\n    )\n    \n    target_coords = {}\n    for chunk in chunk_iter:\n        chunk['target_id'] = chunk['ID'].str.split('_').str[0]\n        for tid, group in chunk.groupby('target_id'):\n            if tid not in target_coords:\n                coords = group[['x_1', 'y_1', 'z_1']].values.astype(np.float32)\n                target_coords[tid] = coords\n            if len(target_coords) >= max_train:\n                break\n        if len(target_coords) >= max_train:\n            break\n        del chunk\n        gc.collect()\n    \n    for _, row in train_df.iterrows():\n        tid = row['target_id']\n        if tid not in target_coords:\n            continue\n        \n        coords = target_coords[tid]\n        seq = row['sequence']\n        \n        if len(coords) != len(seq) or len(coords) < 5:\n            continue\n        \n        quality = compute_structure_quality_fast(coords)\n        \n        if quality > 0.5:\n            templates.append({\n                'coords': coords,\n                'sequence': seq,\n                'length': len(seq),\n                'quality': quality,\n                'source': 'train'\n            })\n        \n        if len(templates) >= max_train:\n            break\n    \n    del target_coords\n    gc.collect()\n    \n    train_count = len(templates)\n    \n    # 2. Load high-quality PDB templates (with progress)\n    print(f\"{train_count} train\", end=\"\", flush=True)\n    \n    pdb_ids = list(cif_cache.pdb_files.keys())\n    np.random.shuffle(pdb_ids)  # Random sampling\n    \n    pdb_loaded = 0\n    pdb_tried = 0\n    pdb_failed = 0\n    \n    # Try more files but with lower quality threshold\n    for pdb_id in pdb_ids[:min(800, len(pdb_ids))]:  # Check first 800\n        if pdb_loaded >= max_pdb:\n            break\n        \n        pdb_tried += 1\n        \n        # Progress indicator\n        if pdb_tried % 100 == 0:\n            print(\".\", end=\"\", flush=True)\n        \n        coords, seq = cif_cache.get_structure(pdb_id)\n        if coords is None or seq is None:\n            pdb_failed += 1\n            continue\n        \n        if len(coords) < 15 or len(coords) > 500:  # Useful size range\n            continue\n        \n        quality = compute_structure_quality_fast(coords)\n        \n        # Lower threshold to get more templates\n        if quality > 0.55:  # Reduced from 0.65\n            templates.append({\n                'coords': coords,\n                'sequence': seq,\n                'length': len(seq),\n                'quality': quality,\n                'source': 'pdb'\n            })\n            pdb_loaded += 1\n    \n    # Sort by quality\n    templates.sort(key=lambda x: x['quality'], reverse=True)\n    \n    print(f\" → {len(templates)} total ({train_count} train, {pdb_loaded} PDB, {pdb_failed} failed) in {time.time()-start:.1f}s\")\n    return templates\n\n# ============================================================\n# TEMPLATE MATCHING\n# ============================================================\n\ndef sequence_similarity_fast(seq1, seq2):\n    \"\"\"Fast sequence similarity with alignment-like scoring\"\"\"\n    n1, n2 = len(seq1), len(seq2)\n    if n1 == 0 or n2 == 0:\n        return 0.0\n    \n    min_len = min(n1, n2)\n    matches = sum(1 for i in range(min_len) if seq1[i] == seq2[i])\n    \n    # Penalize length difference\n    len_penalty = abs(n1 - n2) / max(n1, n2)\n    \n    return (matches / max(n1, n2)) * (1.0 - 0.3 * len_penalty)\n\ndef select_best_templates(query_seq, templates, k=8):\n    \"\"\"Select best templates with scoring\"\"\"\n    if not templates:\n        return []\n    \n    seq_len = len(query_seq)\n    scored = []\n    \n    for t in templates:\n        len_ratio = min(seq_len, t['length']) / max(seq_len, t['length'])\n        if len_ratio < 0.4:\n            continue\n        \n        seq_sim = sequence_similarity_fast(query_seq, t['sequence'])\n        \n        # Boost PDB templates slightly\n        source_bonus = 1.1 if t['source'] == 'pdb' else 1.0\n        \n        score = (0.5 * seq_sim + 0.4 * t['quality'] + 0.1 * len_ratio) * source_bonus\n        \n        scored.append((score, t))\n    \n    scored.sort(key=lambda x: x[0], reverse=True)\n    return scored[:k]\n\n# ============================================================\n# INTERPOLATION\n# ============================================================\n\ndef interpolate_coords(template_coords, target_len):\n    \"\"\"Smooth interpolation\"\"\"\n    if len(template_coords) == target_len:\n        return template_coords.copy()\n    \n    indices = np.linspace(0, len(template_coords) - 1, target_len)\n    result = np.empty((target_len, 3), dtype=np.float32)\n    \n    for dim in range(3):\n        result[:, dim] = np.interp(indices, np.arange(len(template_coords)), \n                                   template_coords[:, dim])\n    \n    return result\n\n# ============================================================\n# ENHANCED REFINEMENT\n# ============================================================\n\ndef refine_structure_enhanced(coords, sequence, max_iter=MAX_ITER_REFINE):\n    \"\"\"Enhanced refinement with base pairing\"\"\"\n    refined = coords.copy().astype(np.float64)\n    n = len(sequence)\n    \n    base_pairs = {'A': 'U', 'U': 'A', 'G': 'C', 'C': 'G'}\n    \n    for iteration in range(max_iter):\n        forces = np.zeros_like(refined)\n        \n        # 1. Bond length forces (strong)\n        for i in range(n - 1):\n            vec = refined[i + 1] - refined[i]\n            dist = np.linalg.norm(vec)\n            \n            if dist > 1e-6:\n                error = dist - BOND_LENGTH\n                force_mag = error / dist\n                force = force_mag * vec * 1.0\n                forces[i] += force\n                forces[i + 1] -= force\n        \n        # 2. Clash avoidance (medium)\n        for i in range(n):\n            for j in range(i + 2, min(i + 15, n)):\n                vec = refined[j] - refined[i]\n                dist = np.linalg.norm(vec)\n                \n                if dist < MIN_CLASH_DIST and dist > 1e-6:\n                    force_mag = (MIN_CLASH_DIST - dist) / (dist ** 2)\n                    force = force_mag * vec / dist * 0.4\n                    forces[i] -= force\n                    forces[j] += force\n        \n        # 3. Base pairing forces (weak but important)\n        if iteration % 2 == 0:  # Every other iteration\n            for i in range(n):\n                if sequence[i] not in base_pairs:\n                    continue\n                partner = base_pairs[sequence[i]]\n                \n                for j in range(i + 4, min(i + 25, n)):\n                    if sequence[j] == partner:\n                        vec = refined[j] - refined[i]\n                        dist = np.linalg.norm(vec)\n                        \n                        if 7.0 < dist < 15.0:\n                            ideal = 10.5\n                            error = dist - ideal\n                            force_mag = error * 0.12  # Stronger base pairing\n                            force = force_mag * vec / (dist + 1e-6)\n                            forces[i] += force\n                            forces[j] -= force\n                        break\n        \n        # Update with annealing\n        lr = LEARNING_RATE * (1.0 - 0.7 * iteration / max_iter)\n        refined += forces * lr\n        \n        # Early stopping\n        if iteration % 30 == 0:\n            force_norm = np.linalg.norm(forces)\n            if force_norm < 0.015:\n                break\n    \n    return refined.astype(np.float32)\n\n# ============================================================\n# RHOFOLD\n# ============================================================\n\nclass RHoFoldPredictor:\n    def __init__(self, msa_cache):\n        self.available = (\n            os.path.exists(RHOFOLD_DIR)\n            and os.path.exists(RHOFOLD_WEIGHTS)\n            and os.path.exists(os.path.join(RHOFOLD_DIR, \"inference.py\"))\n        )\n        self.msa_cache = msa_cache\n        self.success = 0\n        self.fail = 0\n        \n        if self.available:\n            print(f\"✓ RhoFold available (device: {DEVICE})\")\n        else:\n            print(\"✗ RhoFold NOT available\")\n\n    def predict(self, seq, target_id):\n        if not self.available or len(seq) > MAX_RHOFOLD_LEN:\n            return None\n\n        msa_file = self.msa_cache.get(target_id)\n        tmp = tempfile.mkdtemp(prefix=\"rhf_\")\n        \n        try:\n            fas = os.path.join(tmp, \"in.fasta\")\n            \n            if msa_file:\n                shutil.copy(msa_file, fas)\n                use_msa = True\n            else:\n                with open(fas, \"w\") as f:\n                    f.write(f\">{target_id}\\n{seq}\\n\")\n                use_msa = False\n\n            out = os.path.join(tmp, \"out\")\n            os.makedirs(out, exist_ok=True)\n\n            cmd = [\n                \"python\", \"inference.py\",\n                \"--input_fas\", fas,\n                \"--output_dir\", out,\n                \"--ckpt\", RHOFOLD_WEIGHTS,\n                \"--device\", DEVICE,\n                \"--single_seq_pred\", \"0\" if use_msa else \"1\",\n                \"--relax_steps\", \"0\",\n            ]\n\n            print(f\"[{'MSA' if use_msa else 'SEQ'}]\", end=\"\", flush=True)\n            \n            result = subprocess.run(\n                cmd, cwd=RHOFOLD_DIR, capture_output=True,\n                text=True, timeout=RHOFOLD_TIMEOUT\n            )\n\n            pdb = os.path.join(out, \"unrelaxed_model.pdb\")\n            \n            if os.path.exists(pdb):\n                coords = self._parse_pdb(pdb, len(seq))\n                if coords is not None:\n                    self.success += 1\n                    print(\"✓\", end=\" \")\n                    return coords\n            \n            self.fail += 1\n            print(\"✗\", end=\" \")\n            return None\n        except:\n            self.fail += 1\n            print(\"✗\", end=\" \")\n            return None\n        finally:\n            shutil.rmtree(tmp, ignore_errors=True)\n\n    @staticmethod\n    def _parse_pdb(pdb, n):\n        coords = []\n        with open(pdb) as f:\n            for l in f:\n                if l.startswith(\"ATOM\") and l[12:16].strip() == \"C1'\":\n                    coords.append([float(l[30:38]), float(l[38:46]), float(l[46:54])])\n        coords = np.array(coords, dtype=np.float32)\n        return coords if len(coords) == n else None\n\n# ============================================================\n# DIVERSE PREDICTION GENERATION\n# ============================================================\n\ndef generate_helix(n, radius_offset=0.0):\n    \"\"\"Generate A-form helix with variation\"\"\"\n    angle_step = 32.7 * np.pi / 180\n    radius = 9.0 + radius_offset\n    rise = 2.8\n    \n    angles = np.arange(n) * angle_step\n    coords = np.zeros((n, 3), dtype=np.float32)\n    coords[:, 0] = radius * np.cos(angles)\n    coords[:, 1] = radius * np.sin(angles)\n    coords[:, 2] = np.arange(n) * rise\n    \n    return coords\n\ndef generate_predictions(seq, templates, rhofold_pred, seed):\n    \"\"\"Generate 5 diverse, high-quality predictions\"\"\"\n    preds = []\n    n = len(seq)\n    rng = np.random.default_rng(seed)\n    \n    # Pred 1: RhoFold with strong refinement\n    if rhofold_pred is not None:\n        # More iterations for shorter sequences (better quality)\n        iters = 130 if n < 100 else 100\n        refined = refine_structure_enhanced(rhofold_pred, seq, max_iter=iters)\n        preds.append(refined)\n        print(\"[R]\", end=\"\")\n        \n        # Pred 2: RhoFold variant\n        if len(preds) < N_PREDICTIONS:\n            variant = rhofold_pred + rng.normal(0, 0.2, rhofold_pred.shape).astype(np.float32)\n            refined = refine_structure_enhanced(variant, seq, max_iter=iters-20)\n            preds.append(refined)\n            print(\"[Rv]\", end=\"\")\n    \n    # Preds from templates (up to 3)\n    if templates:\n        for i, (score, tmpl) in enumerate(templates[:min(3, N_PREDICTIONS - len(preds))]):\n            coords = interpolate_coords(tmpl['coords'], n)\n            \n            # Smaller perturbation for high-quality templates\n            noise_scale = 0.15 if tmpl['quality'] > 0.8 else 0.3\n            coords += rng.normal(0, noise_scale, coords.shape).astype(np.float32)\n            \n            # More refinement for short sequences\n            iters = 120 if n < 100 else 90\n            refined = refine_structure_enhanced(coords, seq, max_iter=iters)\n            preds.append(refined)\n            src = 'P' if tmpl['source'] == 'pdb' else 'T'\n            print(f\"[{src}{i+1}]\", end=\"\")\n    \n    # Fill remaining with geometry\n    while len(preds) < N_PREDICTIONS:\n        idx = len(preds)\n        coords = generate_helix(n, radius_offset=idx * 0.4)\n        coords += rng.normal(0, 0.7 + idx * 0.2, coords.shape).astype(np.float32)\n        refined = refine_structure_enhanced(coords, seq, max_iter=70)\n        preds.append(refined)\n        print(f\"[G{idx}]\", end=\"\")\n    \n    return preds[:N_PREDICTIONS]\n\n# ============================================================\n# MAIN PREDICTION\n# ============================================================\n\ndef predict_sequence(seq, target_id, templates, rhofold):\n    \"\"\"Predict structure\"\"\"\n    # RhoFold\n    rhofold_pred = None\n    if rhofold.available and len(seq) <= MAX_RHOFOLD_LEN:\n        rhofold_pred = rhofold.predict(seq, target_id)\n    \n    # Templates\n    matched = select_best_templates(seq, templates, k=8)\n    if matched:\n        top_score = matched[0][0]\n        print(f\"[T:{len(matched)},S:{top_score:.2f}]\", end=\" \")\n    \n    # Generate\n    seed = hash(target_id) % (2**32)\n    preds = generate_predictions(seq, matched, rhofold_pred, seed)\n    \n    print(\" ✓\")\n    return preds\n\n# ============================================================\n# MAIN\n# ============================================================\n\ndef create_submission():\n    print(\"=\"*70)\n    print(\"  STANFORD RNA 3D FOLDING - ENHANCED v3\")\n    print(\"=\"*70)\n    \n    print(\"\\n[1/4] Initializing\")\n    msa_cache = MSACache()\n    cif_cache = CIFStructureCache(max_cache=MAX_PDB_CACHE)\n    rhofold = RHoFoldPredictor(msa_cache)\n    \n    print(\"\\n[2/4] Loading templates\")\n    templates = load_templates_hybrid(cif_cache, max_train=MAX_TRAIN_TEMPLATES, max_pdb=MAX_PDB_TEMPLATES)\n    \n    print(f\"\\n[3/4] Loading test data\")\n    test_df = pd.read_csv(f\"{DATA_PATH}/test_sequences.csv\")\n    sub = pd.read_csv(f\"{DATA_PATH}/sample_submission.csv\")\n    print(f\"Test sequences: {len(test_df)}\")\n    \n    print(\"\\n[4/4] Generating predictions\")\n    print(\"=\"*70)\n    \n    row = 0\n    start_time = time.time()\n    \n    for i, r in test_df.iterrows():\n        seq = r[\"sequence\"]\n        target_id = r[\"target_id\"]\n        print(f\"[{i+1:2d}/{len(test_df)}] {target_id:10s} L={len(seq):4d} \", end=\"\")\n        \n        preds = predict_sequence(seq, target_id, templates, rhofold)\n        \n        for j in range(len(seq)):\n            for k in range(N_PREDICTIONS):\n                sub.loc[row, f\"x_{k+1}\"] = preds[k][j, 0]\n                sub.loc[row, f\"y_{k+1}\"] = preds[k][j, 1]\n                sub.loc[row, f\"z_{k+1}\"] = preds[k][j, 2]\n            row += 1\n        \n        if (i + 1) % 5 == 0:\n            gc.collect()\n    \n    elapsed = time.time() - start_time\n    \n    print(\"\\n\" + \"=\"*70)\n    print(\"  SUMMARY\")\n    print(\"=\"*70)\n    total = rhofold.success + rhofold.fail\n    if total > 0:\n        print(f\"RhoFold: {rhofold.success}/{total} ({100*rhofold.success/total:.0f}%)\")\n    print(f\"Templates: {len(templates)}\")\n    print(f\"Time: {elapsed/60:.1f}min ({elapsed/len(test_df):.1f}s/seq)\")\n    \n    # Validate\n    coord_cols = [f\"{a}_{i}\" for i in range(1, N_PREDICTIONS+1) for a in \"xyz\"]\n    if not np.isfinite(sub[coord_cols].values).all():\n        print(\"\\n⚠ Fixing invalid values...\")\n        for col in coord_cols:\n            mask = ~np.isfinite(sub[col])\n            if mask.any():\n                sub.loc[mask, col] = 0.0\n    \n    sub.to_csv(OUTPUT_PATH, index=False)\n    print(f\"\\n✅ Saved: {OUTPUT_PATH}\")\n\nif __name__ == \"__main__\":\n    create_submission()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-18T07:12:43.664056Z","iopub.execute_input":"2026-01-18T07:12:43.664272Z","iopub.status.idle":"2026-01-18T07:41:50.032381Z","shell.execute_reply.started":"2026-01-18T07:12:43.664252Z","shell.execute_reply":"2026-01-18T07:41:50.031636Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Visualisation of RNA Chains ","metadata":{}},{"cell_type":"code","source":"!pip install /kaggle/input/kaleido/kaleido-0.2.1-py2.py3-none-manylinux1_x86_64.whl","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-18T07:45:21.117476Z","iopub.execute_input":"2026-01-18T07:45:21.118186Z","iopub.status.idle":"2026-01-18T07:45:26.762946Z","shell.execute_reply.started":"2026-01-18T07:45:21.11816Z","shell.execute_reply":"2026-01-18T07:45:26.76222Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =======================\n# RNA 3D STRUCTURE VISUALIZATION\n# Kaggle-Compatible Version with Static Images\n# =======================\n\nimport pandas as pd\nimport numpy as np\nimport plotly.graph_objects as go\nfrom plotly.subplots import make_subplots\nimport plotly.io as pio\nfrom IPython.display import Image, display, HTML\nimport os\n\n# =======================\n# SETTINGS\n# =======================\n\nSUBMISSION_PATH = \"/kaggle/working/submission.csv\"\nOUTPUT_DIR = \"/kaggle/working/visualizations\"\n\n# Create output directory\nos.makedirs(OUTPUT_DIR, exist_ok=True)\n\n# Target structures to visualize (in 2x2 grid)\ntargets = [\n    (\"8ZNQ\", 15),      # Small RNA (30 nt)\n    (\"9E74\", 128),     # Medium RNA (255 nt)\n    (\"9JGM\", 105),     # Medium RNA (210 nt)\n    (\"9MME\", 2320),    # Large RNA (4640 nt)\n]\n\n# Nucleotide colors\nnt_color = {\n    \"A\": \"#FF6B6B\",  # Red\n    \"U\": \"#4ECDC4\",  # Teal\n    \"G\": \"#FFE66D\",  # Yellow\n    \"C\": \"#95E1D3\",  # Light blue\n    \"N\": \"#CCCCCC\",  # Gray\n}\n\n# =======================\n# LOAD DATA\n# =======================\n\nprint(\"=\"*70)\nprint(\"  RNA 3D STRUCTURE VISUALIZATION\")\nprint(\"=\"*70)\n\nprint(\"\\nLoading submission data...\")\ndf = pd.read_csv(SUBMISSION_PATH)\nprint(f\"✓ Loaded {len(df)} residues\")\n\nunique_targets = df[\"ID\"].str.split(\"_\").str[0].unique()\nprint(f\"✓ Found {len(unique_targets)} unique targets\")\n\n# =======================\n# HELPER FUNCTIONS\n# =======================\n\ndef extract_chain(df, target_id, pred_idx=1):\n    \"\"\"Extract coordinates and metadata for a target\"\"\"\n    sub = df[df[\"ID\"].str.startswith(target_id + \"_\")].sort_values(\"resid\")\n    \n    coords = sub[[f\"x_{pred_idx}\", f\"y_{pred_idx}\", f\"z_{pred_idx}\"]].values\n    nts = sub[\"resname\"].values\n    resids = sub[\"resid\"].values\n    \n    return coords, nts, resids\n\n# =======================\n# CREATE 2×2 3D VISUALIZATION\n# =======================\n\nprint(\"\\n\" + \"=\"*70)\nprint(\"  CREATING 2×2 3D STRUCTURE GRID\")\nprint(\"=\"*70)\n\nfig = make_subplots(\n    rows=2,\n    cols=2,\n    specs=[[{\"type\": \"scene\"}, {\"type\": \"scene\"}],\n           [{\"type\": \"scene\"}, {\"type\": \"scene\"}]],\n    subplot_titles=[f\"<b>{t}</b> (L={len(extract_chain(df, t)[0])} nt)\" \n                   for t, r in targets],\n    vertical_spacing=0.1,\n    horizontal_spacing=0.05\n)\n\nfor i, (target_id, highlight_res) in enumerate(targets):\n    row = i // 2 + 1\n    col = i % 2 + 1\n    \n    coords, nts, resids = extract_chain(df, target_id, pred_idx=1)\n    \n    print(f\"  [{i+1}/4] Processing {target_id}: {len(coords)} residues\")\n    \n    # Backbone line\n    fig.add_trace(\n        go.Scatter3d(\n            x=coords[:, 0],\n            y=coords[:, 1],\n            z=coords[:, 2],\n            mode=\"lines\",\n            line=dict(color=\"rgba(50,50,50,0.6)\", width=4),\n            showlegend=False,\n            hoverinfo=\"skip\",\n        ),\n        row=row, col=col\n    )\n    \n    # Nucleotide markers\n    colors = [nt_color.get(n, nt_color[\"N\"]) for n in nts]\n    \n    fig.add_trace(\n        go.Scatter3d(\n            x=coords[:, 0],\n            y=coords[:, 1],\n            z=coords[:, 2],\n            mode=\"markers\",\n            marker=dict(\n                size=6,\n                color=colors,\n                line=dict(width=1, color=\"white\")\n            ),\n            text=[f\"{n}{r}\" for n, r in zip(nts, resids)],\n            hovertemplate=\"<b>%{text}</b><br>\" +\n                         \"x: %{x:.2f}<br>\" +\n                         \"y: %{y:.2f}<br>\" +\n                         \"z: %{z:.2f}<br>\" +\n                         \"<extra></extra>\",\n            showlegend=False,\n        ),\n        row=row, col=col\n    )\n    \n    # Highlight residue\n    idx = np.where(resids == highlight_res)[0]\n    if len(idx) > 0:\n        j = idx[0]\n        fig.add_trace(\n            go.Scatter3d(\n                x=[coords[j, 0]],\n                y=[coords[j, 1]],\n                z=[coords[j, 2]],\n                mode=\"markers\",\n                marker=dict(\n                    size=14,\n                    color=\"black\",\n                    symbol=\"diamond\",\n                    line=dict(width=3, color=\"yellow\")\n                ),\n                text=[f\"<b>{nts[j]}{resids[j]}</b>\"],\n                hoverinfo=\"text\",\n                showlegend=False,\n            ),\n            row=row, col=col\n        )\n    \n    # Update scene\n    scene_dict = dict(\n        aspectmode=\"data\",\n        xaxis=dict(showbackground=False, showgrid=True, gridwidth=1, gridcolor=\"lightgray\"),\n        yaxis=dict(showbackground=False, showgrid=True, gridwidth=1, gridcolor=\"lightgray\"),\n        zaxis=dict(showbackground=False, showgrid=True, gridwidth=1, gridcolor=\"lightgray\"),\n        camera=dict(eye=dict(x=1.5, y=1.5, z=1.5))\n    )\n    \n    if i == 0:\n        fig.update_layout(scene=scene_dict)\n    else:\n        fig.update_layout(**{f\"scene{i+1}\": scene_dict})\n\nfig.update_layout(\n    title=dict(\n        text=\"<b>RNA 3D Structures — RhoFold + Template Predictions</b><br>\" +\n             \"<sub>Color: A=red, U=teal, G=yellow, C=lightblue | Diamond=highlighted residue</sub>\",\n        x=0.5,\n        xanchor=\"center\",\n        font=dict(size=18)\n    ),\n    height=1000,\n    width=1400,\n    showlegend=False,\n    paper_bgcolor=\"white\",\n    plot_bgcolor=\"white\"\n)\n\n# Save as HTML\nhtml_path = f\"{OUTPUT_DIR}/structures_3d.html\"\nfig.write_html(html_path)\nprint(f\"\\n✓ Saved interactive 3D plot: {html_path}\")\n\n# Save as static image\nimg_path = f\"{OUTPUT_DIR}/structures_3d.png\"\ntry:\n    fig.write_image(img_path, width=1400, height=1000, scale=2)\n    print(f\"✓ Saved static image: {img_path}\")\n    has_image = True\nexcept Exception as e:\n    print(f\"⚠ Could not save PNG (kaleido issue): {e}\")\n    print(f\"  → HTML file saved instead: {html_path}\")\n    has_image = False\n\n# Display in notebook\nprint(\"\\n\" + \"=\"*70)\nprint(\"  DISPLAYING 3D STRUCTURES\")\nprint(\"=\"*70)\n\n# Read and display the HTML content directly\nwith open(html_path, 'r') as f:\n    html_content = f.read()\n\ndisplay(HTML(\"<h2 style='text-align: center; color: #2c3e50;'>Interactive 3D Structures</h2>\"))\ndisplay(HTML(html_content))\nprint(\"\\n✓ Interactive plot displayed above\")\n\n# =======================\n# CREATE QUALITY METRICS\n# =======================\n\nprint(\"\\n\" + \"=\"*70)\nprint(\"  CREATING QUALITY METRICS\")\nprint(\"=\"*70)\n\nfig_metrics = make_subplots(\n    rows=2,\n    cols=2,\n    subplot_titles=[f\"<b>{t}</b>\" for t, _ in targets],\n    specs=[[{\"type\": \"xy\"}, {\"type\": \"xy\"}],\n           [{\"type\": \"xy\"}, {\"type\": \"xy\"}]],\n    vertical_spacing=0.15,\n    horizontal_spacing=0.12\n)\n\nfor i, (target_id, _) in enumerate(targets):\n    row = i // 2 + 1\n    col = i % 2 + 1\n    \n    coords, _, _ = extract_chain(df, target_id, pred_idx=1)\n    \n    # Bond lengths\n    bond_lengths = np.linalg.norm(coords[1:] - coords[:-1], axis=1)\n    \n    # Histogram\n    fig_metrics.add_trace(\n        go.Histogram(\n            x=bond_lengths,\n            nbinsx=50,\n            marker=dict(color=\"steelblue\", line=dict(width=1, color=\"white\")),\n            showlegend=False\n        ),\n        row=row, col=col\n    )\n    \n    # Ideal line\n    fig_metrics.add_vline(\n        x=5.9, line_dash=\"dash\", line_color=\"red\", line_width=3,\n        annotation_text=\"Ideal (5.9Å)\",\n        row=row, col=col\n    )\n    \n    # Update axes\n    fig_metrics.update_xaxes(title_text=\"C1'-C1' Distance (Å)\", row=row, col=col, range=[3, 10])\n    fig_metrics.update_yaxes(title_text=\"Count\", row=row, col=col)\n    \n    # Stats\n    mean_bond = np.mean(bond_lengths)\n    std_bond = np.std(bond_lengths)\n    print(f\"  {target_id}: mean={mean_bond:.2f}Å, std={std_bond:.2f}Å (ideal=5.9Å)\")\n\nfig_metrics.update_layout(\n    title=dict(\n        text=\"<b>Bond Length Distribution (Quality Check)</b><br>\" +\n             \"<sub>Good structures cluster around 5.9Å (C1'-C1' distance in A-form RNA)</sub>\",\n        x=0.5,\n        xanchor=\"center\",\n        font=dict(size=18)\n    ),\n    height=800,\n    width=1200,\n    showlegend=False,\n    paper_bgcolor=\"white\"\n)\n\n# Save metrics\nhtml_metrics_path = f\"{OUTPUT_DIR}/quality_metrics.html\"\nfig_metrics.write_html(html_metrics_path)\nprint(f\"\\n✓ Saved quality metrics: {html_metrics_path}\")\n\nimg_metrics_path = f\"{OUTPUT_DIR}/quality_metrics.png\"\ntry:\n    fig_metrics.write_image(img_metrics_path, width=1200, height=800, scale=2)\n    print(f\"✓ Saved metrics image: {img_metrics_path}\")\n    has_metrics_image = True\nexcept Exception as e:\n    print(f\"⚠ Could not save PNG: HTML version available\")\n    has_metrics_image = False\n\n# Display\nprint(\"\\n\" + \"=\"*70)\nprint(\"  DISPLAYING QUALITY METRICS\")\nprint(\"=\"*70)\n\n# Read and display the HTML content directly\nwith open(html_metrics_path, 'r') as f:\n    html_metrics_content = f.read()\n\ndisplay(HTML(\"<h2 style='text-align: center; color: #2c3e50;'>Bond Length Quality Metrics</h2>\"))\ndisplay(HTML(html_metrics_content))\nprint(\"\\n✓ Quality metrics displayed above\")\n\n# =======================\n# CREATE SUMMARY TABLE\n# =======================\n\nprint(\"\\n\" + \"=\"*70)\nprint(\"  STRUCTURE SUMMARY\")\nprint(\"=\"*70)\n\nsummary_data = []\nfor target_id, _ in targets:\n    coords, nts, _ = extract_chain(df, target_id, pred_idx=1)\n    bond_lengths = np.linalg.norm(coords[1:] - coords[:-1], axis=1)\n    \n    # Count nucleotides\n    nt_counts = {nt: np.sum(nts == nt) for nt in \"AUGC\"}\n    \n    summary_data.append({\n        \"Target\": target_id,\n        \"Length\": len(coords),\n        \"A\": nt_counts.get(\"A\", 0),\n        \"U\": nt_counts.get(\"U\", 0),\n        \"G\": nt_counts.get(\"G\", 0),\n        \"C\": nt_counts.get(\"C\", 0),\n        \"Mean Bond (Å)\": f\"{np.mean(bond_lengths):.2f}\",\n        \"Std Bond (Å)\": f\"{np.std(bond_lengths):.2f}\",\n        \"Quality\": \"✓ Good\" if abs(np.mean(bond_lengths) - 5.9) < 0.5 else \"⚠ Check\"\n    })\n\nsummary_df = pd.DataFrame(summary_data)\n\nprint(\"\\n\" + summary_df.to_string(index=False))\n\n# Save summary\nsummary_path = f\"{OUTPUT_DIR}/summary_table.csv\"\nsummary_df.to_csv(summary_path, index=False)\nprint(f\"\\n✓ Saved summary: {summary_path}\")\n\n# Display as HTML table\ndisplay(HTML(\"<h2>Summary Table</h2>\"))\ndisplay(HTML(summary_df.to_html(index=False, border=1, justify=\"center\")))\n\n# =======================\n# FINAL MESSAGE\n# =======================\n\nprint(\"\\n\" + \"=\"*70)\nprint(\"  ✅ VISUALIZATION COMPLETE\")\nprint(\"=\"*70)\nprint(f\"\\nFiles saved to: {OUTPUT_DIR}/\")\nprint(\"  - structures_3d.html (interactive 3D)\")\nprint(\"  - structures_3d.png (static image)\")\nprint(\"  - quality_metrics.html (interactive metrics)\")\nprint(\"  - quality_metrics.png (static image)\")\nprint(\"  - summary_table.csv (structure info)\")\nprint(\"\\nTo view in Kaggle:\")\nprint(\"  1. Outputs are displayed above\")\nprint(\"  2. HTML files can be downloaded from Files tab\")\nprint(\"  3. PNG images will appear in notebook output after saving\")\nprint(\"=\"*70)\n\n# Display static images as fallback\nprint(\"\\n\" + \"=\"*70)\nprint(\"  STATIC IMAGE PREVIEWS\")\nprint(\"=\"*70)\n\nif has_image:\n    try:\n        print(\"\\n3D Structures:\")\n        display(Image(filename=img_path))\n    except:\n        print(\"(View the HTML file above for interactive 3D)\")\n\nif has_metrics_image:\n    try:\n        print(\"\\nQuality Metrics:\")\n        display(Image(filename=img_metrics_path))\n    except:\n        print(\"(View the HTML file above for interactive metrics)\")\n\nif not has_image and not has_metrics_image:\n    print(\"\\nℹ PNG export not available, but HTML files are saved!\")\n    print(\"  Download and open the HTML files to view interactive plots\")\n\nprint(\"\\n✓ All visualizations complete!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-18T07:45:37.807404Z","iopub.execute_input":"2026-01-18T07:45:37.808148Z","iopub.status.idle":"2026-01-18T07:45:49.670476Z","shell.execute_reply.started":"2026-01-18T07:45:37.808118Z","shell.execute_reply":"2026-01-18T07:45:49.66967Z"}},"outputs":[],"execution_count":null}]}