{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.10.12"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":87793,"databundleVersionId":12276181,"sourceType":"competition"},{"sourceId":5123458,"sourceType":"datasetVersion","datasetId":2975803},{"sourceId":10923077,"sourceType":"datasetVersion","datasetId":6785143},{"sourceId":11118830,"sourceType":"datasetVersion","datasetId":6933267}],"dockerImageVersionId":30919,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## Make Boltz-1 predictions","metadata":{}},{"cell_type":"code","source":"# pip install fairscale in dependencies fails\n!pip install --no-index /kaggle/input/fairscale-0413/*whl --no-deps","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pickle\nimport urllib.request\nfrom dataclasses import asdict, dataclass\nimport os\nfrom pathlib import Path\nfrom typing import Literal, Optional\n\nimport pandas as pd\nimport click\nimport torch\nfrom pytorch_lightning import Trainer, seed_everything\nfrom pytorch_lightning.strategies import DDPStrategy\nfrom pytorch_lightning.utilities import rank_zero_only\nfrom tqdm import tqdm\n\nfrom boltz.data.module.inference import BoltzInferenceDataModule\nfrom boltz.data.parse.fasta import parse_fasta\nfrom boltz.data.parse.yaml import parse_yaml\nfrom boltz.data.types import Manifest, Record\nfrom boltz.data.write.writer import BoltzWriter\nfrom boltz.model.model import Boltz1\nfrom Bio.PDB.MMCIF2Dict import MMCIF2Dict\n\nCCD_URL = \"https://huggingface.co/boltz-community/boltz-1/resolve/main/ccd.pkl\"\nMODEL_URL = (\n    \"https://huggingface.co/boltz-community/boltz-1/resolve/main/boltz1_conf.ckpt\"\n)\n\nif os.path.exists('/kaggle/input/stanford-rna-3d-folding/test_sequences.csv'):\n    # run on kaggle\n    TEST_SEQUENCES_PATH = '/kaggle/input/stanford-rna-3d-folding/test_sequences.csv'\n    CACHE_PATH = \"/kaggle/input/rna-prediction-boltz/\"\n    SAMPLE_SUBMISSION = '/kaggle/input/stanford-rna-3d-folding/sample_submission.csv'\nelse:\n    # run locally\n    TEST_SEQUENCES_PATH = 'test_sequences.csv'\n    CACHE_PATH = '/tmp/rna-prediction-boltz/'\n    SAMPLE_SUBMISSION = 'sample_submission.csv'\n    \n@dataclass\nclass BoltzProcessedInput:\n    \"\"\"Processed input data.\"\"\"\n\n    manifest: Manifest\n    targets_dir: Path\n\n\n@dataclass\nclass BoltzDiffusionParams:\n    \"\"\"Diffusion process parameters.\"\"\"\n\n    gamma_0: float = 0.605\n    gamma_min: float = 1.107\n    noise_scale: float = 0.901\n    rho: float = 8\n    step_scale: float = 1.638\n    sigma_min: float = 0.0004\n    sigma_max: float = 160.0\n    sigma_data: float = 16.0\n    P_mean: float = -1.2\n    P_std: float = 1.5\n    coordinate_augmentation: bool = True\n    alignment_reverse_diff: bool = True\n    synchronize_sigmas: bool = True\n    use_inference_model_cache: bool = True\n\n\n@rank_zero_only\ndef download(cache: Path) -> None:\n    \"\"\"Download all the required data.\n\n    Parameters\n    ----------\n    cache : Path\n        The cache directory.\n\n    \"\"\"\n    # Download CCD\n    ccd = cache / \"ccd.pkl\"\n    if not ccd.exists():\n        click.echo(\n            f\"Downloading the CCD dictionary to {ccd}. You may \"\n            \"change the cache directory with the --cache flag.\"\n        )\n        urllib.request.urlretrieve(CCD_URL, str(ccd))  # noqa: S310\n\n    # Download model\n    model = cache / \"boltz1_conf.ckpt\"\n    if not model.exists():\n        click.echo(\n            f\"Downloading the model weights to {model}. You may \"\n            \"change the cache directory with the --cache flag.\"\n        )\n        urllib.request.urlretrieve(MODEL_URL, str(model))  # noqa: S310\n\n\ndef check_inputs(\n    data: Path,\n    outdir: Path,\n    override: bool = False,\n) -> list[Path]:\n    \"\"\"Check the input data and output directory.\n\n    If the input data is a directory, it will be expanded\n    to all files in this directory. Then, we check if there\n    are any existing predictions and remove them from the\n    list of input data, unless the override flag is set.\n\n    Parameters\n    ----------\n    data : Path\n        The input data.\n    outdir : Path\n        The output directory.\n    override: bool\n        Whether to override existing predictions.\n\n    Returns\n    -------\n    list[Path]\n        The list of input data.\n\n    \"\"\"\n    click.echo(\"Checking input data.\")\n\n    # Check if data is a directory\n    if data.is_dir():\n        data: list[Path] = list(data.glob(\"*\"))\n\n        # Filter out non .fasta or .yaml files, raise\n        # an error on directory and other file types\n        filtered_data = []\n        for d in data:\n            if d.suffix in (\".fa\", \".fas\", \".fasta\", \".yml\", \".yaml\"):\n                filtered_data.append(d)\n            elif d.is_dir():\n                msg = f\"Found directory {d} instead of .fasta or .yaml.\"\n                raise RuntimeError(msg)\n            else:\n                msg = (\n                    f\"Unable to parse filetype {d.suffix}, \"\n                    \"please provide a .fasta or .yaml file.\"\n                )\n                raise RuntimeError(msg)\n\n        data = filtered_data\n    else:\n        data = [data]\n\n    # Check if existing predictions are found\n    existing = (outdir / \"predictions\").rglob(\"*\")\n    existing = {e.name for e in existing if e.is_dir()}\n\n    # Remove them from the input data\n    if existing and not override:\n        data = [d for d in data if d.stem not in existing]\n        num_skipped = len(existing) - len(data)\n        msg = (\n            f\"Found some existing predictions ({num_skipped}), \"\n            f\"skipping and running only the missing ones, \"\n            \"if any. If you wish to override these existing \"\n            \"predictions, please set the --override flag.\"\n        )\n        click.echo(msg)\n    elif existing and override:\n        msg = \"Found existing predictions, will override.\"\n        click.echo(msg)\n\n    return data\n\n\n@rank_zero_only\ndef process_inputs(  # noqa: C901, PLR0912, PLR0915\n    data: list[Path],\n    out_dir: Path,\n    ccd_path: Path,\n) -> None:\n    \"\"\"Process the input data and output directory.\n\n    Parameters\n    ----------\n    data : list[Path]\n        The input data.\n    out_dir : Path\n        The output directory.\n    ccd_path : Path\n        The path to the CCD dictionary.\n\n    Returns\n    -------\n    BoltzProcessedInput\n        The processed input data.\n\n    \"\"\"\n    click.echo(\"Processing input data.\")\n    existing_records = None\n\n    # Check if manifest exists at output path\n    manifest_path = out_dir / \"processed\" / \"manifest.json\"\n    if manifest_path.exists():\n        click.echo(f\"Found a manifest file at output directory: {out_dir}\")\n\n        manifest: Manifest = Manifest.load(manifest_path)\n        input_ids = [d.stem for d in data]\n        existing_records, processed_ids = zip(\n            *[\n                (record, record.id)\n                for record in manifest.records\n                if record.id in input_ids\n            ]\n        )\n\n        if isinstance(existing_records, tuple):\n            existing_records = list(existing_records)\n\n        # Check how many examples need to be processed\n        missing = len(input_ids) - len(processed_ids)\n        if not missing:\n            click.echo(\"All examples in data are processed. Updating the manifest\")\n            # Dump updated manifest\n            updated_manifest = Manifest(existing_records)\n            updated_manifest.dump(out_dir / \"processed\" / \"manifest.json\")\n            return\n\n        click.echo(f\"{missing} missing ids. Preprocessing these ids\")\n        missing_ids = list(set(input_ids).difference(set(processed_ids)))\n        data = [d for d in data if d.stem in missing_ids]\n        assert len(data) == len(missing_ids)\n\n    # Create output directories\n    structure_dir = out_dir / \"processed\" / \"structures\"\n    predictions_dir = out_dir / \"predictions\"\n\n    out_dir.mkdir(parents=True, exist_ok=True)\n    structure_dir.mkdir(parents=True, exist_ok=True)\n    predictions_dir.mkdir(parents=True, exist_ok=True)\n\n    # Load CCD\n    with ccd_path.open(\"rb\") as file:\n        ccd = pickle.load(file)  # noqa: S301\n\n    if existing_records is not None:\n        click.echo(f\"Found {len(existing_records)} records. Adding them to records\")\n\n    # Parse input data\n    records: list[Record] = existing_records if existing_records is not None else []\n    for path in tqdm(data):\n        try:\n            # Parse data\n            if path.suffix in (\".fa\", \".fas\", \".fasta\"):\n                target = parse_fasta(path, ccd)\n            elif path.suffix in (\".yml\", \".yaml\"):\n                target = parse_yaml(path, ccd)\n            elif path.is_dir():\n                msg = f\"Found directory {path} instead of .fasta or .yaml, skipping.\"\n                raise RuntimeError(msg)\n            else:\n                msg = (\n                    f\"Unable to parse filetype {path.suffix}, \"\n                    \"please provide a .fasta or .yaml file.\"\n                )\n                raise RuntimeError(msg)\n\n            # Keep record\n            records.append(target.record)\n\n            # Dump structure\n            struct_path = structure_dir / f\"{target.record.id}.npz\"\n            target.structure.dump(struct_path)\n\n        except Exception as e:\n            if len(data) > 1:\n                print(f\"Failed to process {path}. Skipping. Error: {e}.\")\n            else:\n                raise e\n\n    # Dump manifest\n    manifest = Manifest(records)\n    manifest.dump(out_dir / \"processed\" / \"manifest.json\")\n\n\ndef predict(\n    data: str,\n    out_dir: str,\n    cache: str = \"~/.boltz\",\n    checkpoint: Optional[str] = None,\n    devices: int = 1,\n    accelerator: str = \"gpu\",\n    recycling_steps: int = 3,\n    sampling_steps: int = 200,\n    diffusion_samples: int = 1,\n    step_scale: float = 1.638,\n    write_full_pae: bool = False,\n    write_full_pde: bool = False,\n    output_format: Literal[\"pdb\", \"mmcif\"] = \"mmcif\",\n    num_workers: int = 2,\n    override: bool = False,\n    seed: Optional[int] = None,\n) -> None:\n    \"\"\"Run predictions with Boltz-1.\"\"\"\n    # If cpu, write a friendly warning\n    if accelerator == \"cpu\":\n        msg = \"Running on CPU, this will be slow. Consider using a GPU.\"\n        click.echo(msg)\n\n    # Set no grad\n    torch.set_grad_enabled(False)\n\n    # Ignore matmul precision warning\n    torch.set_float32_matmul_precision(\"highest\")\n\n    # Set seed if desired\n    if seed is not None:\n        seed_everything(int(seed))\n\n    # Set cache path\n    cache = Path(cache).expanduser()\n    cache.mkdir(parents=True, exist_ok=True)\n\n    # Create output directories\n    data = Path(data).expanduser()\n    out_dir = Path(out_dir).expanduser()\n    out_dir = out_dir / f\"boltz_results_{data.stem}\"\n    out_dir.mkdir(parents=True, exist_ok=True)\n\n    # Download necessary data and model\n    download(cache)\n\n    # Validate inputs\n    data = check_inputs(data, out_dir, override)\n    if not data:\n        click.echo(\"No predictions to run, exiting.\")\n        return\n\n    # Set up trainer\n    strategy = \"auto\"\n    if (isinstance(devices, int) and devices > 1) or (\n        isinstance(devices, list) and len(devices) > 1\n    ):\n        strategy = DDPStrategy()\n        if len(data) < devices:\n            msg = (\n                \"Number of requested devices is greater \"\n                \"than the number of predictions.\"\n            )\n            raise ValueError(msg)\n\n    msg = f\"Running predictions for {len(data)} structure\"\n    msg += \"s\" if len(data) > 1 else \"\"\n    click.echo(msg)\n\n    # Process inputs\n    ccd_path = cache / \"ccd.pkl\"\n    process_inputs(\n        data=data,\n        out_dir=out_dir,\n        ccd_path=ccd_path,\n    )\n\n    # Load processed data\n    processed_dir = out_dir / \"processed\"\n    processed = BoltzProcessedInput(\n        manifest=Manifest.load(processed_dir / \"manifest.json\"),\n        targets_dir=processed_dir / \"structures\"\n    )\n\n    # Create data module\n    data_module = BoltzInferenceDataModule(\n        manifest=processed.manifest,\n        target_dir=processed.targets_dir,\n        msa_dir=processed_dir / \"msa\",\n        num_workers=num_workers,\n    )\n\n    # Load model\n    if checkpoint is None:\n        checkpoint = cache / \"boltz1_conf.ckpt\"\n\n    predict_args = {\n        \"recycling_steps\": recycling_steps,\n        \"sampling_steps\": sampling_steps,\n        \"diffusion_samples\": diffusion_samples,\n        \"write_confidence_summary\": True,\n        \"write_full_pae\": write_full_pae,\n        \"write_full_pde\": write_full_pde,\n    }\n    diffusion_params = BoltzDiffusionParams()\n    diffusion_params.step_scale = step_scale\n\n    # Add steering arguments\n    steering_args = {\n        \"use_steering\": False,\n        \"steering_strength\": 0.0,\n        \"steering_target\": None,\n        \"fk_steering\": False,\n        \"guidance_update\": False\n    }\n\n    model_module: Boltz1 = Boltz1.load_from_checkpoint(\n        checkpoint,\n        strict=True,\n        predict_args=predict_args,\n        map_location=\"cpu\",\n        diffusion_process_args=asdict(diffusion_params),\n        steering_args=steering_args,\n        ema=False,\n    )\n    model_module.eval()\n\n    # Create prediction writer\n    pred_writer = BoltzWriter(\n        data_dir=processed.targets_dir,\n        output_dir=out_dir / \"predictions\",\n        output_format=output_format,\n    )\n\n    trainer = Trainer(\n        default_root_dir=out_dir,\n        strategy=strategy,\n        callbacks=[pred_writer],\n        accelerator=accelerator,\n        devices=devices,\n        precision=32,\n    )\n\n    # Compute predictions\n    trainer.predict(\n        model_module,\n        datamodule=data_module,\n        return_predictions=False,\n    )\n\ndef preprocess_inputs(test_sequences_path: str, data_dir:str='./inputs_prediction'):\n    sub_file = pd.read_csv(test_sequences_path)\n\n    print(sub_file.head())\n\n    names = sub_file['target_id'].tolist()\n    sequences = sub_file['sequence'].tolist()\n\n    os.makedirs(data_dir, exist_ok=True)\n\n    for tmp_id, tmp_sequence in zip(names, sequences):\n        with open(f'{data_dir}/{tmp_id}.yaml', 'w') as f:\n            f.write(\"constraints: []\\n\")\n            f.write(\"sequences:\\n\")\n            f.write(\"- rna:\\n\")\n            f.write(\"    id:\\n\")\n            f.write(\"    - A1\\n\")\n            f.write(f\"    sequence: {tmp_sequence}\")\n\n\ndef get_coords(tmp_id, idx):\n    cif_file = f\"outputs_prediction/boltz_results_inputs_prediction/predictions/{tmp_id}/{tmp_id}_model_{idx}.cif\"\n\n    mmcif_dict = MMCIF2Dict(cif_file)\n\n    entity_poly_seq = mmcif_dict.get(\"_entity_poly_seq.mon_id\", [])\n    sequence = \"\".join(entity_poly_seq)\n    print(\"RNA sequence:\", sequence)\n\n    x_coords = mmcif_dict[\"_atom_site.Cartn_x\"]\n    y_coords = mmcif_dict[\"_atom_site.Cartn_y\"]\n    z_coords = mmcif_dict[\"_atom_site.Cartn_z\"]\n    atom_names = mmcif_dict[\"_atom_site.label_atom_id\"]\n\n    c1_coords = []\n    for i, atom in enumerate(atom_names):\n        if atom == \"C1'\":\n            c1_coords.append((float(x_coords[i]), float(y_coords[i]), float(z_coords[i])))\n    return c1_coords\n\ndef postprocess_outputs(output_path: str):\n    all_preds = os.listdir('outputs_prediction/boltz_results_inputs_prediction/predictions')\n\n    submission = pd.read_csv(SAMPLE_SUBMISSION)\n    idx = 0\n    for tmp_id in all_preds:\n        print('#'*20, f'inferences for {tmp_id}')\n        for idx in range(5):\n            c1_coords = get_coords(tmp_id, idx)\n            submission.loc[submission['ID'].apply(lambda x: tmp_id in x), [f'x_{idx+1}', f'y_{idx+1}', f'z_{idx+1}']] = c1_coords\n        print()\n    submission.to_csv(output_path, index=False)\n    submission['target_id'] = submission['ID'].apply(lambda x: x.split('_')[0])\n    print(submission.groupby('target_id')['x_1'].mean())\n\n\npreprocess_inputs(test_sequences_path=TEST_SEQUENCES_PATH)\npredict(data=\"./inputs_prediction\",\n        out_dir=\"./outputs_prediction\",\n        cache=CACHE_PATH,\n        diffusion_samples=5,\n        seed=42,\n        override=True)\n\npostprocess_outputs(output_path='boltz_submission.csv')\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Make protenix predictons","metadata":{}},{"cell_type":"code","source":"\nimport os\nimport random\nimport pandas as pd\nimport numpy as np\nimport torch\nfrom tqdm import tqdm\nimport warnings\nwarnings.filterwarnings(\"ignore\")  \n\n#time0=time.time()\n\nprint('IMPORT OK !!!!')\n\nif os.path.exists('/kaggle/input/stanford-rna-3d-folding'):\n    # run on kaggle\n    DATA_KAGGLE_DIR = '/kaggle/input/stanford-rna-3d-folding'\n    CHECKPOINT_PATH = '/kaggle/input/protenix-checkpoints/model_v0.2.0.pt'\n    os.environ[\"PROTENIX_DATA_ROOT_DIR\"] = '/kaggle/input/protenix-checkpoints'\nelse:\n    # run locally\n    DATA_KAGGLE_DIR = '.' \n    CHECKPOINT_PATH = '/protenix/release_data/checkpoint/model_v0.2.0.pt'\n    os.environ[\"PROTENIX_DATA_ROOT_DIR\"] = '/protenix/release_data/ccd_cache'\n\nfrom runner.inference import update_inference_configs, InferenceRunner   \nfrom protenix.data.infer_data_pipeline import InferenceDataset\nfrom configs.configs_base import configs as configs_base\nfrom configs.configs_data import data_configs\nfrom configs.configs_inference import inference_configs\nfrom protenix.config.config import parse_configs\n\n\ndef parse_output_to_df(output, seq, target_id):\n    df = []\n    chain_data = []\n    for i, res in enumerate(seq):\n        d={\"ID\": target_id, \"resname\": res, \"resid\": i + 1}\n        for n in range(len(output)):\n            d={**d, f'x_{n+1}': round(output[n,i,0].item(),3),\n                     f'y_{n+1}': round(output[n,i,1].item(),3),\n                     f'z_{n+1}': round(output[n,i,2].item(),3)}\n        chain_data.append(d)\n\n    if len(chain_data)!=0:\n        chain_df = pd.DataFrame(chain_data)\n        df.append(chain_df)\n        ##print(chain_df)\n    return df\n\ndef seed_everything(seed=0):\n    np.random.seed(seed)\n    torch.random.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    random.seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n    print(f\"Random seed set to {seed}\")\n\nseed_everything()\n\nclass DictDataset(InferenceDataset):\n    def __init__(\n        self,\n        seq_list: list,\n        dump_dir: str,\n        id_list: list = None,\n        use_msa: bool = False,\n    ) -> None:\n\n        self.dump_dir = dump_dir\n        self.use_msa = use_msa\n        if isinstance(id_list,type(None)):\n            self.inputs = [{\"sequences\": \n                            [{\"rnaSequence\": \n                                {\"sequence\": seq, \n                                \"count\": 1}}],\n                            \"name\": \"query\"} for seq in seq_list]\n        else:\n            self.inputs = [{\"sequences\": \n                            [{\"rnaSequence\": \n                                {\"sequence\": seq, \n                                \"count\": 1}}],\n                            \"name\": i} for i, seq in zip(id_list,seq_list)]\n\n\n\nconfigs_base[\"use_deepspeed_evo_attention\"] = (\nos.environ.get(\"USE_DEEPSPEED_EVO_ATTENTION\", False) == \"true\")\nconfigs_base[\"model\"][\"N_cycle\"] = 10 #10\nconfigs_base[\"sample_diffusion\"][\"N_sample\"] = 5\nconfigs_base[\"sample_diffusion\"][\"N_step\"] = 200\ninference_configs['load_checkpoint_path']=CHECKPOINT_PATH\nconfigs = {**configs_base, **{\"data\": data_configs}, **inference_configs}\n\nconfigs = parse_configs(\n        configs=configs,\n        fill_required_with_null=True,\n    )\n#print('CONFIGS:')\n#print(configs)\n#print('DATA CONFIGS:')\n#print(data_configs)\nrunner=InferenceRunner(configs)\n\n\ntest_df=pd.read_csv(TEST_SEQUENCES_PATH)\n\ndataset = DictDataset(test_df.sequence, dump_dir='output', id_list=test_df.target_id, use_msa=False)\nnum_data = len(dataset)\nfor i, seq in tqdm(enumerate(test_df.sequence),total=num_data):\n    try:\n        data, atom_array, data_error_message=dataset[i]\n        target_id = data[\"sample_name\"]\n        assert target_id==test_df.target_id[i]\n        assert data_error_message==''\n        \n        new_configs = update_inference_configs(configs, data[\"N_token\"].item())\n        runner.update_model_configs(new_configs)\n        prediction = runner.predict(data)\n        prediction=prediction['coordinate'][:,data['input_feature_dict']['atom_to_tokatom_idx']==12]\n\n        result = parse_output_to_df(prediction, seq, target_id)[0]\n    except Exception as exc:\n        print(exc)\n        print(f\"Error message: {data_error_message}\")        \n        print('Failed to predict', target_id)\n        result=pd.DataFrame(columns=['ID', 'resname', 'resid', \n                                        'x_1', 'y_1', 'z_1', \n                                        'x_2', 'y_2', 'z_2',\n                                        'x_3', 'y_3', 'z_3', \n                                        'x_4', 'y_4', 'z_4', \n                                        'x_5', 'y_5', 'z_5'], \n                                        data=[[target_id, x, j+1] + [0.0]*15 for j, x in enumerate(seq)])\n        \n    result['ID']=result.apply(lambda x: x.ID + '_' + str(x.resid), axis=1)\n    result.to_csv('protenix_submission.csv', index=False, mode='a', header=(i==0))\n    torch.cuda.empty_cache()\n\nprint(pd.read_csv('protenix_submission.csv'))\n\n\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## BAT average of all predictions","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport random\nimport pickle\nfrom tqdm import tqdm\nimport plotly.graph_objects as go\n\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset\nfrom fastai.vision.all import *\n\nfrom scipy.stats import vonmises\nfrom sklearn.mixture import GaussianMixture\n\ndevice = 'cuda' if torch.cuda.is_available() else 'cpu'\n\ndef givens_rotation(M, k, i, j):\n    \"\"\"\n    Compute the ij Givens rotation matrix to zero out M[k, i].\n    \"\"\"\n    a = M[k, i]\n    b = M[k, j]\n    r = np.hypot(a, b)\n    \n    G = np.eye(3)\n    if r == 0: return G\n    \n#   Compute c and s based on the order of i and j\n    if i > j:\n        c = b / r\n        s = a / r\n    else:\n        c = a / r\n        s = -b / r\n    \n#   Construct the Givens rotation matrix\n    G[i, i] = c\n    G[j, j] = c\n    G[i, j] = s\n    G[j, i] = -s\n    \n    return G","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def xyz_to_BAT(xyz):\n    BA = xyz[:-1] - xyz[1:]\n    b = np.sqrt((BA*BA).sum(1))\n\n    a = np.arccos(-np.sum(BA[:-1]*BA[1:],1) / (b[:-1]*b[1:]))\n\n    t = []\n    for k in range(len(xyz)-3):\n#       Scan ABCD\n        if np.isnan(xyz[k:k+4]).sum() == 0:\n            ABCD = xyz[k:k+4].copy()\n#           Center B\n            ABCD -= ABCD[1]\n#           Rotate BC to x axis, BA to xy plane pointing y, get angle of CD in yz plane\n            ABCD = ABCD@givens_rotation(ABCD, 2, 0, 2)\n            ABCD = ABCD@givens_rotation(ABCD, 2, 0, 1)\n            ABCD = ABCD@givens_rotation(ABCD, 0, 1, 2)\n            if ABCD[2,0] < 0: ABCD[:,[0,2]] *= -1\n            if ABCD[0,1] < 0: ABCD[:,1:] *= -1\n            t.append(np.arctan2(ABCD[3,1],ABCD[3,2]))\n        else:\n            t.append(np.nan)\n    return b,a,t","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def BAT_to_xyz(b,a,t):\n    xyz = np.zeros((len(b)+1,3))\n    if len(b)>0:\n        xyz[1,0] = b[0]\n\n    if len(b)>1:\n        xyz[:2] -= xyz[1]\n        xyz[2,0] = -np.cos(a[0])*b[1]\n        xyz[2,1] = np.sin(a[0])*b[1]\n\n    if len(b)>2:\n        for i in range(3,len(b)+1):\n#           Rotate BC pointing x\n            xyz[:i] = xyz[:i]@givens_rotation(xyz[:i], i-1, 1, 0)\n            xyz[:i] = xyz[:i]@givens_rotation(xyz[:i], i-1, 2, 0)\n#           Rotate BA pointing y\n            xyz[:i] = xyz[:i]@givens_rotation(xyz[:i], i-3, 2, 1)\n            if xyz[i-1,0] < 0: xyz[:i,[0,2]] *= -1\n            if xyz[i-3,1] < 0: xyz[:i,1:] *= -1\n#           Put D at corresponding angle\n            xyz[i,0] = -np.cos(a[i-2])*b[i-1] + xyz[i-1,0]\n            xyz[i,1] = np.sin(a[i-2])*b[i-1] + xyz[i-1,1]\n#           Rotate D to corresponding torsion\n            c = np.cos(np.pi/2 - t[i-3])\n            s = np.sin(np.pi/2 - t[i-3])\n            xyz[i,1:] = xyz[i,1:]@np.array([\n                [c,s],\n                [-s,c]\n            ])\n            xyz[:i+1] -= xyz[i-1]\n\n    return xyz","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_df = pd.read_csv('/kaggle/input/stanford-rna-3d-folding/test_sequences.csv')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"boltz_df = pd.read_csv('boltz_submission.csv')\nboltz_df.tail()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"protenix_df = pd.read_csv('protenix_submission.csv')\nprotenix_df.tail()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"xyz = []\nresid = []\nKAPPA_THRESHOLD = 0.5  # Concentration threshold for torsion reliability\n\nfor SEQ_ID in tqdm(test_df.target_id):\n    # Load and prepare data\n    bdf = boltz_df[boltz_df.ID.apply(lambda v: SEQ_ID in v)].sort_values('resid')\n    pdf = protenix_df[protenix_df.ID.apply(lambda v: SEQ_ID in v)].sort_values('resid')\n    resid.append(bdf.ID.values)\n    \n    # Reshape coordinates (n_residues x 5 models x 3 coordinates)\n    bxyz = bdf[['x_1','y_1','z_1','x_2','y_2','z_2','x_3','y_3','z_3','x_4','y_4','z_4','x_5','y_5','z_5']].values.reshape(-1, 5, 3)\n    pxyz = pdf[['x_1','y_1','z_1','x_2','y_2','z_2','x_3','y_3','z_3','x_4','y_4','z_4','x_5','y_5','z_5']].values.reshape(-1, 5, 3)\n    \n    all_b = []\n    all_a = []\n    all_t = []\n    \n    # Convert all predictions to BAT coordinates\n    for i in range(5):\n        # Process Boltzmann predictions\n        b, a, t = xyz_to_BAT(bxyz[:, i])\n        all_b.append(b)\n        all_a.append(a)\n        all_t.append(t)\n        \n        # Process Protenix predictions\n        b, a, t = xyz_to_BAT(pxyz[:, i])\n        all_b.append(b)\n        all_a.append(a)\n        all_t.append(t)\n    \n    # ===== BONDS: Harmonic mean =====\n    all_b_stack = np.stack(all_b)  # Shape: (10, n_bonds)\n    harmonic_bonds = all_b_stack.shape[0] / np.sum(1 / (all_b_stack + 1e-10), axis=0)\n    \n    # ===== ANGLES: Arithmetic mean =====\n    all_a_stack = np.stack(all_a)\n    mean_angles = np.mean(all_a_stack, axis=0)\n    \n    # ===== TORSIONS: Von Mises/GMM =====\n    all_t_stack = np.stack(all_t)  # Shape: (10, n_torsions)\n    mean_torsions = np.zeros(all_t_stack.shape[1])\n    \n    for torsion_idx in range(all_t_stack.shape[1]):\n        torsions = all_t_stack[:, torsion_idx]\n        \n        # Fit Von Mises distribution\n        try:\n            kappa, loc, _ = vonmises.fit(torsions, fscale=1)\n        except:\n            # Fallback to circular mean if fitting fails\n            x = np.mean(np.cos(torsions))\n            y = np.mean(np.sin(torsions))\n            mean_torsions[torsion_idx] = np.arctan2(y, x)\n            continue\n        \n        # Handle low-concentration torsions with GMM\n        if kappa < KAPPA_THRESHOLD:\n            # Convert angles to 2D representation\n            angle_2d = np.column_stack([np.sin(torsions), np.cos(torsions)])\n            \n            # Fit 2-component Gaussian mixture\n            try:\n                gmm = GaussianMixture(n_components=2, covariance_type='full', random_state=42)\n                gmm.fit(angle_2d)\n                \n                # Convert means back to angles\n                gmm_means = np.arctan2(gmm.means_[:, 0], gmm.means_[:, 1])\n                \n                # Select dominant cluster\n                dominant_cluster = np.argmax(gmm.weights_)\n                mean_torsions[torsion_idx] = gmm_means[dominant_cluster]\n            except:\n                # Fallback to circular mean if GMM fails\n                x = np.mean(np.cos(torsions))\n                y = np.mean(np.sin(torsions))\n                mean_torsions[torsion_idx] = np.arctan2(y, x)\n        else:\n            mean_torsions[torsion_idx] = loc\n    \n    # Reconstruct averaged structure\n    xyz.append(BAT_to_xyz(harmonic_bonds, mean_angles, mean_torsions))\n\nxyz = np.concatenate(xyz)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ensemble_df = pd.DataFrame({\n    'ID': np.concatenate(resid),\n    'x': xyz[:, 0],\n    'y': xyz[:, 1],\n    'z': xyz[:, 2]\n}).set_index('ID')\nensemble_df.tail()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Use top 2 predictions from both","metadata":{}},{"cell_type":"code","source":"boltz_preds = pd.read_csv(\"boltz_submission.csv\", index_col=\"ID\")\nprotenix_preds = pd.read_csv(\"protenix_submission.csv\", index_col=\"ID\")\nsubmission = pd.read_csv(SAMPLE_SUBMISSION, index_col=\"ID\")\n\nfor coord in ['x', 'y', 'z']:\n#   Copy first prediction from boltz to submission position 1\n    submission[f'{coord}_1'] = boltz_preds[f'{coord}_1']\n#   Copy first prediction from protenix to submission position 2\n    submission[f'{coord}_2'] = protenix_preds[f'{coord}_1']\n#   Copy second prediction from boltz to submission position 3\n    submission[f'{coord}_3'] = boltz_preds[f'{coord}_2']\n#   Copy second prediction from protenix to submission position 4\n    submission[f'{coord}_4'] = protenix_preds[f'{coord}_2']\n#   Copy ensemble to submission position 5\n    submission[f'{coord}_5'] = ensemble_df[f'{coord}']","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submission.to_csv('submission.csv', index=True)","metadata":{},"outputs":[],"execution_count":null}]}