{"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":"nvidiaTeslaT4","dataSources":[{"sourceId":87793,"databundleVersionId":12024591,"sourceType":"competition"},{"sourceId":6233380,"sourceType":"datasetVersion","datasetId":3580819},{"sourceId":7395079,"sourceType":"datasetVersion","datasetId":4299455},{"sourceId":7639698,"sourceType":"datasetVersion","datasetId":4299272},{"sourceId":10197079,"sourceType":"datasetVersion","datasetId":6300687},{"sourceId":10855324,"sourceType":"datasetVersion","datasetId":6742586},{"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":11419352,"sourceType":"datasetVersion","datasetId":891011},{"sourceId":224703571,"sourceType":"kernelVersion"},{"sourceId":224896926,"sourceType":"kernelVersion"},{"sourceId":232583388,"sourceType":"kernelVersion"},{"sourceId":235531625,"sourceType":"kernelVersion"}],"dockerImageVersionId":30919,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"Refered Notebooks:\n## https://www.kaggle.com/code/nanacat0520/apply-hooks-to-get-embeddings-ensemble-baseline\nthis notebook is an advanced version of this↑\n\nother notebooks refered:(perfectly compatible)\nhttps://www.kaggle.com/code/ogurtsov/rhofold-ribonanzanet-msas-lb-0-215\nhttps://www.kaggle.com/code/shujun717/ribonanzanet-3d-finetune\n\n","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport torch\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport torch\nimport random\nimport pickle\nimport argparse\nimport os\nimport sys\nimport shutil\nimport yaml\nfrom tqdm import tqdm\nfrom torch import nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import GradScaler\nfrom collections import OrderedDict","metadata":{"trusted":true,"scrolled":true,"execution":{"iopub.status.busy":"2025-04-28T13:19:29.944282Z","iopub.execute_input":"2025-04-28T13:19:29.944663Z","iopub.status.idle":"2025-04-28T13:19:33.986724Z","shell.execute_reply.started":"2025-04-28T13:19:29.944629Z","shell.execute_reply":"2025-04-28T13:19:33.986083Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!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","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T13:19:33.987708Z","iopub.execute_input":"2025-04-28T13:19:33.988090Z","iopub.status.idle":"2025-04-28T13:19:53.535521Z","shell.execute_reply.started":"2025-04-28T13:19:33.988068Z","shell.execute_reply":"2025-04-28T13:19:53.534678Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"torch.manual_seed(0)\nnp.random.seed(0)\nrandom.seed(0)\ntrain_sequences=pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/train_sequences.csv\")\ntrain_labels=pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/train_labels.csv\")\ntrain_labels[\"pdb_id\"] = train_labels[\"ID\"].apply(lambda x: x.split(\"_\")[0]+'_'+x.split(\"_\")[1])\ntrain_labels[\"pdb_id\"]\n\n\n\"\"\"\nRibosanetのconfig\n\"\"\"\nconfig = {\n    \"seed\": 0,\n    \"cutoff_date\": \"2020-01-01\",\n    \"test_cutoff_date\": \"2022-05-01\",\n    \"max_len\": 384,\n    \"batch_size\": 1,\n    \"learning_rate\": 1e-4,\n    \"weight_decay\": 0.0,\n    \"mixed_precision\": \"bf16\",\n    \"model_config_path\": \"../working/configs/pairwise.yaml\",  # Adjust path as needed\n    \"epochs\": 10,\n    \"cos_epoch\": 5,\n    \"loss_power_scale\": 1.0,\n    \"max_cycles\": 1,\n    \"grad_clip\": 0.1,\n    \"gradient_accumulation_steps\": 1,\n    \"d_clamp\": 30,\n    \"max_len_filter\": 9999999,\n    \"min_len_filter\": 10, \n    \"structural_violation_epoch\": 50,\n    \"balance_weight\": False,\n}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T13:19:53.537028Z","iopub.execute_input":"2025-04-28T13:19:53.537286Z","iopub.status.idle":"2025-04-28T13:19:53.958185Z","shell.execute_reply.started":"2025-04-28T13:19:53.537265Z","shell.execute_reply":"2025-04-28T13:19:53.957215Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"losses:","metadata":{}},{"cell_type":"code","source":"def calculate_distance_matrix(X,Y,epsilon=1e-4):\n    return (torch.square(X[:,None]-Y[None,:])+epsilon).sum(-1).sqrt()\n\n\ndef dRMSD(pred_x,\n          pred_y,\n          gt_x,\n          gt_y,\n          epsilon=1e-4,Z=10,d_clamp=None):\n    pred_dm=calculate_distance_matrix(pred_x,pred_y)\n    gt_dm=calculate_distance_matrix(gt_x,gt_y)\n    mask=~torch.isnan(gt_dm)\n    mask[torch.eye(mask.shape[0]).bool()]=False\n    if d_clamp is not None:\n        rmsd=(torch.square(pred_dm[mask]-gt_dm[mask])+epsilon).clip(0,d_clamp**2)\n    else:\n        rmsd=torch.square(pred_dm[mask]-gt_dm[mask])+epsilon\n    return rmsd.sqrt().mean()/Z\n\ndef local_dRMSD(pred_x,\n          pred_y,\n          gt_x,\n          gt_y,\n          epsilon=1e-4,Z=10,d_clamp=30):\n    pred_dm=calculate_distance_matrix(pred_x,pred_y)\n    gt_dm=calculate_distance_matrix(gt_x,gt_y)\n    mask=(~torch.isnan(gt_dm))*(gt_dm<d_clamp)\n    mask[torch.eye(mask.shape[0]).bool()]=False\n    rmsd=torch.square(pred_dm[mask]-gt_dm[mask])+epsilon\n    # rmsd=(torch.square(pred_dm[mask]-gt_dm[mask])+epsilon).sqrt()/Z\n    #rmsd=torch.abs(pred_dm[mask]-gt_dm[mask])/Z\n    return rmsd.sqrt().mean()/Z\n\ndef dRMAE(pred_x, pred_y, gt_x, gt_y, epsilon=1e-4, Z=10, d_clamp=None):\n    \"\"\"\n    dRMAE loss with masking for NaNs in gt_x and gt_y.\n    \"\"\"\n    if pred_x.ndim == 3: pred_x = pred_x.squeeze(0)\n    if pred_y.ndim == 3: pred_y = pred_y.squeeze(0)\n    if gt_x.ndim == 3: gt_x = gt_x.squeeze(0)\n    if gt_y.ndim == 3: gt_y = gt_y.squeeze(0)\n\n    assert pred_x.shape == gt_x.shape, f\"Shape mismatch: pred_x {pred_x.shape}, gt_x {gt_x.shape}\"\n\n    # マスク：NaNがある位置を除外\n    valid_mask = ~(torch.isnan(gt_x).any(dim=-1) | torch.isnan(gt_y).any(dim=-1))\n    pred_x = pred_x[valid_mask]\n    pred_y = pred_y[valid_mask]\n    gt_x = gt_x[valid_mask]\n    gt_y = gt_y[valid_mask]\n\n    if pred_x.shape[0] < 2:\n        return torch.tensor(0.0, device=pred_x.device, dtype=pred_x.dtype)  # too short to compute\n\n    pred_dm = torch.cdist(pred_x, pred_y)\n    gt_dm = torch.cdist(gt_x, gt_y)\n\n    mask = torch.ones_like(gt_dm, dtype=torch.bool)\n    mask.fill_diagonal_(False)\n\n    if d_clamp is not None:\n        pred_dm = pred_dm.clamp(0, d_clamp)\n        gt_dm = gt_dm.clamp(0, d_clamp)\n\n    loss = torch.abs(pred_dm[mask] - gt_dm[mask])\n    return loss.mean() / Z\n\n\ndef align_svd_mae(input, target, Z=10):\n    assert input.shape == target.shape, \"Shape mismatch\"\n\n    mask = ~torch.isnan(target).any(dim=-1)\n    input = input[mask]\n    target = target[mask]\n\n    if input.shape[0] < 3:\n        return torch.tensor(0.0, device=input.device, dtype=input.dtype)  # too short to align\n\n    centroid_input = input.mean(dim=0, keepdim=True)\n    centroid_target = target.mean(dim=0, keepdim=True)\n\n    input_centered = input - centroid_input\n    target_centered = target - centroid_target\n\n    cov_matrix = input_centered.T @ target_centered\n    U, _, Vt = torch.svd(cov_matrix)\n    R = Vt @ U.T\n    if torch.det(R) < 0:\n        Vt[-1, :] *= -1\n        R = Vt @ U.T\n\n    aligned_input = (input_centered @ R.T) + centroid_target\n    return torch.abs(aligned_input - target).mean() / Z\n\n\ndef entropy_regularization(alpha, scale=1.0):\n    p = torch.softmax(alpha, dim=0)\n    log_p = torch.log(p + 1e-8)\n    entropy = -(p * log_p).sum()\n    return scale * entropy\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T16:07:15.260343Z","iopub.execute_input":"2025-04-28T16:07:15.260651Z","iopub.status.idle":"2025-04-28T16:07:15.273881Z","shell.execute_reply.started":"2025-04-28T16:07:15.260625Z","shell.execute_reply":"2025-04-28T16:07:15.272986Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Precomputing embeddings from Rhofold&RibonanzaNet","metadata":{}},{"cell_type":"code","source":"\"\"\"\ndef print_model_layers(model, model_name=\"model\"):\n    print(f\"\\n🔎 [{model_name}] named_modules:\")\n    for name, module in model.named_modules():\n        print(f\" - {name}: {type(module).__name__}\")\n\n# RhoFold\nmodel_rho = RhoFold(rho_config)\n#print_model_layers(model_rho, \"RhoFold\")\n\n# RibonanzaNet\nmodel_ribo = finetuned_RibonanzaNet(\n    load_config_from_yaml(\"/kaggle/input/ribonanzanet2d-final/configs/pairwise.yaml\"),\n    pretrained=False\n)\n\n\nprint_model_layers(model_ribo, \"RibonanzaNet\")\n\"\"\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T13:19:53.974776Z","iopub.execute_input":"2025-04-28T13:19:53.975117Z","iopub.status.idle":"2025-04-28T13:19:53.992718Z","shell.execute_reply.started":"2025-04-28T13:19:53.975078Z","shell.execute_reply":"2025-04-28T13:19:53.991880Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================\n# RNA3D Pipeline with Hook-Based Feature Fusion (Improved)\n# ============================\n\nimport os\nimport sys\nimport shutil\nimport torch\nimport numpy as np\nimport pandas as pd\nimport yaml\nimport random\nfrom tqdm import tqdm\nfrom torch import nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import GradScaler\nfrom collections import OrderedDict\nimport torch._dynamo\n\n# -- Torch Dynamism Setup --\ntorch._dynamo.config.suppress_errors = True\ntorch._dynamo.config.verbose = True\ntorch._dynamo.disable()\n\n# -- Reproducibility --\nos.environ[\"CUDA_LAUNCH_BLOCKING\"] = \"1\"\ntorch.manual_seed(0)\nnp.random.seed(0)\nrandom.seed(0)\n\n# ============================\n# RhoFold Setup\n# ============================\nsrc_dir = \"/kaggle/input/rhofold-repo/rhofold\"\ndst_dir = \"/kaggle/working/rhofold\"\nif not os.path.exists(dst_dir):\n    shutil.copytree(src_dir, dst_dir)\n\ndef rewrite_imports(filepath, from_level=1):\n    if not os.path.exists(filepath): return\n    with open(filepath, \"r\") as f:\n        lines = f.readlines()\n    new_lines = []\n    for line in lines:\n        if \"from rhofold.\" in line:\n            new_lines.append(line.replace(\"from rhofold.\", \"from \" + \".\" * from_level))\n        elif \"import rhofold.\" in line:\n            new_lines.append(line.replace(\"import rhofold.\", \"import \" + \".\" * from_level))\n        else:\n            new_lines.append(line)\n    with open(filepath, \"w\") as f:\n        f.writelines(new_lines)\n\nrewrite_imports(os.path.join(dst_dir, \"rhofold.py\"), from_level=1)\nrewrite_imports(os.path.join(dst_dir, \"data\", \"data_pipeline.py\"), from_level=2)\nrewrite_imports(os.path.join(dst_dir, \"utils\", \"tensor_utils.py\"), from_level=2)\nsys.path.insert(0, \"/kaggle/working\")\n\nfrom rhofold.rhofold import RhoFold\nfrom rhofold.config import rhofold_config as rho_config\nfrom rhofold.utils.alphabet import get_features\n\n# ============================\n# RibonanzaNet Setup\n# ============================\nsys.path.append(\"/kaggle/input/ribonanzanet2d-final\")\nfrom Network import RibonanzaNet\n\nclass Config:\n    def __init__(self, **entries):\n        self.__dict__.update(entries)\n        self.entries = 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 FusionHead(nn.Module):\n    def __init__(self, hidden_dim=256, out_dim=3, dropout=0.1):\n        super().__init__()\n        self.block = nn.Sequential(\n            nn.Linear(hidden_dim, hidden_dim), nn.ReLU(), nn.Dropout(dropout),\n            nn.Linear(hidden_dim, hidden_dim), nn.ReLU(), nn.Dropout(dropout),\n            nn.Linear(hidden_dim, out_dim)\n        )\n        self.residual = nn.Linear(hidden_dim, out_dim)\n\n    def forward(self, x):\n        return self.block(x) + self.residual(x)\n\nclass TransformerDecoderHead(nn.Module):\n    def __init__(self, d_model=256, nhead=8, num_layers=2, dropout=0.1):\n        super().__init__()\n        encoder_layer = nn.TransformerEncoderLayer(d_model=d_model, nhead=nhead, dropout=dropout, batch_first=True)\n        self.transformer = nn.TransformerEncoder(encoder_layer, num_layers=num_layers)\n        self.fusion_head = FusionHead(hidden_dim=d_model, out_dim=3)\n\n    def forward(self, x):\n        x = self.transformer(x.unsqueeze(0)).squeeze(0)\n        return self.fusion_head(x)\n\nclass finetuned_RibonanzaNet(RibonanzaNet):\n    def __init__(self, config, pretrained=False):\n        config.dropout = 0.2\n        super().__init__(config)\n        if pretrained:\n            self.load_state_dict(torch.load(\"/kaggle/input/ribonanzanet-3d-finetune/RibonanzaNet-3D-final.pt\", map_location='cpu'), strict=False)\n        self.decoder_head = TransformerDecoderHead(d_model=256)\n\n    def get_finetuned_embeddings(self, src):\n        if src.dim() == 1:\n            src = src.unsqueeze(0)\n        seq_feat, pair_feat = self.get_embeddings(src, torch.ones_like(src).long().to(src.device))\n        return {\n            'sequence_features': seq_feat,\n            'pairwise_features': pair_feat,\n            'xyz_predictions': self.decoder_head(seq_feat)\n        }\n\n    def forward(self, src):\n        sequence_features, _ = self.get_embeddings(src, torch.ones_like(src).long().to(src.device))\n    \n        # 🔧 4D → 3D or 2Dに変換\n        if sequence_features.dim() == 4:\n            sequence_features = sequence_features.squeeze(0).mean(dim=1)\n        elif sequence_features.dim() == 3 and sequence_features.shape[0] == 1:\n            sequence_features = sequence_features.squeeze(0)\n    \n        xyz = self.decoder_head(sequence_features)\n        return xyz\n\n\n\nclass GatedFusionLayer(nn.Module):\n    def __init__(self, input_dim=256):\n        super().__init__()\n        self.gate = nn.Sequential(\n            nn.Linear(input_dim * 2, input_dim), nn.ReLU(),\n            nn.Linear(input_dim, 1), nn.Sigmoid()\n        )\n        self.fusion = nn.Linear(input_dim * 2, input_dim)\n\n    def forward(self, x1, x2):\n        concat = torch.cat([x1, x2], dim=-1)\n        gate_val = self.gate(concat)\n        fusion_input = gate_val * x1 + (1 - gate_val) * x2\n        return self.fusion(torch.cat([x1, x2], dim=-1))\n\nall_layer_outputs = OrderedDict()\n\n\"\"\"\n# ✅ 新しいHook対象層\nRIBONANZANET_HOOK_TARGETS = [\n    # 中間層\n    \"transformer_encoder.3.self_attn.fc\",\n    \"transformer_encoder.3.triangle_update_out.to_out\",\n    \"transformer_encoder.3.triangle_update_in.to_out\",\n    \"transformer_encoder.3.outer_product_mean.proj_down1\",\n    # 後半層\n    \"transformer_encoder.6.triangle_update_out.to_out\",\n    \"transformer_encoder.6.outer_product_mean.proj_down1\",\n    \"transformer_encoder.6.pair_transition.3\",\n    # 前半層\n    \"transformer_encoder.1.triangle_update_in.to_out\",\n]\n\nRHOFOLD_HOOK_TARGETS = [\n    # 中間・後層\n    \"e2eformer.blocks.5.core.tri_att_end.mha.linear_o\",\n    \"e2eformer.blocks.6.core.outer_product_mean.linear_out\",\n    \"e2eformer.blocks.6.core.pair_transition.linear_2\",\n    \"structure_module.ipa.linear_out\",\n    \"structure_module.transition.layers.0.linear_3\",\n    \"structure_module.angle_resnet.linear_out\",\n    # 前半層\n    \"e2eformer.blocks.2.core.tri_att_start.mha.linear_o\",\n    \"e2eformer.blocks.2.core.outer_product_mean.linear_out\",\n    # 中間層\n    \"e2eformer.blocks.4.core.pair_transition.linear_2\",\n]\"\"\"\n\n\"\"\"RIBONANZANET_HOOK_TARGETS = [\n    \"transformer_encoder.3.self_attn.fc\",\n    \"transformer_encoder.3.triangle_update_out.to_out\",\n    \"transformer_encoder.3.triangle_update_in.to_out\",\n    \"transformer_encoder.3.outer_product_mean.proj_down1\",\n    \"transformer_encoder.6.triangle_update_out.to_out\",\n    \"transformer_encoder.6.outer_product_mean.proj_down1\",\n    \"transformer_encoder.6.pair_transition.3\",\n    \"transformer_encoder.1.triangle_update_in.to_out\",\n]\n\nRHOFOLD_HOOK_TARGETS = [\n    \"e2eformer.blocks.5.core.tri_att_end.mha.linear_o\",\n    \"e2eformer.blocks.6.core.outer_product_mean.linear_out\",\n    \"e2eformer.blocks.6.core.pair_transition.linear_2\",\n    \"structure_module.ipa.linear_out\",\n    \"structure_module.transition.layers.0.linear_3\",\n    \"structure_module.angle_resnet.linear_out\",\n    \"e2eformer.blocks.2.core.tri_att_start.mha.linear_o\",\n    \"e2eformer.blocks.2.core.outer_product_mean.linear_out\",\n    \"e2eformer.blocks.4.core.pair_transition.linear_2\",\n]\"\"\"\n\nRIBONANZANET_HOOK_TARGETS = [\n    \"transformer_encoder.3.self_attn.fc\",                    # 中間層・局所情報\n    \"transformer_encoder.3.triangle_update_in.to_out\",      # 中間層・3残基関係\n    \"transformer_encoder.6.outer_product_mean.proj_down1\",  # 最終層・2残基関係の集約\n]\n\nRHOFOLD_HOOK_TARGETS = [\n    \"structure_module.ipa.linear_out\",                         # 3D配置に直接寄与\n    \"structure_module.angle_resnet.linear_out\",                # 二面角出力\n    \"e2eformer.blocks.6.core.outer_product_mean.linear_out\",   # 高次層の残基間関係\n]\n\n\"\"\"def get_full_hook(name):\n    def hook(module, input, output):\n        if isinstance(output, torch.Tensor):\n            all_layer_outputs[name] = output.detach().cpu().clone()\n    return hook\"\"\"\ndef get_full_hook(name):\n    def hook(module, input, output):\n        try:\n            if isinstance(output, tuple):\n                out = [o.detach().cpu().clone() for o in output if isinstance(o, torch.Tensor)]\n                if len(out) > 0:\n                    all_layer_outputs[name] = _sanitize_layer_output(out[0])\n            elif isinstance(output, torch.Tensor):\n                all_layer_outputs[name] = _sanitize_layer_output(output.detach().cpu().clone())\n        except Exception as e:\n            print(f\"[hook error] {name}: {e}\")\n    return hook\n\ndef register_hooks(model):\n    for name, module in model.named_modules():\n        if name in RIBONANZANET_HOOK_TARGETS:\n            module.register_forward_hook(get_full_hook(name))\n\ndef register_rhofold_hooks(model):\n    for name, module in model.named_modules():\n        if name in RHOFOLD_HOOK_TARGETS:\n            module.register_forward_hook(get_full_hook(name))\n\nmodel_rho = RhoFold(rho_config).to(\"cuda\").eval()\nmodel_rho.structure_module.refinenet = None\nmodel_rho.load_state_dict(torch.load(\"/kaggle/input/rhofold-repo/pretrained/RhoFold_pretrained.pt\", map_location=\"cuda\")[\"model\"])\n\ndef get_final_embeddings(fasta_path, model):\n    all_layer_outputs.clear()\n    register_hooks(model, RHOFOLD_HOOK_TARGETS)\n\n    data_dict = get_features(fasta_path, fasta_path)\n\n    with torch.no_grad():\n        _ = model(\n            tokens=data_dict['tokens'].to(\"cuda\"),\n            rna_fm_tokens=data_dict['rna_fm_tokens'].to(\"cuda\"),\n            seq=data_dict['seq']\n        )\n\n    return {\n    k: _sanitize_layer_output(v[0] if isinstance(v, list) else v).cpu()\n    for k, v in all_layer_outputs.items()\n    if isinstance((v[0] if isinstance(v, list) else v), torch.Tensor)\n    }\n\n\n\ndef get_ribonanzanet_embeddings(model, sequence_tensor):\n    all_layer_outputs.clear()\n    register_hooks(model)\n    with torch.no_grad():\n        _ = model(sequence_tensor.to(\"cuda\"))\n    return {\n        k: _sanitize_layer_output(v[0] if isinstance(v, list) else v).cpu()\n        for k, v in all_layer_outputs.items()\n        if k in RIBONANZANET_HOOK_TARGETS\n    }\n\n\ntrain_sequences = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/train_sequences.csv\")\ntrain_labels = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/train_labels.csv\")\ntrain_labels[\"pdb_id\"] = train_labels[\"ID\"].apply(lambda x: x.split(\"_\")[0] + \"_\" + x.split(\"_\")[1])\n\nall_xyz = []\nfor pdb_id in tqdm(train_sequences['target_id']):\n    df = train_labels[train_labels[\"pdb_id\"] == pdb_id]\n    xyz = df[['x_1', 'y_1', 'z_1']].to_numpy().astype('float32')\n    xyz[np.isnan(xyz) | (xyz < -1e17)] = np.nan\n    all_xyz.append(xyz)\n\nconfig_data = {'cutoff_date': '2021-12-01', 'test_cutoff_date': '2022-08-01', 'min_len_filter': 50, 'max_len_filter': 400, 'max_len': 400}\n\nfilter_nan = [\n    (np.isnan(xyz).mean() <= 0.5) and (config_data['min_len_filter'] < len(xyz) < config_data['max_len_filter'])\n    for xyz in all_xyz\n]\nnon_nan_indices = np.where(filter_nan)[0]\ntrain_sequences = train_sequences.loc[non_nan_indices].reset_index(drop=True)\nall_xyz = [all_xyz[i] for i in non_nan_indices]\n\ndata = {\n    \"sequence\": train_sequences['sequence'].tolist(),\n    \"temporal_cutoff\": train_sequences['temporal_cutoff'].tolist(),\n    \"xyz\": all_xyz\n}\n\ncutoff = pd.Timestamp(config_data['cutoff_date'])\ntrain_index = [i for i, d in enumerate(data['temporal_cutoff']) if pd.Timestamp(d) <= cutoff]\n\nclass RNA3D_Dataset(Dataset):\n    def __init__(self, indices, data, model):\n        self.indices = indices\n        self.data = data\n        self.model = model\n        self.tokens = {nt: i for i, nt in enumerate('ACGU')}\n\n    def __len__(self):\n        return len(self.indices)\n\n    def __getitem__(self, idx):\n        idx = self.indices[idx]\n        sequence_str = self.data['sequence'][idx]\n        token_ids = [self.tokens.get(nt, 0) for nt in sequence_str]\n        sequence_tensor = torch.tensor(token_ids, dtype=torch.long).unsqueeze(0)\n\n        if sequence_tensor.size(1) > config_data['max_len']:\n            start = np.random.randint(sequence_tensor.size(1) - config_data['max_len'])\n            sequence_tensor = sequence_tensor[:, start:start + config_data['max_len']]\n            xyz = self.data['xyz'][idx][start:start + config_data['max_len']]\n        else:\n            xyz = self.data['xyz'][idx]\n\n        xyz = torch.tensor(xyz, dtype=torch.float32)\n        ribo_features = get_ribonanzanet_embeddings(self.model, sequence_tensor)\n\n        return {\n            'sequence': sequence_tensor,\n            'xyz': xyz,\n            'ribonanza_features': ribo_features,\n            'target_id': train_sequences.iloc[idx]['target_id']\n        }\n\nribo_config = load_config_from_yaml(\"/kaggle/input/ribonanzanet2d-final/configs/pairwise.yaml\")\nmodel_ribo = finetuned_RibonanzaNet(ribo_config, pretrained=True).to(\"cuda\").eval()\ntrain_dataset = RNA3D_Dataset(train_index, data, model_ribo)\ntrain_loader = DataLoader(train_dataset, batch_size=1, shuffle=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T14:13:42.088129Z","iopub.execute_input":"2025-04-28T14:13:42.088465Z","iopub.status.idle":"2025-04-28T14:13:55.366522Z","shell.execute_reply.started":"2025-04-28T14:13:42.088437Z","shell.execute_reply":"2025-04-28T14:13:55.365782Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Training","metadata":{}},{"cell_type":"code","source":"\"\"\"RIBONANZANET_HOOK_TARGETS = [\n    \"transformer_encoder.3.self_attn.fc\",\n    \"transformer_encoder.3.triangle_update_out.to_out\",\n    \"transformer_encoder.3.triangle_update_in.to_out\",\n    \"transformer_encoder.3.outer_product_mean.proj_down1\",\n    \"transformer_encoder.6.triangle_update_out.to_out\",\n    \"transformer_encoder.6.outer_product_mean.proj_down1\",\n    \"transformer_encoder.6.pair_transition.3\",\n    \"transformer_encoder.1.triangle_update_in.to_out\",\n]\n\nRHOFOLD_HOOK_TARGETS = [\n    \"e2eformer.blocks.5.core.tri_att_end.mha.linear_o\",\n    \"e2eformer.blocks.6.core.outer_product_mean.linear_out\",\n    \"e2eformer.blocks.6.core.pair_transition.linear_2\",\n    \"structure_module.ipa.linear_out\",\n    \"structure_module.transition.layers.0.linear_3\",\n    \"structure_module.angle_resnet.linear_out\",\n    \"e2eformer.blocks.2.core.tri_att_start.mha.linear_o\",\n    \"e2eformer.blocks.2.core.outer_product_mean.linear_out\",\n    \"e2eformer.blocks.4.core.pair_transition.linear_2\",\n]\"\"\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T13:20:15.709253Z","iopub.execute_input":"2025-04-28T13:20:15.709532Z","iopub.status.idle":"2025-04-28T13:20:15.714467Z","shell.execute_reply.started":"2025-04-28T13:20:15.709511Z","shell.execute_reply":"2025-04-28T13:20:15.713731Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================\n# Precompute & Combine Pipeline\n# ============================\nimport os\nimport sys\nimport shutil\nimport torch\nimport numpy as np\nimport pandas as pd\nimport yaml\nimport random\nimport pickle\nfrom tqdm import tqdm\nfrom torch import nn\nfrom collections import OrderedDict\n\n# ============================\n# Utility & Combiner Modules\n# ============================\n\nclass AttentionCombiner(nn.Module):\n    def __init__(self, input_dims, learn_init=False):\n        super().__init__()\n        self.num_layers = len(input_dims)\n        self.input_dims = input_dims\n        self.proj_layers = nn.ModuleList([\n            nn.Linear(dim, input_dims[0]) for dim in input_dims\n        ])\n        init_weights = (\n            torch.ones(self.num_layers) / self.num_layers\n            if not learn_init else torch.randn(self.num_layers)\n        )\n        self.alpha = nn.Parameter(init_weights)\n\n    def forward(self, layer_outputs):\n        projected = [proj(t) for proj, t in zip(self.proj_layers, layer_outputs)]\n        stack = torch.stack(projected, dim=0)  # [N_layers, L, C0]\n        weights = torch.softmax(self.alpha, dim=0)  # [N_layers]\n        weighted = (weights.view(-1, 1, 1) * stack).sum(dim=0)\n        return weighted\n\nclass TransformerCombiner(nn.Module):\n    def __init__(self, input_dims, hidden_dim=256, nhead=8, nlayers=2):\n        super().__init__()\n        self.num_layers = len(input_dims)\n        self.input_proj = nn.ModuleList([\n            nn.Linear(d, hidden_dim) for d in input_dims\n        ])\n        encoder_layer = nn.TransformerEncoderLayer(\n            d_model=hidden_dim, nhead=nhead, batch_first=True\n        )\n        self.transformer = nn.TransformerEncoder(encoder_layer, num_layers=nlayers)\n        self.output_proj = nn.Linear(hidden_dim, hidden_dim)\n\n    def forward(self, layer_outputs):\n        clean_outputs = []\n        for idx, o in enumerate(layer_outputs):\n            if o.dim() == 4:\n                o = o.squeeze(0).mean(dim=1)\n            elif o.dim() == 3 and o.shape[0] == 1:\n                o = o.squeeze(0)\n            if o.dim() != 2:\n                raise ValueError(f\"Layer {idx} has invalid shape: {o.shape}\")\n            for idx, o in enumerate(layer_outputs):\n                o = _sanitize_layer_output(o)\n                clean_outputs.append(o)\n        projected = [proj(o) for proj, o in zip(self.input_proj, clean_outputs)]\n        stack = torch.stack(projected, dim=1)  # [L, N_layers, H]\n        stack = stack.permute(1, 0, 2)         # [N_layers, L, H]\n        fused = self.transformer(stack)   # [N_layers, L, H]\n        fused = fused.mean(dim=0)         # ✅ [L, H]\n        output = self.output_proj(fused)  # → [L, H]\n        return self.output_proj(output)\n\nclass GatedFusionLayer(nn.Module):\n    def __init__(self, input_dim=256):\n        super().__init__()\n        self.gate = nn.Sequential(\n            nn.Linear(input_dim * 2, input_dim),\n            nn.ReLU(),\n            nn.Linear(input_dim, 1),\n            nn.Sigmoid()\n        )\n        self.fusion = nn.Linear(input_dim * 2, input_dim)\n\n    def forward(self, x1, x2):\n        concat = torch.cat([x1, x2], dim=-1)\n        gate_val = self.gate(concat)\n        fusion_input = gate_val * x1 + (1 - gate_val) * x2\n        return self.fusion(torch.cat([x1, x2], dim=-1))\n\nclass FeatureReshaper(nn.Module):\n    def __init__(self, input_shapes, target_dim=256):\n        super().__init__()\n        self.projectors = nn.ModuleDict()\n        self.name_map = {}  # 逆変換用\n\n        for name, shape in input_shapes.items():\n            safe_name = name.replace('.', '__dot__')\n            self.name_map[name] = safe_name\n            self.projectors[safe_name] = nn.Sequential(\n                nn.LayerNorm(shape[-1]),\n                nn.Linear(shape[-1], target_dim)\n            )\n\n    def forward(self, features: dict):\n        reshaped = {}\n        for name, feat in features.items():\n            safe_name = name.replace('.', '__dot__')\n            #print(f\"[FeatureReshaper] Processing {name} (→ {safe_name}) with shape {feat.shape}\")\n            reshaped[name] = self.projectors[safe_name](feat)\n            #print(f\"[FeatureReshaper] Reshaped {name} → {reshaped[name].shape}\")\n        return reshaped\n\n\n# ============================\n# Hook関連ユーティリティ\n# ============================\n\nall_layer_outputs = OrderedDict()\n\n\"\"\"RIBONANZANET_HOOK_TARGETS = [\n    \"transformer_encoder.3.self_attn.fc\",\n    \"transformer_encoder.3.triangle_update_out.to_out\",\n    \"transformer_encoder.3.triangle_update_in.to_out\",\n    \"transformer_encoder.3.outer_product_mean.proj_down1\",\n    \"transformer_encoder.6.triangle_update_out.to_out\",\n    \"transformer_encoder.6.outer_product_mean.proj_down1\",\n    \"transformer_encoder.6.pair_transition.3\",\n    \"transformer_encoder.1.triangle_update_in.to_out\",\n]\n\nRHOFOLD_HOOK_TARGETS = [\n    \"e2eformer.blocks.5.core.tri_att_end.mha.linear_o\",\n    \"e2eformer.blocks.6.core.outer_product_mean.linear_out\",\n    \"e2eformer.blocks.6.core.pair_transition.linear_2\",\n    \"structure_module.ipa.linear_out\",\n    \"structure_module.transition.layers.0.linear_3\",\n    \"structure_module.angle_resnet.linear_out\",\n    \"e2eformer.blocks.2.core.tri_att_start.mha.linear_o\",\n    \"e2eformer.blocks.2.core.outer_product_mean.linear_out\",\n    \"e2eformer.blocks.4.core.pair_transition.linear_2\",\n]\"\"\"\n\nRIBONANZANET_HOOK_TARGETS = [\n    \"transformer_encoder.3.self_attn.fc\",                    # 中間層・局所情報\n    \"transformer_encoder.3.triangle_update_in.to_out\",      # 中間層・3残基関係\n    \"transformer_encoder.6.outer_product_mean.proj_down1\",  # 最終層・2残基関係の集約\n]\n\nRHOFOLD_HOOK_TARGETS = [\n    \"structure_module.ipa.linear_out\",                         # 3D配置に直接寄与\n    \"structure_module.angle_resnet.linear_out\",                # 二面角出力\n    \"e2eformer.blocks.6.core.outer_product_mean.linear_out\",   # 高次層の残基間関係\n]\n\n\n\ndef _sanitize_layer_output(tensor):\n    if tensor.dim() == 4:\n        return tensor.squeeze(0).mean(dim=1)\n    elif tensor.dim() == 3 and tensor.shape[0] == 1:\n        return tensor.squeeze(0)\n    elif tensor.dim() == 3 and tensor.shape[1] == 1:\n        return tensor.squeeze(1)\n    elif tensor.dim() == 2:\n        return tensor\n    else:\n        raise ValueError(f\"Unexpected tensor shape: {tensor.shape}\")\n\n\n\ndef register_hooks(model, hook_targets):\n    for name, module in model.named_modules():\n        if name in hook_targets:\n            module.register_forward_hook(get_full_hook(name))\n\n# ============================\n# Embedding Precomputation\n# ============================\n\ndef precompute_embeddings(model_ribo, model_rho):\n    embeddings_cache = {}\n    tokens = {nt: i for i, nt in enumerate('ACGU')}\n\n    for idx in tqdm(train_index):\n        #print(idx)\n        target_id = train_sequences.iloc[idx]['target_id']\n        sequence_tensor = torch.tensor([\n            tokens.get(nt, 0) for nt in data['sequence'][idx]\n        ], dtype=torch.long).unsqueeze(0)\n\n        if target_id in embeddings_cache:\n            print(\"in_cashe\")\n            continue\n\n        fasta_path = f\"/kaggle/input/stanford-rna-3d-folding/MSA/{target_id}.MSA.fasta\"\n        if not os.path.exists(fasta_path):\n            print(\"No!\")\n            continue\n\n        try:\n            all_layer_outputs.clear()\n            register_hooks(model_rho, RHOFOLD_HOOK_TARGETS)\n            rho_embs = get_final_embeddings(fasta_path, model_rho)\n\n            rho_valid = {\n                k: _sanitize_layer_output(v.cpu())\n                for k, v in rho_embs.items()\n                if isinstance(v, torch.Tensor)\n            }\n            \"\"\"for k, v in rho_valid.items():\n                print(f\"[hook shape check] {k} → shape: {v.shape}\")\"\"\"\n\n            all_layer_outputs.clear()\n            #print('register_hooks')\n            register_hooks(model_ribo, RIBONANZANET_HOOK_TARGETS)\n            with torch.no_grad():\n                #print('model_ribo')\n                _ = model_ribo(sequence_tensor.cuda())\n                #print('model ribo done')\n            ribo_valid = {\n                k: v.cpu()\n                for k, v in all_layer_outputs.items()\n                if k in RIBONANZANET_HOOK_TARGETS and isinstance(v, torch.Tensor)\n            }\n            \"\"\"for k, v in ribo_valid.items():\n                print(f\"[hook shape check] {k} → shape: {v.shape}\")\"\"\"\n            #print(\"starting FeatureReshape..\")\n            reshaper_rho = FeatureReshaper(input_shapes={k: v.shape for k, v in rho_valid.items()}).cuda()\n            \n            reshaper_ribo = FeatureReshaper(input_shapes={k: v.shape for k, v in ribo_valid.items()}).cuda()\n            #print(\"resh_rho\",reshape_rho.shape)\n            #print(\"resh_ribo\",reshape_ribo.shape)\n            rho_reshaped = reshaper_rho({\n                k: _sanitize_layer_output(v.cuda()) for k, v in rho_valid.items()\n            })\n            ribo_reshaped = reshaper_ribo({\n                k: _sanitize_layer_output(v.cuda()) for k, v in ribo_valid.items()\n            })\n            #print(\"resh_rho\",reshape_rho.shape)\n            #print(\"resh_ribo\",reshape_ribo.shape)\n\n            embeddings_cache[target_id] = {\n                'rho_emb': rho_reshaped,\n                'ribo_emb': ribo_reshaped\n            }\n\n            torch.cuda.empty_cache()\n            #print(f\"✅ Extracted features from {fasta_path}: {[k for k in all_layer_outputs.keys()]}\")\n\n        except Exception as e:\n            print(f\"❌ Error {target_id}: {e}\")\n            continue\n\n    return embeddings_cache\n\noutput_dir = \"/kaggle/working/embeddings_cache\"\nos.makedirs(output_dir, exist_ok=True)\n\ndef save_cache_pickle(embeddings_cache, filename=\"embeddings_cache.pkl\"):\n    file_path = os.path.join(output_dir, filename)\n    with open(file_path, \"wb\") as f:\n        pickle.dump(embeddings_cache, f)\n    return file_path\n\ndef load_cache_pickle(filename=\"embeddings_cache.pkl\"):\n    file_path = os.path.join(output_dir, filename)\n    with open(file_path, \"rb\") as f:\n        return pickle.load(f)\n\nribo_config = load_config_from_yaml(\"/kaggle/input/ribonanzanet2d-final/configs/pairwise.yaml\")\nmodel_ribo = finetuned_RibonanzaNet(ribo_config, pretrained=False).cuda()\nmodel_ribo.load_state_dict(torch.load(\"/kaggle/input/ribonanzanet-3d-finetune/RibonanzaNet-3D.pt\"), strict=False)\nmodel_ribo.eval()\n\n_example_fasta = \"/kaggle/input/stanford-rna-3d-folding/MSA/17RA_A.MSA.fasta\"\n_example_target_id = \"17RA_A\"\n_example_rho_embs = get_final_embeddings(_example_fasta, model=model_rho)\n\nprint(f\"[{_example_target_id}] Rho keys: {list(_example_rho_embs.keys())}\")\n_example_layer_dict = {\n    k: _sanitize_layer_output(v[0] if isinstance(v, list) else v)\n    for k, v in _example_rho_embs.items()\n    if isinstance((v[0] if isinstance(v, list) else v), torch.Tensor)\n}\nprint(\"✅ Shapes to be passed to TransformerCombiner:\")\nfor k, t in _example_layer_dict.items():\n    print(f\" - {k}: {t.shape}\")\n_input_dims = [x.shape[-1] for x in _example_layer_dict.values()]\n\nattention_combiner = TransformerCombiner(input_dims=_input_dims).cuda()\nfusion_module = GatedFusionLayer().cuda().train()\n\noptimizer = torch.optim.Adam(\n    list(fusion_module.parameters()) + list(attention_combiner.parameters()),\n    lr=1e-4\n)\n\nprint(\"🚀 Starting precompute_embeddings()\")\nembeddings_cache = precompute_embeddings(model_ribo, model_rho)\n\ntry:\n    pickle_path = save_cache_pickle(embeddings_cache)\nexcept Exception as e:\n    print(f\"❌ Failed to save embeddings: {e}\")\n\nembeddings_cache = load_cache_pickle()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T14:13:57.811579Z","iopub.execute_input":"2025-04-28T14:13:57.811899Z","iopub.status.idle":"2025-04-28T15:21:23.948456Z","shell.execute_reply.started":"2025-04-28T14:13:57.811875Z","shell.execute_reply":"2025-04-28T15:21:23.947791Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================\n# Dataset to load embeddings with AttentionCombiner\n# ============================\nfrom copy import deepcopy\n\nclass RNA3D_Dataset_WithCache(Dataset):\n    def __init__(self, indices, data, embeddings_cache):\n        self.indices = indices\n        self.data = data\n        self.tokens = {nt: i for i, nt in enumerate('ACGU')}\n        self.embeddings_cache = embeddings_cache\n\n    def __len__(self):\n        return len(self.indices)\n\n    def __getitem__(self, idx):\n        idx = self.indices[idx]\n        sequence_str = self.data['sequence'][idx]\n        sequence = torch.tensor([self.tokens.get(nt, 0) for nt in sequence_str], dtype=torch.long).unsqueeze(0)\n        xyz = torch.tensor(self.data['xyz'][idx], dtype=torch.float32)\n        target_id = train_sequences.iloc[idx]['target_id']\n\n        embeddings = deepcopy(self.embeddings_cache.get(target_id, None))\n        if embeddings is None or \"rho_emb\" not in embeddings or \"ribo_emb\" not in embeddings:\n            raise ValueError(f\"Embeddings for {target_id} not found!\")\n\n        L = sequence.shape[1]\n        max_len = config_data['max_len']\n        if L > max_len:\n            start = np.random.randint(L - max_len)\n            end = start + max_len\n            sequence = sequence[:, start:end]\n            xyz = xyz[start:end]\n\n            # crop embeddings to same region\n            for k in embeddings['rho_emb']:\n                v = embeddings['rho_emb'][k]\n                if isinstance(v, torch.Tensor) and v.dim() == 2 and v.shape[0] == L:\n                    embeddings['rho_emb'][k] = v[start:end]\n            for k in embeddings['ribo_emb']:\n                v = embeddings['ribo_emb'][k]\n                if isinstance(v, torch.Tensor) and v.dim() == 2 and v.shape[0] == L:\n                    embeddings['ribo_emb'][k] = v[start:end]\n\n        return {\n            'sequence': sequence,\n            'xyz': xyz,\n            'target_id': target_id,\n            'embeddings': embeddings\n        }\n\n\n# ============================\n# Helper to extract consistent [L, C] shaped embeddings\n# ============================\ndef extract_valid_rho_embs(rho_emb_dict, length):\n    return [\n        (v[0] if isinstance(v, list) else v)\n        for v in rho_emb_dict.values()\n        if isinstance((v[0] if isinstance(v, list) else v), torch.Tensor)\n        and (v[0] if isinstance(v, list) else v).dim() == 2\n        and (v[0] if isinstance(v, list) else v).shape[0] == length\n    ]\n\n# ============================\n# Create dataloader\n# ============================\ntrain_dataset = RNA3D_Dataset_WithCache(train_index, data, embeddings_cache)\ntrain_loader = DataLoader(train_dataset, batch_size=1, shuffle=True)\n\n# ============================\n# Load Models\n# ============================\nmodel_ribo = finetuned_RibonanzaNet(load_config_from_yaml(\"/kaggle/input/ribonanzanet2d-final/configs/pairwise.yaml\"), pretrained=False).cuda()\nmodel_ribo.load_state_dict(torch.load(\"/kaggle/input/ribonanzanet-3d-finetune/RibonanzaNet-3D.pt\"), strict=False)\nmodel_ribo.eval()\n\nFIXED_KEYS_RHO =RHOFOLD_HOOK_TARGETS\nFIXED_KEYS_RIBO = RIBONANZANET_HOOK_TARGETS\n\nexample = next(iter(embeddings_cache.values()))\n#input_dims_rho = [example['rho_emb'][k].shape[-1] for k in FIXED_KEYS_RHO if k in example['rho_emb']]\n#input_dims_ribo = [example['ribo_emb'][k].shape[-1] for k in FIXED_KEYS_RIBO if k in example['ribo_emb']]\n\nexample = next(iter(embeddings_cache.values()))\nprint(\"📦 FIXED_KEYS_RHO:\", FIXED_KEYS_RHO)\nprint(\"📦 actual keys in example['rho_emb']:\", list(example['rho_emb'].keys()))\ninput_dims_rho = []\nfor k in FIXED_KEYS_RHO:\n    if k not in example['rho_emb']:\n        print(f\"⚠️ Missing RHO key: {k}\")\n    else:\n        print(f\"✅ Found RHO key: {k}, shape: {example['rho_emb'][k].shape}\")\n        input_dims_rho.append(example['rho_emb'][k].shape[-1])\n\ninput_dims_ribo = []\nfor k in FIXED_KEYS_RIBO:\n    if k not in example['ribo_emb']:\n        print(f\"⚠️ Missing RIBO key: {k}\")\n    else:\n        print(f\"✅ Found RIBO key: {k}, shape: {example['ribo_emb'][k].shape}\")\n        input_dims_ribo.append(example['ribo_emb'][k].shape[-1])\n\n\n#fusion_module = FusionLayer().cuda().train()\nfusion_module = GatedFusionLayer().cuda().train()\n\n\"\"\"# 注意：input_dims は一例（最初のエンベディングを用いる）\nexample_embs = next(iter(embeddings_cache.values()))['rho_emb']\nexample_len = list(embeddings_cache.values())[0]['ribo_emb'][\"transformer_encoder.3.self_attn.fc\"].shape[0]\ninput_dims = [\n    v.shape[-1] for v in extract_valid_rho_embs(example_embs, example_len)\n]\n#attention_combiner = AttentionCombiner(input_dims=input_dims).cuda()\n\"\"\"\nattention_combiner = TransformerCombiner(input_dims=input_dims).cuda()\n\n# ============================\n# Optimizer and Scheduler\n# ============================\noptimizer = torch.optim.Adam(\n    list(fusion_module.parameters()) + list(attention_combiner.parameters()),\n    lr=1e-4\n)\n\nscaler = GradScaler()\n\nepochs = 50\ncos_epoch = 30\nsteps_per_epoch = len(train_loader)\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=(epochs - cos_epoch) * steps_per_epoch)\n\n# ==========================\n# Training loop with dual attention fusion (RhoFold + RibonanzaNet)\n# ==========================\n\n# 動的に次元を取得（両方）\nsample_entry = next(iter(embeddings_cache.values()))\nsample_rho_emb = [\n    (v[0] if isinstance(v, list) else v)\n    for v in sample_entry['rho_emb'].values()\n    if (v[0] if isinstance(v, list) else v).dim() == 2\n]\nsample_ribo_emb = [\n    (v[0] if isinstance(v, list) else v)\n    for v in sample_entry['ribo_emb'].values()\n    if (v[0] if isinstance(v, list) else v).dim() == 2\n]\n\n#input_dims_rho = [t.shape[-1] for t in sample_rho_emb]\n#input_dims_ribo = [t.shape[-1] for t in sample_ribo_emb]\n\n\n\n# input_dims 計算時も完全にこれに従う\nexample = next(iter(embeddings_cache.values()))\ninput_dims_rho = [example['rho_emb'][k].shape[-1] for k in FIXED_KEYS_RHO if k in example['rho_emb']]\ninput_dims_ribo = [example['ribo_emb'][k].shape[-1] for k in FIXED_KEYS_RIBO if k in example['ribo_emb']]\n\n# モジュール初期化\n#attention_rho = AttentionCombiner(input_dims=input_dims_rho).cuda()\n#attention_ribo = AttentionCombiner(input_dims=input_dims_ribo).cuda()\n\n\nattention_rho = TransformerCombiner(input_dims=input_dims_rho).cuda()\nattention_ribo = TransformerCombiner(input_dims=input_dims_ribo).cuda()\n#fusion_module = FusionLayer().cuda().train()\nfusion_module = GatedFusionLayer().cuda().train()\n\noptimizer = torch.optim.Adam(\n    list(fusion_module.parameters()) + \n    list(attention_rho.parameters()) + \n    list(attention_ribo.parameters()),\n    lr=1e-4\n)\n\n\n# 学習ループ\nfor epoch in range(epochs):\n    model_ribo.eval()\n    fusion_module.train()\n    tbar = tqdm(train_loader, desc=f\"Epoch {epoch+1}/{epochs}\")\n    total_loss = 0\n    oom = 0\n\n    for idx, batch in enumerate(tbar):\n        try:\n            sequence = batch['sequence'].squeeze(0).cuda()\n            gt_xyz = batch['xyz'].squeeze(0).cuda()\n            if torch.isnan(gt_xyz).any():\n                print(f\"⚠️ NaN detected in gt_xyz at batch {idx}, skipping\")\n                continue\n\n            embeddings = batch['embeddings']\n            if embeddings is None:\n                continue\n\n            def _get_valid_embeddings(emb_dict, fixed_keys, gt_xyz):\n                valid_list = []\n                for k in fixed_keys:\n                    if k not in emb_dict:\n                        continue\n                    v = emb_dict[k]\n                    v = v[0] if isinstance(v, list) else v\n                    if v.dim() >= 2 and v.shape[0] == gt_xyz.shape[0]:\n                        valid_list.append(_sanitize_layer_output(v))\n                return valid_list\n            \n            rho_list = []\n            for k in FIXED_KEYS_RHO:\n                v = embeddings['rho_emb'][k]\n                if isinstance(v, list):\n                    v = v[0]\n                v = _sanitize_layer_output(v)\n                assert v.shape[0] == gt_xyz.shape[0], f\"[rho] shape mismatch: {k} → {v.shape[0]} vs {gt_xyz.shape[0]}\"\n                rho_list.append(v)\n            \n            ribo_list = []\n            for k in FIXED_KEYS_RIBO:\n                v = embeddings['ribo_emb'][k]\n                if isinstance(v, list):\n                    v = v[0]\n                v = _sanitize_layer_output(v)\n                assert v.shape[0] == gt_xyz.shape[0], f\"[ribo] shape mismatch: {k} → {v.shape[0]} vs {gt_xyz.shape[0]}\"\n                ribo_list.append(v)\n\n            assert all([v.shape[0] == gt_xyz.shape[0] for v in rho_list]), \"Rho shape mismatch\"\n            assert all([v.shape[0] == gt_xyz.shape[0] for v in ribo_list]), \"Ribo shape mismatch\"\n            if epoch == 0 and idx < 5:\n                print(f\"\\n🧪 Batch {idx}: shape inspection\")\n                print(\"📏 gt_xyz.shape:\", gt_xyz.shape)\n            \n                for k, v in embeddings['rho_emb'].items():\n                    v_ = v[0] if isinstance(v, list) else v\n                    v_ = _sanitize_layer_output(v_)  # ✅ ここが必須！！\n                    print(f\"[rho] {k} → shape: {v_.shape}, valid={v_.dim() == 2 and v_.shape[0] == gt_xyz.shape[0]}\")\n\n                    \n                for k, v in embeddings['ribo_emb'].items():\n                    v_ = v[0] if isinstance(v, list) else v\n                    v_ = _sanitize_layer_output(v_)\n                    print(f\"[ribo] {k} → shape: {v_.shape}, valid={v_.dim() == 2 and v_.shape[0] == gt_xyz.shape[0]}\")\n\n            if len(rho_list) != len(input_dims_rho) or len(ribo_list) != len(input_dims_ribo):\n                print(\"rholist_len\",len(rho_list))\n                print(\"rho_imput_dims\",len(input_dims_rho))\n                print(\"ribo_list\",len(ribo_list))\n                print(\"ribo_imp_dims\",len(input_dims_ribo))\n                print(f\"⚠️ Layer count mismatch at batch {idx}, skipping\")\n                continue\n            print(\"🧪 calling attention_rho\")\n            print(f\"[DEBUG] attention_rho = {type(attention_rho)}\")\n            rho_emb = attention_rho(rho_list).cuda()\n            ribo_emb = attention_ribo(ribo_list).cuda()\n            print(\"🧪 rho_emb shape:\", rho_emb.shape)\n            print(\"🧪 ribo_emb shape:\", ribo_emb.shape)\n            print(\"🧪 gt_xyz shape:\", gt_xyz.shape)\n            print(f\"🔍 Batch {idx}: rho_list shapes:\")\n            for i, x in enumerate(rho_list):\n                print(f\"  rho_list[{i}].shape = {x.shape}\")\n            \n            fused_embedding = fusion_module(rho_emb, ribo_emb)\n            #pred_xyz = model_ribo.xyz_predictor(fused_embedding).squeeze(0)\n            pred_xyz = model_ribo.decoder_head(fused_embedding).squeeze(0)\n\n\n            if epoch == 0 and idx == 0:\n                print(\"📏 pred_xyz.shape:\", pred_xyz.shape)\n                print(\"📏 gt_xyz.shape:\", gt_xyz.shape)\n\n            loss = dRMAE(pred_xyz, pred_xyz, gt_xyz, gt_xyz) + align_svd_mae(pred_xyz, gt_xyz)\n            #loss = 0.5 * dRMAE(pred_xyz, pred_xyz, gt_xyz, gt_xyz)+0.5 * align_svd_mae(pred_xyz, gt_xyz) +0.01 * entropy_regularization(attention_combiner)\n\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(fusion_module.parameters(), 1.0)\n            optimizer.step()\n            optimizer.zero_grad()\n\n            if (epoch + 1) > cos_epoch:\n                scheduler.step()\n\n            total_loss += loss.item()\n            tbar.set_postfix(loss=total_loss / (tbar.n + 1), OOM=oom)\n\n        except RuntimeError as e:\n            print(f\"❌ OOM or Error at batch {tbar.n}: {e}\")\n            oom += 1\n            torch.cuda.empty_cache()\n            import gc; gc.collect()\n            continue\n\n    print(f\"✅ Epoch {epoch+1} completed. Avg Loss: {total_loss / len(tbar):.6f}\")\n\n# 保存\ntorch.save(fusion_module.state_dict(), \"fusion_module_finetuned.pt\")\ntorch.save(attention_rho.state_dict(), \"attention_rho.pt\")\ntorch.save(attention_ribo.state_dict(), \"attention_ribo.pt\")\n\nprint(\"✅ 最終 pred_xyz の shape:\", pred_xyz.shape)\nprint(\"🧾 pred_xyz の中身（先頭5個）:\", pred_xyz[:5].detach().cpu().numpy())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T16:02:25.061992Z","iopub.execute_input":"2025-04-28T16:02:25.062314Z","iopub.status.idle":"2025-04-28T16:05:57.452235Z","shell.execute_reply.started":"2025-04-28T16:02:25.062290Z","shell.execute_reply":"2025-04-28T16:05:57.451067Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"def get_final_embeddings(fasta_path, model):\n    all_layer_outputs.clear()\n    register_hooks(model, RHOFOLD_HOOK_TARGETS)\n\n    data_dict = get_features(fasta_path, fasta_path)\n\n    with torch.no_grad():\n        _ = model(\n            tokens=data_dict['tokens'].to(\"cuda\"),\n            rna_fm_tokens=data_dict['rna_fm_tokens'].to(\"cuda\"),\n            seq=data_dict['seq']\n        )\n\n    return {\n        k: (v[0] if isinstance(v, list) else v).cpu()\n        for k, v in all_layer_outputs.items()\n        if isinstance(v, (torch.Tensor, list))\n    }\n### def precompute_embeddings(model_ribo, model_rho):\n    embeddings_cache = {}\n    tokens = {nt: i for i, nt in enumerate('ACGU')}\n\n    for idx in tqdm(train_index):\n        target_id = train_sequences.iloc[idx]['target_id']\n        sequence_tensor = torch.tensor([\n            tokens.get(nt, 0) for nt in data['sequence'][idx]\n        ], dtype=torch.long).unsqueeze(0)\n\n        if target_id in embeddings_cache:\n            continue\n\n        fasta_path = f\"/kaggle/input/stanford-rna-3d-folding/MSA/{target_id}.MSA.fasta\"\n        if not os.path.exists(fasta_path):\n            continue\n\n        try:\n            all_layer_outputs.clear()\n            register_hooks(model_rho, RHOFOLD_HOOK_TARGETS)\n            rho_embs = get_final_embeddings(fasta_path, model_rho=model_rho)\n            rho_valid = {\n                k: (v[0] if isinstance(v, list) else v).cpu()\n                for k, v in rho_embs.items()\n                if isinstance((v[0] if isinstance(v, list) else v), torch.Tensor)\n            }\n\n            all_layer_outputs.clear()\n            register_hooks(model_ribo, RIBONANZANET_HOOK_TARGETS)\n            with torch.no_grad():\n                _ = model_ribo(sequence_tensor.cuda())\n            ribo_valid = {\n                k: (v[0] if isinstance(v, list) else v).cpu()\n                for k, v in all_layer_outputs.items()\n                if k in RIBONANZANET_HOOK_TARGETS and isinstance((v[0] if isinstance(v, list) else v), torch.Tensor)\n            }\n\n            # Reshape features\n            reshaper_rho = FeatureReshaper(input_shapes={k: v.shape for k, v in rho_valid.items()}).cuda()\n            reshaper_ribo = FeatureReshaper(input_shapes={k: v.shape for k, v in ribo_valid.items()}).cuda()\n\n            rho_reshaped = reshaper_rho({k: v.cuda() for k, v in rho_valid.items()})\n            ribo_reshaped = reshaper_ribo({k: v.cuda() for k, v in ribo_valid.items()})\n\n            # 結果保存\n            embeddings_cache[target_id] = {\n                'rho_emb': rho_reshaped,\n                'ribo_emb': ribo_reshaped\n            }\n\n            torch.cuda.empty_cache()\n\n        except Exception as e:\n            print(f\"[\\u274c Error] {target_id}: {e}\")\n            continue\n\n    return embeddings_cache\nInference & Visualization","metadata":{}},{"cell_type":"code","source":"# ===============================\n# Prediction with dual AttentionCombiner + FusionModule\n# ===============================\n\n#fusion_module = FusionLayer().to(\"cuda\")\nfusion_module = GatedFusionLayer().to(\"cuda\")\nfusion_ckpt_path = \"/kaggle/working/fusion_module_finetuned.pt\"\nfusion_module.load_state_dict(torch.load(fusion_ckpt_path, map_location=\"cuda\"))\nfusion_module.eval()\n\n# 🔄 AttentionCombiner (Rho & Ribo) を input_dims ベースで再構築\nsample_entry = next(iter(embeddings_cache.values()))\nsample_rho_emb = [\n    (v[0] if isinstance(v, list) else v)\n    for v in sample_entry['rho_emb'].values()\n    if (v[0] if isinstance(v, list) else v).dim() == 2\n]\nsample_ribo_emb = [\n    (v[0] if isinstance(v, list) else v)\n    for v in sample_entry['ribo_emb'].values()\n    if (v[0] if isinstance(v, list) else v).dim() == 2\n]\n\ninput_dims_rho = [t.shape[-1] for t in sample_rho_emb]\ninput_dims_ribo = [t.shape[-1] for t in sample_ribo_emb]\n\n#attention_rho = AttentionCombiner(input_dims=input_dims_rho).to(\"cuda\")\n#attention_ribo = AttentionCombiner(input_dims=input_dims_ribo).to(\"cuda\")\nattention_rho = TransformerCombiner(input_dims=input_dims_rho).to(\"cuda\")\nattention_ribo = TransformerCombiner(input_dims=input_dims_ribo).to(\"cuda\")\n\nrho_ckpt = \"/kaggle/working/attention_rho.pt\"\nribo_ckpt = \"/kaggle/working/attention_ribo.pt\"\nif os.path.exists(rho_ckpt):\n    attention_rho.load_state_dict(torch.load(rho_ckpt, map_location=\"cuda\"))\nif os.path.exists(ribo_ckpt):\n    attention_ribo.load_state_dict(torch.load(ribo_ckpt, map_location=\"cuda\"))\nattention_rho.eval()\nattention_ribo.eval()\n\n# ===============================\n# 推論関数\n# ===============================\ndef tokenize_sequence(input_sequence):\n    tokens = {nt: i for i, nt in enumerate('ACGU')}\n    sequence = [tokens.get(nt, 0) for nt in input_sequence]\n    return torch.tensor(sequence).unsqueeze(0)\n\ndef process_dataset(sequence_dict, output_dir=\"/kaggle/working/predictions\"):\n    os.makedirs(output_dir, exist_ok=True)\n\n    for i, (name, sequence) in enumerate(sequence_dict.items()):\n        print(f\"\\n🧬 [{i+1}/{len(sequence_dict)}] Processing {name}\")\n        try:\n            fasta_path = f\"/kaggle/input/stanford-rna-3d-folding/MSA/{name}.MSA.fasta\"\n            if not os.path.exists(fasta_path):\n                print(f\"❌ {fasta_path} not found. Skipping.\")\n                continue\n\n            # --- RhoFold embeddings ---\n            rho_embs = get_final_embeddings(fasta_path, model=model_rho)\n            rho_layers = [\n                _sanitize_layer_output(v[0] if isinstance(v, list) else v)\n                for v in rho_embs.values()\n                if (v[0] if isinstance(v, list) else v).dim() >= 2\n            ]\n\n            if len(rho_layers) != len(input_dims_rho):\n                print(f\"⚠️ Rho Layer count mismatch in {name}, skipping.\")\n                continue\n\n            rho_emb_combined = attention_rho(rho_layers).to(\"cuda\")\n\n            # --- Ribonanza embeddings ---\n            input_tensor = tokenize_sequence(sequence).to(\"cuda\")\n            with torch.no_grad():\n                ribo_embs = model_ribo.get_finetuned_embeddings(input_tensor)\n            ribo_layers = [\n                _sanitize_layer_output(v[0] if isinstance(v, list) else v)\n                for v in ribo_embs['ribo_emb'].values()\n                if (v[0] if isinstance(v, list) else v).dim() >= 2\n            ]\n\n            if len(ribo_layers) != len(input_dims_ribo):\n                print(f\"⚠️ Ribo Layer count mismatch in {name}, skipping.\")\n                continue\n\n            ribo_emb_combined = attention_ribo(ribo_layers).to(\"cuda\")\n\n            # --- fusion & prediction ---\n            with torch.no_grad():\n                fused_embedding = fusion_module(rho_emb_combined, ribo_emb_combined)\n                xyz = model_ribo.xyz_predictor(fused_embedding).squeeze(0).cpu().numpy()\n\n            df_xyz = pd.DataFrame(xyz, columns=[\"x_1\", \"y_1\", \"z_1\"])\n            csv_path = os.path.join(output_dir, f\"{name}_predicted.csv\")\n            df_xyz.to_csv(csv_path, index=False)\n            print(f\"✅ Saved: {csv_path}\")\n\n        except RuntimeError as e:\n            print(f\"❌ RuntimeError in {name}: {e}\")\n            torch.cuda.empty_cache()\n            import gc; gc.collect()\n            continue\n\n        del rho_embs, rho_emb_combined, input_tensor, ribo_embs, ribo_emb_combined, fused_embedding, xyz\n        torch.cuda.empty_cache()\n        import gc; gc.collect()\n\n# ===============================\n# 推論対象のロード\n# ===============================\ndf_test = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/test_sequences.csv\")\nsequence_dict = dict(zip(df_test[\"target_id\"], df_test[\"sequence\"]))\nprocess_dataset(sequence_dict)\n\n# ===============================\n# R1138 特別処理 (OOM回避用)\n# ===============================\ninput_csv = \"/kaggle/input/ribonanzanet-3d-inference/submission.csv\"\noutput_csv = \"/kaggle/working/predictions/R1138_predicted.csv\"\ndf_r1138 = pd.read_csv(input_csv)\ndf_r1138 = df_r1138[df_r1138[\"ID\"].str.startswith(\"R1138\")][[\"x_1\", \"y_1\", \"z_1\"]]\ndf_r1138.to_csv(output_csv, index=False)\nprint(f\"✅ R1138 のデータを {output_csv} に保存しました。\")\n\n# ===============================\n# 提出用形式への変換（全予測CSVを1つにマージ）\n# ===============================\nimport glob\nimport re\n\nrows = []\ninput_dir = \"/kaggle/working/predictions\"\n\nfor csv_path in sorted(glob.glob(os.path.join(input_dir, \"*_predicted.csv\"))):\n    name = os.path.basename(csv_path).replace(\"_predicted.csv\", \"\")\n    df = pd.read_csv(csv_path)\n    sequence = sequence_dict.get(name, \"\")\n    for i in range(len(df)):\n        rows.append({\n            \"ID\": f\"{name}_{i+1}\",\n            \"resname\": sequence[i] if i < len(sequence) else \"N\",\n            \"resid\": i + 1,\n            \"x_1\": df.loc[i, \"x_1\"],\n            \"y_1\": df.loc[i, \"y_1\"],\n            \"z_1\": df.loc[i, \"z_1\"],\n        })\n\ndf_all = pd.DataFrame(rows)\ndf_all['sort_key'] = df_all['ID'].apply(lambda x: tuple(map(int, re.findall(r'\\d+', x))))\ndf_all = df_all.sort_values(by='sort_key').drop(columns='sort_key')\ndf_all.reset_index(drop=True, inplace=True)\ndf_all.to_csv(\"/kaggle/working/fused_predictions_all.csv\", index=False)\nprint(\"✅ 全予測を fused_predictions_all.csv に保存しました！\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T09:18:14.161770Z","iopub.execute_input":"2025-04-28T09:18:14.162075Z","iopub.status.idle":"2025-04-28T09:18:14.202187Z","shell.execute_reply.started":"2025-04-28T09:18:14.162049Z","shell.execute_reply":"2025-04-28T09:18:14.201042Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\n\nfused_path = \"/kaggle/working/fused_predictions_all.csv\"\nsample_path = \"/kaggle/input/ribonanzanet-3d-inference/submission.csv\"\noutput_path = \"/kaggle/working/submission_filled.csv\"\n\nnum_copies = 5\n\n\ndf_fused = pd.read_csv(fused_path)\ndf_sample = pd.read_csv(sample_path)\n\nfor i in range(1, num_copies + 1):\n    df_fused[f\"x_{i}\"] = df_fused[\"x_1\"]\n    df_fused[f\"y_{i}\"] = df_fused[\"y_1\"]\n    df_fused[f\"z_{i}\"] = df_fused[\"z_1\"]\n\n\nreplace_cols = [f\"{axis}_{i}\" for i in range(1, num_copies + 1) for axis in ['x', 'y', 'z']]\n\n\ndf_merged = df_sample.drop(columns=replace_cols).merge(\n    df_fused[[\"ID\"] + replace_cols],\n    on=\"ID\",\n    how=\"left\"\n)\n\n\ndf_merged.to_csv(output_path, index=False)\nprint(f\"✅ 置換後の submission を {output_path} に保存しました。\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T08:56:31.670060Z","iopub.status.idle":"2025-04-28T08:56:31.670470Z","shell.execute_reply":"2025-04-28T08:56:31.670297Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"try:\n    import Bio\nexcept:\n    #for rhofold+ #####################\n    !pip install biopython\n    !pip install ml-collections\n    !pip install python-box\n    !pip install dm-tree\n    !pip install openmm[cuda12]\n\nfrom copy import deepcopy\n\nimport pandas as pd\nfrom Bio.PDB import Atom, Model, Chain, Residue, Structure, PDBParser\nfrom Bio import SeqIO\nimport os, sys\nimport re\nimport numpy as np\nimport pandas as pd\n\nimport matplotlib\nimport matplotlib.pyplot as plt\n!pip install kagglehub\nimport kagglehub\n\nprint('IMPORT OK !!!!')\nPYTHON = sys.executable\nprint('PYTHON',PYTHON)\n\nusalign_path = kagglehub.dataset_download('metric/usalign')\nUSALIGN = \\\n'/kaggle/working//USalign'\n#'<your us align path>/USalign'\n\nos.system('cp /kaggle/input/usalign/USalign /kaggle/working/')\nos.system(' chmod u+x /kaggle/working//USalign')\n\n\nsubmission = df_merged\nLABEL_DF = pd.read_csv('/kaggle/input/stanford-rna-3d-folding/train_labels.csv')\n\nLABEL_DF[\"pdb_id\"] = LABEL_DF[\"ID\"].apply(lambda x: x.split(\"_\")[0]+'_'+x.split(\"_\")[1])\nsubmission['submission_id'] = submission['ID'].str.split('_').str[0]","metadata":{"trusted":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2025-04-28T08:56:31.671581Z","iopub.status.idle":"2025-04-28T08:56:31.671919Z","shell.execute_reply":"2025-04-28T08:56:31.671752Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sys.path.append('/kaggle/usr/lib/stanfordrna_submission_visualizer_ribonanza')\nfrom stanfordrna_submission_visualizer_ribonanza import submission_to_visual_2","metadata":{"trusted":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2025-04-28T08:56:31.672634Z","iopub.status.idle":"2025-04-28T08:56:31.672940Z","shell.execute_reply":"2025-04-28T08:56:31.672832Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submission_to_visual_2(submission)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T08:56:31.673867Z","iopub.status.idle":"2025-04-28T08:56:31.674200Z","shell.execute_reply":"2025-04-28T08:56:31.674044Z"}},"outputs":[],"execution_count":null}]}