{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":87793,"databundleVersionId":11553390,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":10855324,"sourceType":"datasetVersion","datasetId":6742586},{"sourceId":10880419,"sourceType":"datasetVersion","datasetId":6760509},{"sourceId":11118830,"sourceType":"datasetVersion","datasetId":6933267},{"sourceId":11377895,"sourceType":"datasetVersion","datasetId":7123615},{"sourceId":11378050,"sourceType":"datasetVersion","datasetId":7123736},{"sourceId":11378078,"sourceType":"datasetVersion","datasetId":7123758},{"sourceId":11378106,"sourceType":"datasetVersion","datasetId":7123782},{"sourceId":11378145,"sourceType":"datasetVersion","datasetId":7123815},{"sourceId":11378156,"sourceType":"datasetVersion","datasetId":7123825},{"sourceId":11378239,"sourceType":"datasetVersion","datasetId":7123884},{"sourceId":11378256,"sourceType":"datasetVersion","datasetId":7123898},{"sourceId":11378268,"sourceType":"datasetVersion","datasetId":7123908},{"sourceId":11378426,"sourceType":"datasetVersion","datasetId":7124032},{"sourceId":233426810,"sourceType":"kernelVersion"}],"dockerImageVersionId":31011,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# # This Python 3 environment comes with many helpful analytics libraries installed\n# # It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# # For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# # Input data files are available in the read-only \"../input/\" directory\n# # For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\n# for dirname, _, filenames in os.walk('/kaggle/input'):\n#     for filename in filenames:\n#         print(os.path.join(dirname, filename))\n\n# # You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# # You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-12T14:11:39.619705Z","iopub.execute_input":"2025-04-12T14:11:39.620485Z","iopub.status.idle":"2025-04-12T14:11:39.873850Z","shell.execute_reply.started":"2025-04-12T14:11:39.620441Z","shell.execute_reply":"2025-04-12T14:11:39.873132Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport sys\nimport pandas as pd\nimport numpy as np\nimport torch\nimport re\nimport time\nimport warnings\nimport yaml\nimport gc\nfrom tqdm.notebook import tqdm as tqdm_nb  # Use notebook tqdm","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T14:11:39.875068Z","iopub.execute_input":"2025-04-12T14:11:39.875411Z","iopub.status.idle":"2025-04-12T14:11:43.439381Z","shell.execute_reply.started":"2025-04-12T14:11:39.875393Z","shell.execute_reply":"2025-04-12T14:11:43.438834Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Suppress warnings\nwarnings.filterwarnings(\"ignore\")\n\n# Timing and Memory Utilities\nstart_time_global = time.time()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T14:11:43.439918Z","iopub.execute_input":"2025-04-12T14:11:43.440167Z","iopub.status.idle":"2025-04-12T14:11:43.443918Z","shell.execute_reply.started":"2025-04-12T14:11:43.440153Z","shell.execute_reply":"2025-04-12T14:11:43.443292Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def time_to_str(t, mode='min'):\n    if mode == 'min':\n        t_int = int(t) / 60\n        hr = int(t_int // 60)\n        min_val = int(t_int % 60)\n        return '%2d hr %02d min' % (hr, min_val)\n    elif mode == 'sec':\n        t_int = int(t)\n        min_val = int(t_int // 60)\n        sec = int(t_int % 60)\n        return '%2d min %02d sec' % (min_val, sec)\n    else:\n        raise NotImplementedError\n\ndef gpu_memory_use():\n    if torch.cuda.is_available():\n        allocated = torch.cuda.memory_allocated(0) / 1024**3\n        reserved = torch.cuda.memory_reserved(0) / 1024**3 # Cache\n        return int(round(allocated)), int(round(reserved))\n    else:\n        return 0, 0\n\n# --- Configuration ---\nENSEMBLE_MODELS = ['protenix', 'drfold', 'ribonanzanet']\nNUM_CONFORMATIONS = 5\nKAGGLE_INPUT_DIR = '/kaggle/input'\nKAGGLE_WORKING_DIR = '/kaggle/working'\nOUTPUT_FILE = 'submission.csv'\n\n# Paths to model-specific inputs (adjust if necessary based on your Kaggle dataset names)\nPROTENIX_CHECKPOINTS = f'{KAGGLE_INPUT_DIR}/protenix-checkpoints'\nDRFOLD_DUMMY_DIR = f'{KAGGLE_INPUT_DIR}/hengck23-drfold2-dummy-00'\nUSALIGN_INPUT = f'{KAGGLE_INPUT_DIR}/usalign'\nRIBONANZANET_2D_DIR = f'{KAGGLE_INPUT_DIR}/ribonanzanet2d-final'\nRIBONANZANET_3D_WEIGHTS = f'{KAGGLE_INPUT_DIR}/ribonanzanet-3d-finetune'\nBIOPYTHON_WHL = f'{KAGGLE_INPUT_DIR}/biopython/biopython-1.85-cp310-cp310-manylinux_2_17_x86_64.manylinux2014_x86_64.whl'\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T14:11:43.445528Z","iopub.execute_input":"2025-04-12T14:11:43.446126Z","iopub.status.idle":"2025-04-12T14:11:43.458350Z","shell.execute_reply.started":"2025-04-12T14:11:43.446108Z","shell.execute_reply":"2025-04-12T14:11:43.457728Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!python --version","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T14:11:43.458873Z","iopub.execute_input":"2025-04-12T14:11:43.459027Z","iopub.status.idle":"2025-04-12T14:11:43.600898Z","shell.execute_reply.started":"2025-04-12T14:11:43.459013Z","shell.execute_reply":"2025-04-12T14:11:43.600268Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# !pip install /kaggle/input/biotite311/biotite-1.0.1-cp311-cp311-manylinux_2_17_x86_64.manylinux2014_x86_64.whl","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T14:11:43.601900Z","iopub.execute_input":"2025-04-12T14:11:43.602166Z","iopub.status.idle":"2025-04-12T14:11:43.605982Z","shell.execute_reply.started":"2025-04-12T14:11:43.602134Z","shell.execute_reply":"2025-04-12T14:11:43.605292Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# !pip install /kaggle/input/d/lordpatil/protenix/protenix-0.4.6-py3-none-any.whl","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T14:11:43.606809Z","iopub.execute_input":"2025-04-12T14:11:43.607630Z","iopub.status.idle":"2025-04-12T14:11:43.619153Z","shell.execute_reply.started":"2025-04-12T14:11:43.607612Z","shell.execute_reply":"2025-04-12T14:11:43.618570Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import warnings\nwarnings.filterwarnings(\"ignore\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T14:11:43.619845Z","iopub.execute_input":"2025-04-12T14:11:43.620045Z","iopub.status.idle":"2025-04-12T14:11:43.631683Z","shell.execute_reply.started":"2025-04-12T14:11:43.620022Z","shell.execute_reply":"2025-04-12T14:11:43.630957Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- Installations and Setup ---\nprint(\"--- Installing Dependencies ---\")\n# Protenix dependencies (assuming it needs to be installed)\n# Note: Protenix installation might be complex and require specific versions.\n# Ensure the protenix library and its dependencies are available in the Kaggle environment\n# or install them here if needed. The original notebook commented these out for submission.\n!pip install /kaggle/input/biopython311/biopython-1.85-cp311-cp311-manylinux_2_17_x86_64.manylinux2014_x86_64.whl\n!pip install --no-deps /kaggle/input/d/lordpatil/protenix/protenix-0.4.6-py3-none-any.whl \n\n!pip install /kaggle/input/ml-collections/ml_collections-1.0.0-py3-none-any.whl\n!pip install --no-deps /kaggle/input/biotite311/biotite-1.0.1-cp311-cp311-manylinux_2_17_x86_64.manylinux2014_x86_64.whl\n!pip install /kaggle/input/rdkit311x86/rdkit-2024.9.4-cp311-cp311-manylinux_2_28_x86_64.whl\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T14:11:43.632485Z","iopub.execute_input":"2025-04-12T14:11:43.632707Z","iopub.status.idle":"2025-04-12T14:12:00.196255Z","shell.execute_reply.started":"2025-04-12T14:11:43.632682Z","shell.execute_reply":"2025-04-12T14:12:00.195539Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# DRFold dependencies\ntry:\n    import Bio\n    print(\"Biopython already installed.\")\nexcept ImportError:\n    print(\"Installing Biopython...\")\n    !pip install \"{BIOPYTHON_WHL}\"\n\nprint(\"Dependency check/installation complete.\")\n\n# --- Setup Environment for Models ---\nprint(\"\\n--- Setting up Model Environments ---\")\n\n# Protenix setup\nif 'protenix' in ENSEMBLE_MODELS:\n    print(\"Setting up Protenix environment...\")\n    os.environ['PROTENIX_DATA_ROOT_DIR'] = PROTENIX_CHECKPOINTS\n    os.makedirs('/af3-dev', exist_ok=True)\n    if not os.path.exists('/af3-dev/release_data'):\n        os.symlink(PROTENIX_CHECKPOINTS, '/af3-dev/release_data', target_is_directory=True)\n    print(\"Protenix files linked:\")\n    !ls /af3-dev/release_data/","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T14:12:00.198512Z","iopub.execute_input":"2025-04-12T14:12:00.198718Z","iopub.status.idle":"2025-04-12T14:12:00.340212Z","shell.execute_reply.started":"2025-04-12T14:12:00.198699Z","shell.execute_reply":"2025-04-12T14:12:00.339300Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# DRFold setup\nif 'drfold' in ENSEMBLE_MODELS:\n    print(\"\\nSetting up DRFold environment...\")\n    USALIGN_EXEC = f'{KAGGLE_WORKING_DIR}/USalign'\n    if not os.path.exists(USALIGN_EXEC):\n        print(\"Copying USalign...\")\n        os.system(f'cp {USALIGN_INPUT}/USalign {KAGGLE_WORKING_DIR}/')\n        os.system(f'chmod +x {USALIGN_EXEC}')\n        print(\"USalign copied and made executable.\")\n    else:\n        print(\"USalign already exists.\")\n    # Add DRFold code path\n    sys.path.append(f'{DRFOLD_DUMMY_DIR}/drfold2/cfg_97')\n\n# RibonanzaNet setup\nif 'ribonanzanet' in ENSEMBLE_MODELS:\n     print(\"\\nSetting up RibonanzaNet environment...\")\n     sys.path.append(RIBONANZANET_2D_DIR)\n\nprint(\"Model environment setup complete.\")\n\n# --- Helper Functions ---\n\n# Generic dotdict\nclass dotdict(dict):\n    __setattr__ = dict.__setitem__\n    __delattr__ = dict.__delitem__\n    def __getattr__(self, name):\n        try:\n            return self[name]\n        except KeyError:\n            raise AttributeError(name)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T14:12:00.341366Z","iopub.execute_input":"2025-04-12T14:12:00.341705Z","iopub.status.idle":"2025-04-12T14:12:00.379804Z","shell.execute_reply.started":"2025-04-12T14:12:00.341673Z","shell.execute_reply":"2025-04-12T14:12:00.379008Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# DataFrame formatting function (common requirement)\ndef format_predictions_to_df(target_id, sequence, coords_list):\n    \"\"\"\n    Formats predictions into the submission DataFrame structure.\n\n    Args:\n        target_id (str): The ID of the RNA target.\n        sequence (str): The RNA sequence.\n        coords_list (list): A list of numpy arrays, where each array is (L, 3)\n                             representing one conformation's C1' coordinates.\n                             Should contain NUM_CONFORMATIONS arrays.\n\n    Returns:\n        pd.DataFrame: DataFrame rows for this sequence.\n    \"\"\"\n    L = len(sequence)\n    if not coords_list or len(coords_list) != NUM_CONFORMATIONS:\n        print(f\"Warning: Incorrect number of conformations for {target_id}. Expected {NUM_CONFORMATIONS}, got {len(coords_list)}. Filling with zeros.\")\n        # Create zero coordinates if prediction failed or returned wrong number\n        coords_list = [np.zeros((L, 3), dtype=np.float32) for _ in range(NUM_CONFORMATIONS)]\n        \n    # Ensure all coordinate arrays have the correct length\n    processed_coords = []\n    for i, coords in enumerate(coords_list):\n        if coords.shape[0] != L:\n             print(f\"Warning: Coordinate length mismatch for {target_id}, conformation {i+1}. Expected {L}, got {coords.shape[0]}. Padding/truncating.\")\n             # Simple padding/truncating - adjust if needed\n             new_coords = np.zeros((L, 3), dtype=np.float32)\n             len_to_copy = min(L, coords.shape[0])\n             new_coords[:len_to_copy, :] = coords[:len_to_copy, :]\n             processed_coords.append(new_coords)\n        else:\n             processed_coords.append(coords)\n             \n    df_rows = []\n    for i in range(L):\n        row_data = {\n            'ID': f'{target_id}_{i + 1}',\n            'resname': sequence[i],\n            'resid': i + 1\n        }\n        for k in range(NUM_CONFORMATIONS):\n            coords = processed_coords[k]\n            row_data[f'x_{k+1}'] = coords[i, 0]\n            row_data[f'y_{k+1}'] = coords[i, 1]\n            row_data[f'z_{k+1}'] = coords[i, 2]\n        df_rows.append(row_data)\n    return pd.DataFrame(df_rows)\n\n# Function to create zero predictions for a sequence\ndef create_zero_predictions_df(target_id, sequence):\n    print(f\"Creating zero predictions for {target_id} due to failure.\")\n    zero_coords = [np.zeros((len(sequence), 3), dtype=np.float32) for _ in range(NUM_CONFORMATIONS)]\n    return format_predictions_to_df(target_id, sequence, zero_coords)\n\n# --- Load Test Data ---\nprint(\"\\n--- Loading Test Data ---\")\ntest_df = pd.read_csv(f\"{KAGGLE_INPUT_DIR}/stanford-rna-3d-folding/test_sequences.csv\")\nprint(f\"Loaded {len(test_df)} test sequences.\")\nprint(test_df.head())\n\n# Dictionary to store predictions from each model\nmodel_predictions = {}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T14:12:00.380954Z","iopub.execute_input":"2025-04-12T14:12:00.381281Z","iopub.status.idle":"2025-04-12T14:12:00.761209Z","shell.execute_reply.started":"2025-04-12T14:12:00.381252Z","shell.execute_reply":"2025-04-12T14:12:00.760509Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# === Protenix Prediction ===\ndef predict_protenix(test_sequences_df):\n    print(\"\\n--- Running Protenix Predictions ---\")\n    try:\n        from runner.inference import update_inference_configs, InferenceRunner\n        from protenix.data.infer_data_pipeline import InferenceDataset\n        from configs.configs_base import configs as configs_base\n        from configs.configs_data import data_configs\n        from configs.configs_inference import inference_configs\n        from protenix.config.config import parse_configs\n    except ImportError as e:\n        print(f\"Protenix import failed: {e}. Skipping Protenix.\")\n        return None\n\n    # Protenix DictDataset specific to its data loading\n    class ProtenixDictDataset(InferenceDataset):\n        def __init__(self, seq_list: list, id_list: list, dump_dir: str = 'output', use_msa: bool = False) -> None:\n            self.dump_dir = dump_dir\n            self.use_msa = use_msa\n            self.inputs = [{\n                \"sequences\": [{\"rnaSequence\": {\"sequence\": seq, \"count\": 1}}],\n                \"name\": i\n            } for i, seq in zip(id_list, seq_list)]\n\n    all_preds_df = pd.DataFrame()\n    try:\n        # Configure Protenix\n        np.random.seed(0)\n        torch.manual_seed(0)\n        if torch.cuda.is_available():\n            torch.cuda.manual_seed_all(0)\n\n        configs_base[\"use_deepspeed_evo_attention\"] = (os.environ.get(\"USE_DEEPSPEED_EVO_ATTENTION\", \"false\").lower() == \"true\")\n        configs_base[\"model\"][\"N_cycle\"] = 10 # As in notebook\n        configs_base[\"sample_diffusion\"][\"N_sample\"] = NUM_CONFORMATIONS # Ensure 5 samples\n        configs_base[\"sample_diffusion\"][\"N_step\"] = 200 # As in notebook\n        inference_configs['load_checkpoint_path'] = f'{PROTENIX_CHECKPOINTS}/model_v0.2.0.pt'\n        configs = {**configs_base, **{\"data\": data_configs}, **inference_configs}\n        configs = parse_configs(configs=configs, fill_required_with_null=True)\n\n        runner = InferenceRunner(configs)\n        print(\"Protenix runner initialized.\")\n\n        dataset = ProtenixDictDataset(\n            test_sequences_df.sequence.tolist(),\n            test_sequences_df.target_id.tolist()\n        )\n        print(f\"Protenix dataset created with {len(dataset)} samples.\")\n        \n        for i in tqdm_nb(range(len(dataset)), desc=\"Protenix\"):\n            target_id = test_sequences_df.target_id[i]\n            seq = test_sequences_df.sequence[i]\n            try:\n                data, _, data_error_message = dataset[i]\n                if data_error_message != '':\n                    print(f\"Protenix data error for {target_id}: {data_error_message}\")\n                    raise ValueError(data_error_message)\n\n                # Protenix uses 'N_token' for config update\n                N_token = data.get(\"N_token\")\n                if N_token is None:\n                     print(f\"Warning: 'N_token' not found in Protenix data for {target_id}. Cannot update configs dynamically. Using defaults.\")\n                     # Handle cases where N_token might be missing, maybe use sequence length?\n                     # Or skip updating configs, which might be okay for inference.\n                     new_configs = configs # Use original configs\n                else:\n                     new_configs = update_inference_configs(configs, N_token.item())\n                \n                runner.update_model_configs(new_configs)\n                \n                with torch.no_grad():\n                    prediction_dict = runner.predict(data)\n\n                # Extract C1' atoms (atom index 12 in protenix output)\n                # Output shape: [N_sample, N_residue, 3]\n                atom_to_tokatom_idx = data.get('input_feature_dict', {}).get('atom_to_tokatom_idx', None)\n                if atom_to_tokatom_idx is None:\n                    raise ValueError(f\"Could not find 'atom_to_tokatom_idx' in Protenix data for {target_id}\")\n                    \n                c1_prime_coords = prediction_dict['coordinate'][:, atom_to_tokatom_idx == 12]\n                \n                if c1_prime_coords.shape[0] != NUM_CONFORMATIONS or c1_prime_coords.shape[1] != len(seq):\n                     print(f\"Warning: Protenix output shape mismatch for {target_id}. Expected ({NUM_CONFORMATIONS}, {len(seq)}, 3), got {c1_prime_coords.shape}\")\n                     # Attempt to reshape or handle error, for now, raise to fall back to zeros\n                     raise ValueError(\"Shape mismatch in Protenix output\")\n\n                coords_list = [c1_prime_coords[k].cpu().numpy() for k in range(NUM_CONFORMATIONS)]\n                seq_df = format_predictions_to_df(target_id, seq, coords_list)\n\n            except Exception as e:\n                print(f\"ERROR predicting with Protenix for {target_id}: {e}\")\n                seq_df = create_zero_predictions_df(target_id, seq)\n\n            all_preds_df = pd.concat([all_preds_df, seq_df], ignore_index=True)\n            if torch.cuda.is_available():\n                torch.cuda.empty_cache()\n                \n        print(\"Protenix predictions finished.\")\n        return all_preds_df\n\n    except Exception as e:\n        print(f\"FATAL ERROR during Protenix prediction setup or loop: {e}\")\n        # Return a DataFrame of zeros for all test sequences if setup fails\n        all_preds_df = pd.DataFrame()\n        for i in range(len(test_sequences_df)):\n            target_id = test_sequences_df.target_id[i]\n            seq = test_sequences_df.sequence[i]\n            seq_df = create_zero_predictions_df(target_id, seq)\n            all_preds_df = pd.concat([all_preds_df, seq_df], ignore_index=True)\n        return all_preds_df\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T14:12:00.762069Z","iopub.execute_input":"2025-04-12T14:12:00.762421Z","iopub.status.idle":"2025-04-12T14:12:00.775423Z","shell.execute_reply.started":"2025-04-12T14:12:00.762398Z","shell.execute_reply":"2025-04-12T14:12:00.774850Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# === DRFold Prediction ===\ndef predict_drfold(test_sequences_df):\n    print(\"\\n--- Running DRFold Predictions ---\")\n    try:\n        from EvoMSA2XYZ.Model import MSA2XYZ\n        from RNALM2.Model import RNA2nd\n        from data import parse_seq, Get_base, BASE_COOR, write_frame_coor_to_pdb, parse_pdb_to_xyz\n    except ImportError as e:\n        print(f\"DRFold import failed: {e}. Skipping DRFold.\")\n        return None\n\n    MAX_LENGTH_DRFOLD = 480 # From notebook\n    DEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\n    \n    # DRFold specific helpers\n    def make_data_drfold(seq):\n        aa_type = parse_seq(seq)\n        base = Get_base(seq, BASE_COOR)\n        seq_idx = np.arange(len(seq)) + 1\n        msa = aa_type[None, :]\n        msa = torch.from_numpy(msa)\n        # The notebook duplicates the MSA - replicate that behavior\n        msa = torch.cat([msa, msa], 0)\n        msa = torch.nn.functional.one_hot(msa.long(), 6).float()\n        base_x = torch.from_numpy(base).float()\n        seq_idx = torch.from_numpy(seq_idx).long()\n        return msa, base_x, seq_idx\n\n    def coord_to_df_drfold(sequence, coord_list, target_id):\n        # Adapts solution_to_submit_df logic for a single sequence\n        L = len(sequence)\n        df_rows = []\n        for i in range(L):\n            row_data = {\n                'ID': f'{target_id}_{i + 1}',\n                'resname': sequence[i],\n                'resid': i + 1\n            }\n            for k in range(NUM_CONFORMATIONS):\n                 coords = coord_list[k]\n                 if coords.shape[0] != L: # Handle length mismatch from padding/truncating\n                     print(f\"Warning: DRFold internal coord mismatch for {target_id}, conf {k+1}. Len {coords.shape[0]} vs seq {L}.\")\n                     # Provide zeros if mismatch is severe, otherwise pad/truncate? Using zeros for safety.\n                     row_data[f'x_{k+1}'] = 0.0\n                     row_data[f'y_{k+1}'] = 0.0\n                     row_data[f'z_{k+1}'] = 0.0\n                 else:\n                    row_data[f'x_{k+1}'] = coords[i, 0]\n                    row_data[f'y_{k+1}'] = coords[i, 1]\n                    row_data[f'z_{k+1}'] = coords[i, 2]\n            df_rows.append(row_data)\n        return pd.DataFrame(df_rows)\n\n    all_preds_df = pd.DataFrame()\n    try:\n        # Load RNALM model (needed by MSA2XYZ)\n        rnalm = RNA2nd(dict(s_in_dim=5, z_in_dim=2, s_dim=512, z_dim=128, N_elayers=18))\n        rnalm_file = f'{DRFOLD_DUMMY_DIR}/RCLM/epoch_67000'\n        print(f\"Loading RNALM from {rnalm_file}\")\n        rnalm.load_state_dict(torch.load(rnalm_file, map_location='cpu', weights_only=True), strict=False)\n        rnalm = rnalm.to(DEVICE).eval()\n        print(\"RNALM loaded.\")\n\n        msa2xyz_models = []\n        model_indices = [0, 1, 2, 8, 9] # Checkpoints used in the notebook\n        for k in model_indices:\n            msa2xyz = MSA2XYZ(dict(seq_dim=6, msa_dim=7, N_ensemble=1, N_cycle=8, m_dim=64, s_dim=64, z_dim=64))\n            msa2xyz_file = f'{DRFOLD_DUMMY_DIR}/cfg_97/model_{k}'\n            print(f\"Loading MSA2XYZ from {msa2xyz_file}\")\n            msa2xyz.load_state_dict(torch.load(msa2xyz_file, map_location='cpu', weights_only=True), strict=True)\n            msa2xyz.msaxyzone.premsa.rnalm = rnalm # Inject RNALM\n            msa2xyz = msa2xyz.to(DEVICE).eval()\n            msa2xyz_models.append(msa2xyz)\n            print(f\"MSA2XYZ model {k} loaded.\")\n\n        for i, row in tqdm_nb(test_sequences_df.iterrows(), total=len(test_sequences_df), desc=\"DRFold\"):\n            target_id = row.target_id\n            sequence = row.sequence\n            L = len(sequence)\n            coords_list = []\n\n            try:\n                # Handle potential length truncation/padding like in the notebook\n                if L > MAX_LENGTH_DRFOLD:\n                    # The original notebook picked a random start index. For reproducibility/submission,\n                    # let's just take the first MAX_LENGTH_DRFOLD residues.\n                    # Or, should we process the whole sequence if possible? Let's try processing all first.\n                    # If memory error, we might need to truncate.\n                    # The original notebook logic:\n                    # i0 = np.random.choice(L - MAX_LENGTH_DRFOLD + 1)\n                    # i1 = i0 + MAX_LENGTH_DRFOLD\n                    # seq_processed = sequence[i0:i1]\n                    # L_processed = len(seq_processed)\n                    # print(f\"Warning: Sequence {target_id} too long ({L}), truncating to {MAX_LENGTH_DRFOLD}\")\n                    # Let's process the full sequence first. If it fails, we'll know.\n                    seq_processed = sequence\n                    L_processed = L\n                    i0, i1 = 0, L # Keep track for potential later padding (though not needed if processing full seq)\n                else:\n                    seq_processed = sequence\n                    L_processed = L\n                    i0, i1 = 0, L\n\n                msa, base_x, seq_idx = make_data_drfold(seq_processed)\n                msa, base_x, seq_idx = msa.to(DEVICE), base_x.to(DEVICE), seq_idx.to(DEVICE)\n\n                for model_idx, msa2xyz_model in enumerate(msa2xyz_models):\n                    with torch.no_grad():\n                        out = msa2xyz_model.pred(msa, seq_idx, None, base_x, np.array(list(seq_processed)))\n                    \n                    # C1' coordinate is the second atom in the frame (index 1)\n                    coord_c1_prime = out['coor'][:, 1, :].cpu().numpy() \n\n                    if L != L_processed: # Handle padding if truncation occurred (not currently happening)\n                         padded_coord = np.zeros((L, 3), dtype=np.float32)\n                         padded_coord[i0:i1, :] = coord_c1_prime\n                         coords_list.append(padded_coord)\n                    else:\n                        coords_list.append(coord_c1_prime)\n                        \n                # Convert list of coordinates to the required format DataFrame\n                seq_df = format_predictions_to_df(target_id, sequence, coords_list)\n            \n            except Exception as e:\n                 print(f\"ERROR predicting with DRFold for {target_id}: {e}\")\n                 seq_df = create_zero_predictions_df(target_id, sequence)\n            \n            all_preds_df = pd.concat([all_preds_df, seq_df], ignore_index=True)\n            if torch.cuda.is_available():\n                torch.cuda.empty_cache()\n                \n        print(\"DRFold predictions finished.\")\n        return all_preds_df\n\n    except Exception as e:\n        print(f\"FATAL ERROR during DRFold prediction setup or loop: {e}\")\n        # Return a DataFrame of zeros for all test sequences if setup fails\n        all_preds_df = pd.DataFrame()\n        for i in range(len(test_sequences_df)):\n            target_id = test_sequences_df.target_id[i]\n            seq = test_sequences_df.sequence[i]\n            seq_df = create_zero_predictions_df(target_id, seq)\n            all_preds_df = pd.concat([all_preds_df, seq_df], ignore_index=True)\n        return all_preds_df\n\n\n# === RibonanzaNet Prediction ===\ndef predict_ribonanzanet(test_sequences_df):\n    print(\"\\n--- Running RibonanzaNet Predictions ---\")\n    try:\n        # Need RNADataset, Config, finetuned_RibonanzaNet from the notebook setup\n        from Network import RibonanzaNet # Assuming Network.py is in the path\n        from torch.utils.data import Dataset\n        \n        # RibonanzaNet specific dataset\n        class RNADatasetRibonanza(Dataset):\n            def __init__(self, data):\n                self.data = data\n                self.tokens = {nt: i for i, nt in enumerate('ACGU')}\n            def __len__(self): return len(self.data)\n            def __getitem__(self, idx):\n                sequence = [self.tokens[nt] for nt in (self.data.loc[idx, 'sequence'])]\n                sequence = np.array(sequence)\n                sequence = torch.tensor(sequence, dtype=torch.long) # Ensure long type\n                return {'sequence': sequence, 'target_id': self.data.loc[idx, 'target_id']}\n        \n        # Load config and model\n        config_rz = load_config_from_yaml(f\"{RIBONANZANET_2D_DIR}/configs/pairwise.yaml\")\n        \n        # Re-define the finetuned model class locally\n        class finetuned_RibonanzaNet(RibonanzaNet):\n             def __init__(self, config, pretrained=False):\n                 config.dropout=0.0 # Set dropout to 0 for eval consistency, but notebook uses 0.2? Let's follow notebook for train() runs.\n                 # Notebook uses 0.2 for train runs, 0.0 for eval runs\n                 super(finetuned_RibonanzaNet, self).__init__(config)\n                 if pretrained: # Not used here, loading finetuned weights below\n                     self.load_state_dict(torch.load(f\"{RIBONANZANET_2D_DIR}/RibonanzaNet.pt\", map_location='cpu'))\n                 self.dropout = nn.Dropout(0.0) # Default dropout for eval\n                 self.xyz_predictor = nn.Linear(256, 3) # Assuming input dim is 256\n\n             def forward(self, src):\n                 # Use the base class method to get embeddings\n                 sequence_features, pairwise_features = self.get_embeddings(src, torch.ones_like(src).long().to(src.device))\n                 # Apply dropout ONLY if in training mode\n                 sequence_features = self.dropout(sequence_features) \n                 xyz = self.xyz_predictor(sequence_features)\n                 return xyz\n\n        model_rz = finetuned_RibonanzaNet(config_rz, pretrained=False)\n        model_rz.load_state_dict(torch.load(f\"{RIBONANZANET_3D_WEIGHTS}/RibonanzaNet-3D.pt\", map_location='cpu'))\n        model_rz = model_rz.to('cuda' if torch.cuda.is_available() else 'cpu')\n        print(\"RibonanzaNet model loaded.\")\n\n        dataset_rz = RNADatasetRibonanza(test_sequences_df)\n        print(f\"RibonanzaNet dataset created with {len(dataset_rz)} samples.\")\n        \n        all_preds_df = pd.DataFrame()\n        \n        for i in tqdm_nb(range(len(dataset_rz)), desc=\"RibonanzaNet\"):\n            item = dataset_rz[i]\n            target_id = item['target_id']\n            seq_str = test_sequences_df.sequence[i] # Get string sequence for formatting\n            src = item['sequence'].unsqueeze(0).to(model_rz.device) # Add batch dim and move to device\n            \n            coords_list_np = []\n            try:\n                # Run 4 times in train mode (with dropout 0.2)\n                model_rz.dropout = torch.nn.Dropout(0.2) # Set dropout for train mode runs\n                model_rz.train()\n                for _ in range(NUM_CONFORMATIONS - 1):\n                     with torch.no_grad(): # Still no gradient needed for inference\n                          xyz = model_rz(src).squeeze(0) # Remove batch dim\n                     coords_list_np.append(xyz.cpu().numpy())\n                \n                # Run 1 time in eval mode (with dropout 0.0)\n                model_rz.dropout = torch.nn.Dropout(0.0) # Set dropout to 0 for eval\n                model_rz.eval()\n                with torch.no_grad():\n                    xyz = model_rz(src).squeeze(0) # Remove batch dim\n                coords_list_np.append(xyz.cpu().numpy())\n                \n                seq_df = format_predictions_to_df(target_id, seq_str, coords_list_np)\n\n            except Exception as e:\n                 print(f\"ERROR predicting with RibonanzaNet for {target_id}: {e}\")\n                 seq_df = create_zero_predictions_df(target_id, seq_str)\n\n            all_preds_df = pd.concat([all_preds_df, seq_df], ignore_index=True)\n            if torch.cuda.is_available():\n                torch.cuda.empty_cache()\n                \n        print(\"RibonanzaNet predictions finished.\")\n        return all_preds_df\n\n    except Exception as e:\n        print(f\"FATAL ERROR during RibonanzaNet prediction setup or loop: {e}\")\n        # Return a DataFrame of zeros for all test sequences if setup fails\n        all_preds_df = pd.DataFrame()\n        for i in range(len(test_sequences_df)):\n            target_id = test_sequences_df.target_id[i]\n            seq = test_sequences_df.sequence[i]\n            seq_df = create_zero_predictions_df(target_id, seq)\n            all_preds_df = pd.concat([all_preds_df, seq_df], ignore_index=True)\n        return all_preds_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T14:12:00.776194Z","iopub.execute_input":"2025-04-12T14:12:00.776443Z","iopub.status.idle":"2025-04-12T14:12:00.803628Z","shell.execute_reply.started":"2025-04-12T14:12:00.776428Z","shell.execute_reply":"2025-04-12T14:12:00.803002Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# Predict using each model\nif 'protenix' in ENSEMBLE_MODELS:\n    df_protenix = predict_protenix(test_df)\n    if df_protenix is not None:\n        model_predictions['protenix'] = df_protenix\n    gc.collect()\n    if torch.cuda.is_available(): torch.cuda.empty_cache()\n\nif 'drfold' in ENSEMBLE_MODELS:\n    df_drfold = predict_drfold(test_df)\n    if df_drfold is not None:\n        model_predictions['drfold'] = df_drfold\n    gc.collect()\n    if torch.cuda.is_available(): torch.cuda.empty_cache()\n\nif 'ribonanzanet' in ENSEMBLE_MODELS:\n    df_ribonanzanet = predict_ribonanzanet(test_df)\n    if df_ribonanzanet is not None:\n        model_predictions['ribonanzanet'] = df_ribonanzanet\n    gc.collect()\n    if torch.cuda.is_available(): torch.cuda.empty_cache()\n\n# --- Ensembling Predictions ---\nprint(\"\\n--- Ensembling Predictions ---\")\n\nactive_models = list(model_predictions.keys())\nnum_active_models = len(active_models)\n\nif num_active_models == 0:\n    print(\"ERROR: No models produced predictions. Cannot ensemble.\")\n    # Create a dummy submission with zeros if needed\n    final_submission_df = pd.DataFrame()\n    for i in range(len(test_df)):\n        target_id = test_df.target_id[i]\n        seq = test_df.sequence[i]\n        seq_df = create_zero_predictions_df(target_id, seq)\n        final_submission_df = pd.concat([final_submission_df, seq_df], ignore_index=True)\n\nelif num_active_models == 1:\n    print(f\"Warning: Only one model ({active_models[0]}) produced predictions. Using its output directly.\")\n    final_submission_df = model_predictions[active_models[0]]\n\nelse:\n    print(f\"Ensembling predictions from: {active_models}\")\n    # Use the first available model's prediction as the base structure (ID, resname, resid)\n    base_df = model_predictions[active_models[0]][['ID', 'resname', 'resid']].copy()\n    \n    # Initialize columns for averaged coordinates\n    for k in range(1, NUM_CONFORMATIONS + 1):\n        base_df[f'x_{k}'] = 0.0\n        base_df[f'y_{k}'] = 0.0\n        base_df[f'z_{k}'] = 0.0\n\n    # Sum coordinates from all active models\n    for model_name in active_models:\n        df = model_predictions[model_name]\n        # Ensure alignment (should be correct if generated from test_df in order)\n        if not base_df['ID'].equals(df['ID']):\n             print(f\"Warning: ID mismatch for model {model_name}. Attempting to align...\")\n             df = df.set_index('ID')\n             df = df.reindex(base_df['ID']).reset_index()\n             # Fill NaNs that might result from reindexing (e.g., if one model failed on a sequence the base didn't)\n             # This shouldn't happen with the current error handling, but good practice.\n             df = df.fillna(0.0) \n\n        for k in range(1, NUM_CONFORMATIONS + 1):\n            base_df[f'x_{k}'] += df[f'x_{k}'].astype(float)\n            base_df[f'y_{k}'] += df[f'y_{k}'].astype(float)\n            base_df[f'z_{k}'] += df[f'z_{k}'].astype(float)\n\n    # Average the coordinates\n    for k in range(1, NUM_CONFORMATIONS + 1):\n        base_df[f'x_{k}'] /= num_active_models\n        base_df[f'y_{k}'] /= num_active_models\n        base_df[f'z_{k}'] /= num_active_models\n        \n    final_submission_df = base_df\n    print(\"Ensembling complete.\")\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T14:12:00.804388Z","iopub.execute_input":"2025-04-12T14:12:00.804615Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- Save Submission ---\nprint(f\"\\n--- Saving Final Submission to {OUTPUT_FILE} ---\")\nprint(final_submission_df.head())\nprint(f\"Shape: {final_submission_df.shape}\")\n# Check for NaNs before saving\nif final_submission_df.isnull().values.any():\n    print(\"Warning: NaNs found in the final submission DataFrame. Filling with 0.\")\n    final_submission_df = final_submission_df.fillna(0.0)\n\nfinal_submission_df.to_csv(OUTPUT_FILE, index=False)\n\nprint(\"\\n--- Script Finished ---\")\ntotal_runtime = time.time() - start_time_global\nprint(f\"Total Runtime: {time_to_str(total_runtime, mode='sec')}\")\nmem_alloc, mem_res = gpu_memory_use()\nprint(f\"Final GPU Memory (GB): Allocated={mem_alloc}, Reserved={mem_res}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}