{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":87793,"databundleVersionId":11228175,"sourceType":"competition"},{"sourceId":6233380,"sourceType":"datasetVersion","datasetId":3580819},{"sourceId":7395079,"sourceType":"datasetVersion","datasetId":4299455},{"sourceId":7639698,"sourceType":"datasetVersion","datasetId":4299272},{"sourceId":10878276,"sourceType":"datasetVersion","datasetId":6758842},{"sourceId":10878463,"sourceType":"datasetVersion","datasetId":6759157},{"sourceId":10880297,"sourceType":"datasetVersion","datasetId":6760419},{"sourceId":10880353,"sourceType":"datasetVersion","datasetId":6760463},{"sourceId":10880374,"sourceType":"datasetVersion","datasetId":6760482},{"sourceId":10880419,"sourceType":"datasetVersion","datasetId":6760509},{"sourceId":10923666,"sourceType":"datasetVersion","datasetId":6791491},{"sourceId":10923782,"sourceType":"datasetVersion","datasetId":6791561},{"sourceId":10923886,"sourceType":"datasetVersion","datasetId":6791607},{"sourceId":10923997,"sourceType":"datasetVersion","datasetId":6791661},{"sourceId":224703571,"sourceType":"kernelVersion"}],"dockerImageVersionId":30919,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"Ribonanza + RhoFold + Rfam","metadata":{}},{"cell_type":"code","source":"# Step 0: Install Dependencies Alongside the Script\n!pip install /kaggle/input/openmm/OpenMM-8.2.0-cp310-cp310-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl\n!pip install /kaggle/input/simtk-0-1/simtk-0.1.0-py2.py3-none-any.whl\n!pip install /kaggle/input/pytest-runner/pytest_runner-6.0.1-py3-none-any.whl\n!pip install /kaggle/input/biopython/biopython-1.85-cp310-cp310-manylinux_2_17_x86_64.manylinux2014_x86_64.whl\n!pip install /kaggle/input/ml-collections/ml_collections-1.0.0-py3-none-any.whl\n\nimport os\nimport sys\nimport pandas as pd\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport yaml  # Added yaml import\nfrom pathlib import Path\nfrom Bio.PDB import PDBParser\nimport subprocess\nimport warnings\n\n# Suppress torch.cross warning\nwarnings.filterwarnings(\"ignore\", category=UserWarning, module=\"rhofold.utils.rigid_utils\")\n\n# Step 1: Setup Environment\ndownload_dir = \"/kaggle/input/stanford-rna-3d-folding/\"\noutput_dir = \"/kaggle/working/\"\nribonanza_dir = \"/kaggle/input/ribonanzanet2d-final/\"\nconfig_dir = \"/kaggle/input/ribonanzanet2d-final/configs/\"\nrhofold_dir = \"/kaggle/input/rhofold-repo/\"\nrfam_dir = \"/kaggle/input/rfam-allignment/\"\nos.makedirs(output_dir, exist_ok=True)\nos.makedirs(os.path.join(output_dir, \"test_fasta\"), exist_ok=True)\nos.makedirs(os.path.join(output_dir, \"out\"), exist_ok=True)\n\n# Copy and modify inference.py for weights_only=True\nos.system(f\"cp -r {rhofold_dir}/* {output_dir}\")\ninference_path = os.path.join(output_dir, \"inference.py\")\nwith open(inference_path, \"r\") as f:\n    lines = f.readlines()\nfor i, line in enumerate(lines):\n    if \"model.load_state_dict(torch.load(config.ckpt\" in line:\n        lines[i] = line.replace(\"torch.load(config.ckpt, map_location=torch.device('cpu'))['model']\", \n                                \"torch.load(config.ckpt, map_location=torch.device('cpu'), weights_only=True)['model']\")\nwith open(inference_path, \"w\") as f:\n    f.writelines(lines)\n\nprint(\"Files in competition data:\", os.listdir(download_dir))\nprint(\"Files in rhofold-repo dir:\", os.listdir(rhofold_dir))\nprint(\"Files in rfam-allignment dir:\", os.listdir(rfam_dir))\n\n# Step 2: Load Test Sequences (full set)\nkaggle_seq_df = pd.read_csv(os.path.join(download_dir, \"test_sequences.csv\"))\nsequences = kaggle_seq_df[\"sequence\"].tolist()\ntarget_ids = kaggle_seq_df[\"target_id\"].tolist()\n\n# Write FASTA files\nfor target_id, seq in zip(target_ids, sequences):\n    with open(f\"{output_dir}/test_fasta/{target_id}.fasta\", \"w\") as f:\n        f.write(f\">{target_id}\\n{seq}\\n\")\n\n# Step 3: Load Rfam Alignment\nrfam_path = os.path.join(rfam_dir, \"rfam_alignment.sto\")\nrfam_dict = {}\nif os.path.exists(rfam_path):\n    def parse_stockholm(file_path):\n        sequences = {}\n        with open(file_path, \"r\") as f:\n            lines = f.readlines()\n            current_id = None\n            for line in lines:\n                if line.startswith(\"#=GS\"):\n                    parts = line.split()\n                    if len(parts) > 1:\n                        current_id = parts[1]\n                elif not line.startswith(\"#\") and not line.startswith(\"//\") and line.strip():\n                    parts = line.split()\n                    if len(parts) > 1 and current_id:\n                        seq = parts[-1].replace(\".\", \"\")\n                        sequences[current_id] = seq\n        return sequences\n    rfam_seqs = parse_stockholm(rfam_path)\n    rfam_dict = {k: v for k, v in rfam_seqs.items()}\n    print(\"Rfam data loaded:\", list(rfam_dict.keys())[:5])\n\n# Step 4: RhoFold Predictions (≤ 350 residues)\nrho_preds = []\ncheckpoint_path = os.path.join(rhofold_dir, \"pretrained\", \"RhoFold_pretrained.pt\")\nif not os.path.exists(checkpoint_path):\n    print(f\"Error: Checkpoint {checkpoint_path} not found.\")\nelse:\n    for i, (seq, target_id) in enumerate(zip(sequences, target_ids)):\n        if len(seq) <= 350:\n            print(f\"Running RhoFold inference.py for {target_id} (length: {len(seq)})\")\n            fasta_file = f\"{output_dir}/test_fasta/{target_id}.fasta\"\n            out_dir = f\"{output_dir}/out/{target_id}\"\n            os.makedirs(out_dir, exist_ok=True)\n            device = \"cuda:0\" if len(seq) < 300 else \"cpu\"\n            cmd = f\"python {inference_path} --relax_steps 5 --input_fas {fasta_file} --single_seq_pred True --output_dir {out_dir} --device {device} --ckpt {checkpoint_path}\"\n            result = subprocess.run(cmd, shell=True, capture_output=True, text=True)\n            print(f\"RhoFold stdout: {result.stdout}\")\n            print(f\"RhoFold stderr: {result.stderr}\")\n            coords_list = []\n            for step in range(1, 6):  # 5 relaxation steps\n                pdb_file = f\"{out_dir}/relaxed_{step}_model.pdb\"\n                if Path(pdb_file).exists():\n                    parser = PDBParser()\n                    structure = parser.get_structure('RNA', pdb_file)\n                    c1_coords = []\n                    for model in structure:\n                        for chain in model:\n                            for residue in chain:\n                                if residue.get_resname() in ['A', 'U', 'G', 'C']:\n                                    try:\n                                        c1_coords.append(residue[\"C1'\"].get_coord())\n                                    except KeyError:\n                                        c1_coords.append([0.0, 0.0, 0.0])\n                    coords_list.append(np.array(c1_coords))\n                else:\n                    coords_list.append(np.zeros((len(seq), 3)))\n            rho_preds.append({\"sequence\": seq, \"coords\": coords_list[:5]})  # Ensure 5\n        else:\n            rho_preds.append({\"sequence\": seq, \"coords\": None})\nnp.save(os.path.join(output_dir, \"rhofold_predictions.npy\"), rho_preds)\n\n# Step 5: RibonanzaNet Predictions (> 350 residues)\nsys.path.append(ribonanza_dir)\nfrom Network import RibonanzaNet\n\nclass Config:\n    def __init__(self, **entries):\n        self.__dict__.update(entries)\n\ndef load_config_from_yaml(file_path):\n    with open(file_path, 'r') as file:\n        config = yaml.safe_load(file)\n    return Config(**config)\n\nclass finetuned_RibonanzaNet(RibonanzaNet):\n    def __init__(self, config):\n        config.dropout = 0.2\n        super(finetuned_RibonanzaNet, self).__init__(config)\n        self.dropout = nn.Dropout(0.0)\n        self.xyz_predictor = nn.Linear(256, 3)\n\n    def forward(self, src):\n        sequence_features, pairwise_features = self.get_embeddings(src, torch.ones_like(src).long().to(src.device))\n        xyz = self.xyz_predictor(sequence_features)\n        return xyz\n\nconfig_path = os.path.join(config_dir, \"pairwise.yaml\")\nribo_model = finetuned_RibonanzaNet(load_config_from_yaml(config_path))\nribo_model_path = \"/kaggle/input/ribonanzanet-3d-finetune/RibonanzaNet-3D.pt\"\nribo_model.load_state_dict(torch.load(ribo_model_path, map_location=\"cuda:0\", weights_only=True))\nribo_model.to(\"cuda:0\")\nribo_model.eval()\n\ndef encode_sequence(seq):\n    mapping = {\"A\": 0, \"U\": 1, \"G\": 2, \"C\": 3}\n    tokens = [mapping.get(n, 0) for n in seq.upper()]\n    return torch.tensor(tokens, dtype=torch.long).unsqueeze(0)  # [1, L]\n\nribo_preds = []\nfor i, seq in enumerate(sequences):\n    if len(seq) > 350:\n        print(f\"Processing RibonanzaNet for sequence {i+1}/{len(sequences)} (length: {len(seq)})\")\n        inputs = encode_sequence(seq).to(\"cuda:0\")\n        coords_list = []\n        for _ in range(5):  # 5 distinct predictions\n            with torch.no_grad():\n                xyz = ribo_model(inputs).squeeze(0)\n                coords_list.append(xyz.cpu().numpy())\n        ribo_preds.append({\"sequence\": seq, \"coords\": coords_list})\n    else:\n        ribo_preds.append({\"sequence\": seq, \"coords\": None})\nnp.save(os.path.join(output_dir, \"ribonanzanet_predictions.npy\"), ribo_preds)\n\n# Step 6: Neural Network Refiner\nclass Refiner(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.fc = nn.Sequential(\n            nn.Linear(3, 64), nn.ReLU(),\n            nn.Linear(64, 64), nn.ReLU(),\n            nn.Linear(64, 3)\n        )\n\n    def forward(self, x):\n        return self.fc(x)\n\nrefiner = Refiner().to(\"cuda:0\")\noptimizer = torch.optim.Adam(refiner.parameters(), lr=1e-3)\nif rfam_dict:\n    rfam_seq = next(iter(rfam_dict.values()), sequences[0][:350])  # Fallback to test seq\n    dummy_coords = torch.randn(len(rfam_seq), 3).to(\"cuda:0\")\n    for _ in range(50):\n        refiner.train()\n        refined = refiner(dummy_coords)\n        loss = ((refined - dummy_coords) ** 2).mean()\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\nrefiner.eval()\n\n# Step 7: Combine and Format Submission\nrho_preds = np.load(os.path.join(output_dir, \"rhofold_predictions.npy\"), allow_pickle=True)\nribo_preds = np.load(os.path.join(output_dir, \"ribonanzanet_predictions.npy\"), allow_pickle=True)\n\nsubmission_data = []\nfor i, (seq, target_id) in enumerate(zip(sequences, target_ids)):\n    rho_coords = rho_preds[i][\"coords\"] if rho_preds[i][\"coords\"] is not None else None\n    ribo_coords = ribo_preds[i][\"coords\"] if ribo_preds[i][\"coords\"] is not None else None\n    coords_list = rho_coords if rho_coords is not None else (ribo_coords if ribo_coords is not None else [np.zeros((len(seq), 3))] * 5)\n    \n    # Ensure 5 predictions\n    if len(coords_list) < 5:\n        coords_list += [coords_list[-1]] * (5 - len(coords_list))\n    \n    # Refine coordinates\n    refined_coords_list = []\n    for coords in coords_list[:5]:\n        coords_tensor = torch.tensor(coords, dtype=torch.float32).to(\"cuda:0\")\n        with torch.no_grad():\n            refined_coords = refiner(coords_tensor).cpu().numpy()\n        refined_coords_list.append(refined_coords)\n    \n    for j, base in enumerate(seq, 1):\n        row = [f\"{target_id}_{j}\", base, j]\n        for k in range(5):\n            coord = refined_coords_list[k][j-1] if j-1 < len(refined_coords_list[k]) else [0.0, 0.0, 0.0]\n            row.extend([float(coord[0]), float(coord[1]), float(coord[2])])\n        submission_data.append(row)\n\ncolumns = [\"ID\", \"resname\", \"resid\"] + [f\"{dim}_{i+1}\" for i in range(5) for dim in [\"x\", \"y\", \"z\"]]\nsubmission_df = pd.DataFrame(submission_data, columns=columns)\nsubmission_df.to_csv(os.path.join(output_dir, \"submission.csv\"), index=False)\nprint(\"Submission file created:\", submission_df.tail())\nprint(\"Submission file saved to /kaggle/working/submission.csv. Please submit manually.\")","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}