{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":87793,"databundleVersionId":12024591,"sourceType":"competition"},{"sourceId":311741,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":264400,"modelId":285488}],"dockerImageVersionId":31011,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input.'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport torch\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport random\nimport pickle\nimport os\nimport sys\nfrom tqdm import tqdm\nimport math\nimport yaml\nfrom collections import defaultdict, Counter\nfrom torch.utils.data import Dataset, DataLoader\nimport torch.nn as nn\nimport torch.utils.checkpoint as checkpoint ","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"config = {\n    \"seed\": 0,\n    \"cutoff_date\": \"2020-01-01\",\n    \"test_cutoff_date\": \"2022-05-01\",\n    \"max_len\": 256, # Max sequence length for cropping\n    \"batch_size\": 1, # Keep at 1 due to potential memory constraints\n    \"learning_rate\": 1e-4, # Original: 2e-4, but 1e-4 is also common\n    \"weight_decay\": 0.0,\n    \"mixed_precision\": \"bf16\", # \"bf16\" or \"fp16\" or None\n    \"model_config_path\": \"/kaggle/input/ribonanzanet2/pytorch/alpha/1/pairwise.yaml\",\n    \"epochs\": 1, # SET THIS TO A HIGHER VALUE FOR ACTUAL TRAINING (e.g., 10, 20, 50)\n    \"cos_epoch\": 0, # Start cosine annealing after this many epochs\n    \"loss_power_scale\": 1.0,\n    \"max_cycles\": 1,\n    \"grad_clip\": 0.1,\n    \"gradient_accumulation_steps\": 1, # Increase if batch_size=1 causes OOM with grad_clip\n    \"d_clamp\": 30,\n    \"max_len_filter\": 9999999, # Max length of sequence to consider from dataset\n    \"min_len_filter\": 10,     # Min length of sequence to consider\n    \"structural_violation_epoch\": 50, # Epoch to start considering structural violation loss (if implemented)\n    \"balance_weight\": False,\n    \"n_times\": 1000, # Number of diffusion timesteps\n    \"msa_base_path\": \"/kaggle/input/stanford-rna-3d-folding/MSA\", # Path to MSA files\n}\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"torch.manual_seed(config['seed'])\nnp.random.seed(config['seed'])\nrandom.seed(config['seed'])\nif torch.cuda.is_available():\n    torch.cuda.manual_seed_all(config['seed'])","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"--- Loading and Preprocessing Competition Data ---\")\n\ntrain_sequences_df=pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/train_sequences.csv\")\ntrain_labels_df=pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/train_labels.csv\")\ntrain_labels_df[\"pdb_id\"] = train_labels_df[\"ID\"].apply(lambda x: x.split(\"_\")[0] + '_' + x.split(\"_\")[1])\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nall_xyz_coords = []\nfor pdb_id in tqdm(train_sequences_df['target_id'], desc=\"Processing PDB IDs\"):\n    df_subset = train_labels_df[train_labels_df[\"pdb_id\"] == pdb_id]\n    xyz = df_subset[['x_1', 'y_1', 'z_1']].to_numpy().astype('float32')\n    xyz[xyz < -1e17] = float('NaN')  # Replace placeholder with NaN\n    all_xyz_coords.append(xyz)\n\n# Filter data\nfilter_mask = []\nmax_seq_len_found = 0\nfor i, xyz in enumerate(all_xyz_coords):\n    seq_len = len(train_sequences_df['sequence'][i])\n    if seq_len > max_seq_len_found:\n        max_seq_len_found = seq_len\n    \n    filter_mask.append(\n        (np.sum(np.isnan(xyz)) / xyz.size <= 0.5) and \\\n        (seq_len < config['max_len_filter']) and \\\n        (seq_len > config['min_len_filter'])\n    )\n\nprint(f\"Longest sequence in raw train data: {max_seq_len_found}\")\n\nfilter_mask = np.array(filter_mask)\nfiltered_indices = np.arange(len(filter_mask))[filter_mask]\n\ntrain_sequences_filtered_df = train_sequences_df.loc[filtered_indices].reset_index(drop=True)\nall_xyz_filtered = [all_xyz_coords[i] for i in filtered_indices]\n\ndata_dict = {\n    \"sequence\": train_sequences_filtered_df['sequence'].to_list(),\n    \"target_id\": train_sequences_filtered_df['target_id'].to_list(), # Added for MSA lookup\n    \"temporal_cutoff\": train_sequences_filtered_df['temporal_cutoff'].to_list(),\n    \"description\": train_sequences_filtered_df['description'].to_list(),\n    \"all_sequences\": train_sequences_filtered_df['all_sequences'].to_list(),\n    \"xyz\": all_xyz_filtered\n}\nprint(f\"Number of sequences after filtering: {len(data_dict['sequence'])}\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"--- Performing Temporal Split ---\")\nall_indices_for_split = np.arange(len(data_dict['sequence']))\ncutoff_timestamp = pd.Timestamp(config['cutoff_date'])\ntest_cutoff_timestamp = pd.Timestamp(config['test_cutoff_date'])\n\ntrain_indices_split = [i for i, date_str in enumerate(data_dict['temporal_cutoff']) \n                       if pd.Timestamp(date_str) <= cutoff_timestamp]\nval_indices_split = [i for i, date_str in enumerate(data_dict['temporal_cutoff']) \n                     if pd.Timestamp(date_str) > cutoff_timestamp and pd.Timestamp(date_str) <= test_cutoff_timestamp]\n\nprint(f\"Train size after split: {len(train_indices_split)}\")\nprint(f\"Validation size after split: {len(val_indices_split)}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"--- Defining Dataset Class with MSA ---\")\ndef get_msa_profile(target_id, msa_base_path_arg): # Renamed arg for clarity\n    \"\"\"\n    Parses the MSA file for a given target_id and computes the frequency profile.\n\n    Args:\n        target_id (str): The target identifier (e.g., \"1SCL_A\").\n        msa_base_path_arg (str): The base directory where MSA files are stored.\n\n    Returns:\n        np.ndarray or None: The frequency profile (alignment_length, 5)\n                             for (A, C, G, U, -), or None if MSA not found or empty.\n    \"\"\"\n    msa_file_path = os.path.join(msa_base_path_arg, f\"{target_id}.MSA.fasta\")\n    \n    aligned_sequences = []\n    current_sequence = \"\"\n\n    if not os.path.exists(msa_file_path):\n        # print(f\"Warning: MSA file not found for {target_id} at {msa_file_path}\")\n        return None \n\n    with open(msa_file_path, 'r') as f:\n        for line in f:\n            line = line.strip()\n            if not line:\n                continue\n            if line.startswith(\">\"):\n                if current_sequence:\n                    aligned_sequences.append(current_sequence)\n                current_sequence = \"\"\n            else:\n                current_sequence += line\n        if current_sequence:\n            aligned_sequences.append(current_sequence)\n\n    if not aligned_sequences:\n        # print(f\"Warning: No sequences found in MSA file for {target_id}\")\n        return None\n\n    num_msa_sequences = len(aligned_sequences)\n    alignment_length = len(aligned_sequences[0])\n    \n    if not all(len(s) == alignment_length for s in aligned_sequences):\n        # print(f\"Warning: MSA sequences for {target_id} have inconsistent lengths.\")\n        return None \n\n    alphabet = ['A', 'C', 'G', 'U', '-']\n    profile_counts = np.zeros((alignment_length, len(alphabet)), dtype=np.int32)\n\n    for pos_idx in range(alignment_length):\n        col_chars = [seq[pos_idx].upper() for seq in aligned_sequences]\n        counts = Counter(col_chars)\n        for char_idx, char_alpha in enumerate(alphabet):\n            profile_counts[pos_idx, char_idx] = counts.get(char_alpha, 0)\n\n    if num_msa_sequences > 0:\n        frequency_profile = profile_counts.astype(np.float32) / num_msa_sequences\n    else:\n        frequency_profile = profile_counts.astype(np.float32) \n\n    return frequency_profile\n\n\nclass RNA3D_Dataset_MSA(Dataset):\n    def __init__(self, indices, data_dict_input, current_config):\n        self.indices = indices\n        self.data = data_dict_input\n        # Store specific config values needed as direct attributes\n        self.cfg_max_len = current_config['max_len']\n        self.cfg_msa_base_path = current_config['msa_base_path'] # Explicitly store this\n        \n        self.tokens = defaultdict(lambda: 4) \n        self.tokens.update({'A':0,'C':1,'G':2,'U':3})\n        if 'target_id' not in self.data: \n            raise ValueError(\"Data dictionary must contain 'target_id' list.\")\n\n    def __len__(self): \n        return len(self.indices)\n    \n    def __getitem__(self, idx):\n        actual_idx = self.indices[idx] \n        sequence_str = self.data['sequence'][actual_idx]\n        tokenized_sequence = torch.tensor([self.tokens[nt] for nt in sequence_str], dtype=torch.long)\n        xyz_coords = torch.tensor(np.array(self.data['xyz'][actual_idx]), dtype=torch.float32)\n        current_target_id = self.data['target_id'][actual_idx]\n        \n        # Use the explicitly stored msa_base_path from __init__\n        msa_profile = get_msa_profile(current_target_id, self.cfg_msa_base_path) \n        \n        if msa_profile is None:\n            # Fallback: create a one-hot encoding of the sequence if MSA is missing\n            msa_alphabet_size = 5 \n            msa_profile = np.zeros((len(tokenized_sequence), msa_alphabet_size), dtype=np.float32)\n            for i, token_val in enumerate(tokenized_sequence):\n                if token_val < 4: \n                    msa_profile[i, token_val] = 1.0\n        \n        msa_profile_tensor = torch.tensor(msa_profile, dtype=torch.float32)\n\n        seq_len = len(tokenized_sequence)\n        if seq_len > self.cfg_max_len: # Use stored max_len\n            crop_start = np.random.randint(0, seq_len - self.cfg_max_len + 1)\n            crop_end = crop_start + self.cfg_max_len\n            tokenized_sequence = tokenized_sequence[crop_start:crop_end]\n            xyz_coords = xyz_coords[crop_start:crop_end]\n            msa_profile_tensor = msa_profile_tensor[crop_start:crop_end, :]\n        \n        first_valid_idx = -1\n        for i in range(len(xyz_coords)):\n            if (~torch.isnan(xyz_coords[i])).all(): \n                first_valid_idx = i\n                break\n        if first_valid_idx != -1: \n            xyz_coords = xyz_coords - xyz_coords[first_valid_idx]\n        \n        return {\n            'sequence': tokenized_sequence, \n            'xyz': xyz_coords, \n            'msa_profile': msa_profile_tensor, \n            'target_id': current_target_id\n        }","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"--- Creating DataLoaders ---\")\ntrain_dataset = RNA3D_Dataset_MSA(train_indices_split, data_dict, config)\nval_dataset = RNA3D_Dataset_MSA(val_indices_split, data_dict, config)\n\ntrain_loader = DataLoader(train_dataset, batch_size=config['batch_size'], shuffle=True, num_workers=2, pin_memory=True)\nval_loader = DataLoader(val_dataset, batch_size=config['batch_size'], shuffle=False, num_workers=2, pin_memory=True)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sys.path.append(\"/kaggle/input/ribonanzanet2/pytorch/alpha/1\")\n\nfrom Network import *\nimport yaml","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"--- Defining Model Architecture ---\")\n\n# Define the path to the RNet2 model code\nrnet2_code_path = \"/kaggle/input/ribonanzanet2/pytorch/alpha/1\"\n\n# Check if the path and Network.py exist (for debugging)\nif not os.path.exists(rnet2_code_path):\n    print(f\"ERROR: RNet2 code path does not exist: {rnet2_code_path}\")\n    print(\"Please ensure the 'RibonanzaNet2' Kaggle Model is added as input and the path is correct.\")\nelif not os.path.exists(os.path.join(rnet2_code_path, \"Network.py\")):\n    print(f\"ERROR: Network.py not found in {rnet2_code_path}\")\n    print(\"Please check the structure of the added 'RibonanzaNet2' Kaggle Model input.\")\nelse:\n    print(f\"Found RNet2 code path: {rnet2_code_path}\")\n    # Add RNet2 code path to system path\n    if rnet2_code_path not in sys.path: # Avoid adding multiple times\n        sys.path.append(rnet2_code_path)\n\n    try:\n        from Network import RibonanzaNet, MultiHeadAttention # Assuming these are in Network.py\n        print(\"Successfully imported RibonanzaNet and MultiHeadAttention from Network.py\")\n    except ImportError as e:\n        print(f\"ImportError: {e}\")\n        print(\"Could not import from Network.py. Check sys.path and file contents.\")\n        raise # Re-raise the error to stop execution if import fails\n\n# Ensure global_notebook_config is defined (from your earlier cells)\n# This is a placeholder, ensure it's correctly defined and populated in your notebook\nif 'global_notebook_config' not in globals():\n    print(\"ERROR: global_notebook_config is not defined. Please define it in a previous cell.\")\n    # Example definition (make sure paths and values are correct for your setup)\n    global_notebook_config = {\n        \"model_config_path\": \"/kaggle/input/ribonanzanet2/pytorch/alpha/1/pairwise.yaml\",\n        \"n_times\": 1000, \n        # ... other necessary configurations from your notebook's config cell ...\n    }\n    print(\"Defined a placeholder global_notebook_config. PLEASE VERIFY ITS CONTENTS.\")\n\n\nclass SinusoidalPosEmb(nn.Module):\n    def __init__(self, dim):\n        super().__init__()\n        self.dim = dim\n    def forward(self, x):\n        device = x.device\n        half_dim = self.dim // 2\n        emb = math.log(10000) / (half_dim - 1)\n        emb = torch.exp(torch.arange(half_dim, device=device) * -emb)\n        emb = x[:, None] * emb[None, :]\n        emb = torch.cat((emb.sin(), emb.cos()), dim=-1)\n        return emb\n\nclass finetuned_RibonanzaNet_MSA(RibonanzaNet):\n    def __init__(self, model_config_obj, msa_feature_dim=5): \n        model_config_obj.dropout = 0.1\n        model_config_obj.use_grad_checkpoint = True \n        super(finetuned_RibonanzaNet_MSA, self).__init__(model_config_obj)\n        \n        weights_path = \"/kaggle/input/ribonanzanet2/pytorch/alpha/1/pytorch_model_fsdp.bin\"\n        if not os.path.exists(weights_path):\n            print(f\"ERROR: Pre-trained weights not found at {weights_path}\")\n            raise FileNotFoundError(f\"Pre-trained weights not found at {weights_path}\")\n        self.load_state_dict(torch.load(weights_path, map_location='cpu'))\n        print(\"Successfully loaded pre-trained RNet2 weights.\")\n        \n        self.dropout_ft = nn.Dropout(0.0) \n        decoder_dim = 768\n        \n        self.msa_projector = nn.Linear(msa_feature_dim, model_config_obj.ninp // 4) \n        adapted_input_dim = model_config_obj.ninp + (model_config_obj.ninp // 4)\n\n        self.adaptor = nn.Sequential(nn.Linear(adapted_input_dim, decoder_dim), nn.LayerNorm(decoder_dim))\n\n        self.structure_module_stack = nn.ModuleList([\n            SimpleStructureModule(d_model=decoder_dim, nhead=12, dim_feedforward=decoder_dim*4, \n                                  pairwise_dimension=model_config_obj.pairwise_dimension, dropout=0.0) \n            for _ in range(6)\n        ])\n        self.xyz_embedder = nn.Linear(3, decoder_dim)\n        self.xyz_norm = nn.LayerNorm(decoder_dim)\n        self.xyz_predictor = nn.Linear(decoder_dim, 3)\n        self.distogram_predictor = nn.Sequential(nn.LayerNorm(model_config_obj.pairwise_dimension),\n                                                 nn.Linear(model_config_obj.pairwise_dimension, 40))\n        self.time_embedder = SinusoidalPosEmb(decoder_dim)\n        self.time_mlp = nn.Sequential(nn.Linear(decoder_dim, decoder_dim), nn.ReLU(), nn.Linear(decoder_dim, decoder_dim))\n        self.time_norm = nn.LayerNorm(decoder_dim)\n        self.distance2pairwise = nn.Linear(1, model_config_obj.pairwise_dimension, bias=False)\n\n    # Removed custom_checkpoint method, will call torch.utils.checkpoint.checkpoint directly\n\n    def embed_pairwise_distance(self, pairwise_features, xyz): \n        dist_matrix = xyz[:, None, :, :] - xyz[:, :, None, :]\n        dist_matrix = (dist_matrix**2).sum(-1).clamp(min=1e-6, max=37**2).sqrt() \n        dist_matrix = dist_matrix[:, :, :, None]\n        pairwise_features = pairwise_features + self.distance2pairwise(dist_matrix)\n        return pairwise_features\n\n    def forward(self, src_tokens, msa_profile_input, noisy_xyz, timestep_t):\n        padding_mask = torch.ones_like(src_tokens).long().to(src_tokens.device)\n        sequence_features, pairwise_features = self.get_embeddings(src_tokens, padding_mask)\n        \n        projected_msa = self.msa_projector(msa_profile_input)\n        concatenated_sequence_features = torch.cat((sequence_features, projected_msa), dim=-1)\n        \n        adapted_sequence_features = self.adaptor(concatenated_sequence_features)\n        \n        distogram_pred = self.distogram_predictor(pairwise_features)\n\n        current_batch_size = noisy_xyz.shape[0]\n        final_sequence_features = adapted_sequence_features.repeat(current_batch_size, 1, 1)\n        final_pairwise_features_init = pairwise_features.expand(current_batch_size, -1, -1, -1) # Renamed\n\n        # Use a lambda or a locally defined function for checkpointing this specific call\n        def _checkpointed_embed_pairwise_dist_fn(pf, nz): # pf: pairwise_features, nz: noisy_xyz\n            return self.embed_pairwise_distance(pf, nz)\n        \n        final_pairwise_features = checkpoint.checkpoint(\n            _checkpointed_embed_pairwise_dist_fn,\n            final_pairwise_features_init, # Pass initial pairwise features\n            noisy_xyz,\n            use_reentrant=False # PyTorch recommendation for newer versions\n        )\n        \n        time_embedding = self.time_embedder(timestep_t).unsqueeze(1)\n        tgt_representation = self.xyz_norm(final_sequence_features + self.xyz_embedder(noisy_xyz) + time_embedding)\n        tgt_representation = self.time_norm(tgt_representation + self.time_mlp(tgt_representation))\n\n        for layer_module in self.structure_module_stack: # Renamed 'layer' to 'layer_module'\n            # For checkpointing nn.Module instances, it's usually direct\n            # The layer_module.forward expects a tuple as its single argument\n            tgt_representation = checkpoint.checkpoint(\n                layer_module, # Pass the module instance directly\n                (tgt_representation, final_sequence_features, final_pairwise_features, noisy_xyz, None),\n                use_reentrant=False\n            )\n        \n        predicted_noise_or_coords = self.xyz_predictor(tgt_representation)\n        return predicted_noise_or_coords, distogram_pred\n\n    def denoise(self, src_tokens, msa_profile_input, noisy_xyz, timestep_t):\n        padding_mask = torch.ones_like(src_tokens).long().to(src_tokens.device)\n        sequence_features, pairwise_features = self.get_embeddings(src_tokens, padding_mask)\n\n        projected_msa = self.msa_projector(msa_profile_input)\n        concatenated_sequence_features = torch.cat((sequence_features, projected_msa), dim=-1)\n        adapted_sequence_features = self.adaptor(concatenated_sequence_features)\n\n        current_batch_size = noisy_xyz.shape[0] \n        final_sequence_features = adapted_sequence_features.expand(current_batch_size, -1, -1) \n        final_pairwise_features = pairwise_features.expand(current_batch_size, -1, -1, -1)\n\n        final_pairwise_features = self.embed_pairwise_distance(final_pairwise_features, noisy_xyz) \n\n        time_embedding = self.time_embedder(timestep_t).unsqueeze(1)\n        tgt_representation = self.xyz_norm(final_sequence_features + self.xyz_embedder(noisy_xyz) + time_embedding)\n        tgt_representation = self.time_norm(tgt_representation + self.time_mlp(tgt_representation))\n\n        for layer_module in self.structure_module_stack: # Renamed 'layer' to 'layer_module'\n            tgt_representation = layer_module((tgt_representation, final_sequence_features, final_pairwise_features, noisy_xyz, None)) \n        \n        predicted_noise_or_coords = self.xyz_predictor(tgt_representation)\n        return predicted_noise_or_coords\n\nclass SimpleStructureModule(nn.Module):\n    def __init__(self, d_model, nhead, dim_feedforward, pairwise_dimension, dropout=0.1):\n        super(SimpleStructureModule, self).__init__()\n        self.self_attn = MultiHeadAttention(d_model, nhead, d_model//nhead, d_model//nhead, dropout=dropout)\n        self.linear1 = nn.Linear(d_model, dim_feedforward)\n        self.dropout_ffn = nn.Dropout(dropout) \n        self.linear2 = nn.Linear(dim_feedforward, d_model)\n        self.norm1 = nn.LayerNorm(d_model)\n        self.norm2 = nn.LayerNorm(d_model)\n        self.dropout_res = nn.Dropout(dropout) \n        self.pairwise2heads = nn.Linear(pairwise_dimension, nhead, bias=False)\n        self.pairwise_norm = nn.LayerNorm(pairwise_dimension)\n        self.activation = nn.GELU()\n\n    def forward(self, inputs_tuple): \n        tgt, src_seq_features_ignored, pairwise_features_in, pred_coords_t_ignored, src_mask_attn_ignored = inputs_tuple\n        \n        pairwise_bias = self.pairwise2heads(self.pairwise_norm(pairwise_features_in)).permute(0, 3, 1, 2)\n        \n        res = tgt\n        tgt, attention_weights = self.self_attn(tgt, tgt, tgt, mask=pairwise_bias, src_mask=None)\n        tgt = res + self.dropout_res(tgt)\n        tgt = self.norm1(tgt)\n        \n        res = tgt\n        tgt = self.linear2(self.dropout_ffn(self.activation(self.linear1(tgt))))\n        tgt = res + self.dropout_res(tgt) \n        tgt = self.norm2(tgt)\n        return tgt\n\nclass Config:\n    def __init__(self, **entries): self.__dict__.update(entries); self.entries = entries\n    def print(self): print(self.entries)\n\ndef load_config_from_yaml(file_path):\n    if not os.path.exists(file_path):\n        print(f\"ERROR: Model YAML config not found at {file_path}\")\n        raise FileNotFoundError(f\"Model YAML config not found at {file_path}\")\n    with open(file_path, 'r') as file: config_data = yaml.safe_load(file)\n    print(f\"Successfully loaded model YAML config from {file_path}\")\n    return Config(**config_data)\n\nmodel_arch_config = load_config_from_yaml(global_notebook_config['model_config_path'])\n\nmodel = finetuned_RibonanzaNet_MSA(model_arch_config, msa_feature_dim=5) \nif torch.cuda.is_available():\n    model = model.cuda()\nprint(\"Model finetuned_RibonanzaNet_MSA instantiated.\")\n\nclass Diffusion(nn.Module):\n    def __init__(self, model_instance, n_times=1000, beta_minmax=[1e-4, 2e-2]):\n        super(Diffusion, self).__init__()\n        self.n_times = n_times\n        self.model = model_instance \n        beta_1, beta_T = beta_minmax\n        betas = torch.linspace(start=beta_1, end=beta_T, steps=n_times)\n        self.sqrt_betas = torch.sqrt(betas)\n        self.alphas = 1 - betas\n        self.sqrt_alphas = torch.sqrt(self.alphas)\n        alpha_bars = torch.cumprod(self.alphas, dim=0)\n        self.sqrt_one_minus_alpha_bars = torch.sqrt(1 - alpha_bars)\n        self.sqrt_alpha_bars = torch.sqrt(alpha_bars)\n\n    def extract(self, a, t, x_shape):\n        b, *_ = t.shape; out = a.gather(-1, t)\n        return out.reshape(b, *((1,) * (len(x_shape) - 1)))\n\n    def random_rotation_matrix_batch(self, B, device):\n        A = torch.randn(B, 3, 3, device=device)\n        Q, R = torch.linalg.qr(A)\n        det = torch.det(Q)\n        Q[det < 0, :, 0] *= -1 \n        return Q \n\n    def make_noisy(self, x_zeros_coords, t_timesteps):\n        x_zeros_norm = x_zeros_coords / 35.0 \n        rotation_matrices = self.random_rotation_matrix_batch(x_zeros_norm.shape[0], x_zeros_norm.device)\n        x_zeros_rotated = torch.bmm(x_zeros_norm, rotation_matrices) \n        epsilon_noise = torch.randn_like(x_zeros_rotated).to(x_zeros_rotated.device)\n        sqrt_alpha_bar_t = self.extract(self.sqrt_alpha_bars.to(x_zeros_rotated.device), t_timesteps, x_zeros_rotated.shape)\n        sqrt_one_minus_alpha_bar_t = self.extract(self.sqrt_one_minus_alpha_bars.to(x_zeros_rotated.device), t_timesteps, x_zeros_rotated.shape)\n        noisy_sample = x_zeros_rotated * sqrt_alpha_bar_t + epsilon_noise * sqrt_one_minus_alpha_bar_t\n        return noisy_sample.detach(), epsilon_noise\n\n    def forward(self, true_coords_batch, src_token_batch, msa_profile_batch):\n        num_noise_levels_per_sample = 48 \n        x_zeros_repeated = true_coords_batch.repeat(num_noise_levels_per_sample, 1, 1)\n        t_timesteps = torch.randint(low=0, high=self.n_times, size=(num_noise_levels_per_sample,)).long().to(true_coords_batch.device)\n        perturbed_xyz, actual_noise_added = self.make_noisy(x_zeros_repeated, t_timesteps)\n        # The model's forward method expects src_token_batch and msa_profile_batch \n        # to be [1, L, Dims] (or whatever the original batch size is, here 1)\n        # and it handles the repeat/expand internally if noisy_xyz has a larger first dimension.\n        # Let's ensure the model's forward method is consistent with this.\n        # The current finetuned_RibonanzaNet_MSA.forward takes [B,L] for src_tokens, [B,L,F] for msa_profile\n        # and [N_noise, L, 3] for noisy_xyz, [N_noise] for timestep_t.\n        # It then repeats sequence_features to match N_noise. This is correct.\n        predicted_output, distogram_prediction = self.model(\n            src_token_batch, \n            msa_profile_batch, \n            perturbed_xyz,     \n            t_timesteps       \n        )\n        return perturbed_xyz, actual_noise_added, predicted_output, distogram_prediction\n\n    def denoise_at_t_step(self, x_t_coords, src_tokens_cond, msa_profile_cond, current_timestep_val, t_idx_val):\n        N_samples = x_t_coords.shape[0]\n        z_noise = torch.randn_like(x_t_coords) if t_idx_val > 1 else torch.zeros_like(x_t_coords)\n        z_noise = z_noise.to(src_tokens_cond.device)\n        epsilon_predicted = self.model.denoise(src_tokens_cond, msa_profile_cond, x_t_coords, current_timestep_val)\n        alpha_t = self.extract(self.alphas.to(x_t_coords.device), current_timestep_val, x_t_coords.shape)\n        sqrt_alpha_t = self.extract(self.sqrt_alphas.to(x_t_coords.device), current_timestep_val, x_t_coords.shape)\n        sqrt_one_minus_alpha_bar_t = self.extract(self.sqrt_one_minus_alpha_bars.to(x_t_coords.device), current_timestep_val, x_t_coords.shape)\n        sqrt_beta_t = self.extract(self.sqrt_betas.to(x_t_coords.device), current_timestep_val, x_t_coords.shape)\n        x_t_minus_1 = (1/sqrt_alpha_t)*(x_t_coords-((1-alpha_t)/sqrt_one_minus_alpha_bar_t)*epsilon_predicted)+sqrt_beta_t*z_noise\n        return x_t_minus_1\n\n    def sample_structures(self, src_tokens_sample, msa_profile_sample, num_samples_to_gen):\n        device = src_tokens_sample.device\n        seq_len = src_tokens_sample.shape[1]\n        x_t_current = torch.randn((num_samples_to_gen, seq_len, 3), device=device)\n        \n        padding_mask_sample = torch.ones_like(src_tokens_sample).long()\n        with torch.no_grad():\n             # Ensure get_embeddings is called on self.model (the finetuned_RibonanzaNet_MSA instance)\n             _, pairwise_features_static = self.model.get_embeddings(src_tokens_sample, padding_mask_sample)\n        distogram_pred_static = self.model.distogram_predictor(pairwise_features_static).squeeze(0)\n        \n        for t_idx_val in range(self.n_times - 1, -1, -1): \n            current_timestep_val = torch.tensor([t_idx_val]).repeat(num_samples_to_gen).long().to(device)\n            with torch.no_grad():\n                x_t_current = self.denoise_at_t_step(x_t_current, src_tokens_sample, msa_profile_sample, current_timestep_val, t_idx_val)\n        x_0_final_coords = x_t_current * 35.0 \n        return x_0_final_coords, distogram_pred_static\n\n# Ensure global_notebook_config is defined before this line\ndiffusion_model = Diffusion(model, n_times=global_notebook_config['n_times'])\nif torch.cuda.is_available():\n    diffusion_model = diffusion_model.cuda()\nprint(\"--- Model and Diffusion Wrapper Initialized ---\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"--- Setting up Optimizer and Scheduler ---\")\noptimizer = torch.optim.Adam(diffusion_model.model.parameters(), lr=config['learning_rate'], weight_decay=config['weight_decay'])\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=(config['epochs'] - config['cos_epoch']) * len(train_loader) // config['gradient_accumulation_steps'])\ncriterion_distogram = torch.nn.CrossEntropyLoss(reduction='none') # For distogram auxiliary loss\n\n# Automatic Mixed Precision Scaler\nif config['mixed_precision'] in ['fp16', 'bf16']:\n    scaler = torch.cuda.amp.GradScaler(enabled=True)\nelse:\n    scaler = None","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"--- Defining Loss and Metric Functions ---\")\ndef calculate_pairwise_distances(coords_tensor, epsilon=1e-6): # Renamed\n    # coords_tensor: (L, 3) or (B, L, 3)\n    # Returns (L,L) or (B,L,L) distance matrix\n    diff = coords_tensor.unsqueeze(-2) - coords_tensor.unsqueeze(-3)\n    return (torch.square(diff).sum(dim=-1) + epsilon).sqrt()\n\ndef drmsd_loss_masked(pred_coords, true_coords, mask=None, d_clamp=None, Z=10): # dRMAE in original, often called dRMSD\n    # pred_coords, true_coords: (L, 3)\n    # mask: (L, L) boolean, True for pairs to include\n    \n    pred_dm = calculate_pairwise_distances(pred_coords)\n    true_dm = calculate_pairwise_distances(true_coords)\n\n    if d_clamp is not None:\n        pred_dm = torch.clamp(pred_dm, max=d_clamp)\n        true_dm = torch.clamp(true_dm, max=d_clamp)\n\n    pair_errors = torch.abs(pred_dm - true_dm)\n    \n    if mask is None:\n        # Default mask: non-NaN in true_dm and exclude diagonal\n        mask = ~torch.isnan(true_dm)\n        diag_mask = torch.eye(true_dm.shape[-1], dtype=torch.bool, device=true_dm.device)\n        if pred_dm.ndim == 3: # Batch dimension\n            diag_mask = diag_mask.unsqueeze(0).expand(pred_dm.shape[0],-1,-1)\n        mask = mask & ~diag_mask\n        \n    if mask.sum() == 0: return torch.tensor(0.0, device=pred_coords.device) # Avoid division by zero\n    \n    return pair_errors[mask].mean() / Z # Normalize by Z as in original\n\ndef kabsch_rmsd_masked(pred_coords, true_coords, true_mask=None):\n    # pred_coords, true_coords: (L, 3)\n    # true_mask: (L) boolean, True for residues to include in alignment/RMSD\n    \n    if true_mask is None:\n        true_mask = ~torch.isnan(true_coords[:, 0]) # Mask based on first coordinate\n        \n    if true_mask.sum() < 3: # Need at least 3 points for SVD\n        return torch.tensor(float('nan'), device=pred_coords.device)\n\n    pred_points = pred_coords[true_mask]\n    true_points = true_coords[true_mask]\n\n    pred_centroid = pred_points.mean(dim=0, keepdim=True)\n    true_centroid = true_points.mean(dim=0, keepdim=True)\n\n    pred_centered = pred_points - pred_centroid\n    true_centered = true_points - true_centroid\n\n    cov_matrix = pred_centered.T @ true_centered\n    try:\n        U, S, Vt = torch.linalg.svd(cov_matrix)\n    except torch._C._LinAlgError: # Catch SVD convergence error\n        return torch.tensor(float('nan'), device=pred_coords.device)\n\n\n    R = Vt.T @ U.T # Note: In PyTorch SVD, Vt is already V.T\n    if torch.det(R) < 0:\n        Vt_corrected = Vt.clone()\n        Vt_corrected[-1, :] *= -1\n        R = Vt_corrected.T @ U.T\n    \n    aligned_pred = pred_centered @ R.T + true_centroid # Use R.T for rotation\n    \n    rmsd = torch.sqrt(torch.square(aligned_pred - true_points).sum(dim=-1).mean())\n    return rmsd\n\ndef lddt_score_torch(pred_coords, true_coords, true_mask=None, cutoffs=[0.5, 1.0, 2.0, 4.0], neighborhood_cutoff=15.0):\n    # pred_coords, true_coords: (L, 3)\n    # true_mask: (L) boolean\n    if true_mask is None:\n        true_mask = ~torch.isnan(true_coords[:, 0])\n\n    if true_mask.sum() < 2: return torch.tensor(0.0, device=pred_coords.device)\n\n    pred_subset = pred_coords[true_mask]\n    true_subset = true_coords[true_mask]\n    \n    L_subset = true_subset.shape[0]\n    \n    true_dist_matrix = calculate_pairwise_distances(true_subset)\n    pred_dist_matrix = calculate_pairwise_distances(pred_subset)\n    \n    # Create neighborhood mask\n    # Exclude diagonal and distances > neighborhood_cutoff\n    neighbor_mask = (true_dist_matrix < neighborhood_cutoff) & (true_dist_matrix > 1e-6) # Exclude self (diagonal)\n    \n    scores_per_residue = []\n    for i in range(L_subset):\n        res_neighbor_mask = neighbor_mask[i]\n        if res_neighbor_mask.sum() == 0: # No neighbors within cutoff\n            scores_per_residue.append(torch.tensor(0.0, device=pred_coords.device)) # Or handle as per official lDDT\n            continue\n            \n        d_true = true_dist_matrix[i, res_neighbor_mask]\n        d_pred = pred_dist_matrix[i, res_neighbor_mask]\n        \n        abs_diff = torch.abs(d_true - d_pred)\n        \n        threshold_scores = []\n        for cutoff_val in cutoffs:\n            threshold_scores.append((abs_diff < cutoff_val).float().mean())\n        scores_per_residue.append(torch.stack(threshold_scores).mean())\n        \n    if not scores_per_residue: return torch.tensor(0.0, device=pred_coords.device)\n    return torch.stack(scores_per_residue).mean()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"--- Starting Training Loop ---\")\nbest_val_rmsd = float('inf')\n\nfor epoch in range(config['epochs']):\n    diffusion_model.train() # Set model to training mode\n    \n    total_epoch_loss_noise = 0\n    total_epoch_loss_distogram = 0\n    \n    optimizer.zero_grad() # Zero gradients at the beginning of accumulation\n\n    progress_bar = tqdm(enumerate(train_loader), total=len(train_loader), desc=f\"Epoch {epoch+1}/{config['epochs']}\")\n    for batch_idx, batch in progress_bar:\n        src_tokens = batch['sequence'].cuda(non_blocking=True)\n        true_xyz = batch['xyz'].cuda(non_blocking=True) # Shape: [1, L, 3]\n        msa_prof = batch['msa_profile'].cuda(non_blocking=True) # Shape: [1, L, 5]\n        \n        # Create a mask for valid (non-NaN) coordinates in the true_xyz\n        # This mask should be [1, L, 3] and then broadcasted or applied carefully\n        coord_mask_for_loss = ~torch.isnan(true_xyz) # True where coords are valid\n        \n        # Ensure true_xyz does not contain NaNs before passing to make_noisy, or handle them inside\n        # For simplicity, let's assume make_noisy can handle NaNs by centering/normalizing based on valid points\n        # Or, we can impute NaNs before noising, but the loss should only be on valid points.\n        # The original notebook's loss calculation implicitly handled this by masking after noise prediction.\n        \n        # The diffusion.forward now returns: perturbed_xyz, actual_noise_added, predicted_output, distogram_prediction\n        # predicted_output is the predicted noise (or coords, depending on model's final layer)\n        \n        # Use autocast for forward pass if mixed precision is enabled\n        if scaler:\n            with torch.cuda.amp.autocast(dtype=torch.bfloat16 if config['mixed_precision']=='bf16' else torch.float16):\n                _, actual_noise, predicted_noise, distogram_pred = diffusion_model(true_xyz, src_tokens, msa_prof)\n                \n                # Noise prediction loss (MSE) - only on valid coordinates\n                # Ensure actual_noise and predicted_noise have the same shape [num_noise_levels, L, 3]\n                # And coord_mask_for_loss is broadcastable, e.g. [1, L, 3] -> [num_noise_levels, L, 3]\n                expanded_coord_mask = coord_mask_for_loss.repeat(predicted_noise.shape[0], 1, 1)\n                \n                loss_noise = torch.square(actual_noise - predicted_noise)\n                loss_noise = loss_noise[expanded_coord_mask].mean() # Only average over valid, masked points\n\n                # Distogram loss\n                # true_pairwise_dist needs to be calculated from true_xyz and binned\n                true_pairwise_dist_matrix = calculate_pairwise_distances(true_xyz.squeeze(0)) # Remove batch dim for calc\n                distogram_target = true_pairwise_dist_matrix.clamp(min=2, max=39).long() - 2 # Bins 0-37 for dists 2-39A\n                \n                distogram_mask_loss = ~torch.isnan(true_pairwise_dist_matrix) & (true_pairwise_dist_matrix > 1e-6) # Exclude self, NaNs\n                distogram_pred_squeezed = distogram_pred.squeeze(0) # Remove batch dim [1,L,L,40] -> [L,L,40]\n\n                if distogram_mask_loss.sum() > 0:\n                    loss_distogram = criterion_distogram(\n                        distogram_pred_squeezed[distogram_mask_loss], # (N_valid_pairs, num_bins)\n                        distogram_target[distogram_mask_loss]        # (N_valid_pairs)\n                    ).mean()\n                else:\n                    loss_distogram = torch.tensor(0.0, device=loss_noise.device)\n\n                combined_loss = loss_noise + 0.2 * loss_distogram\n                \n            # Scale loss and backward pass\n            scaler.scale(combined_loss / config['gradient_accumulation_steps']).backward()\n        else: # No mixed precision\n            _, actual_noise, predicted_noise, distogram_pred = diffusion_model(true_xyz, src_tokens, msa_prof)\n            expanded_coord_mask = coord_mask_for_loss.repeat(predicted_noise.shape[0], 1, 1)\n            loss_noise = torch.square(actual_noise - predicted_noise)[expanded_coord_mask].mean()\n            \n            true_pairwise_dist_matrix = calculate_pairwise_distances(true_xyz.squeeze(0))\n            distogram_target = true_pairwise_dist_matrix.clamp(min=2, max=39).long() - 2\n            distogram_mask_loss = ~torch.isnan(true_pairwise_dist_matrix) & (true_pairwise_dist_matrix > 1e-6)\n            distogram_pred_squeezed = distogram_pred.squeeze(0)\n            if distogram_mask_loss.sum() > 0:\n                loss_distogram = criterion_distogram(distogram_pred_squeezed[distogram_mask_loss], distogram_target[distogram_mask_loss]).mean()\n            else:\n                loss_distogram = torch.tensor(0.0, device=loss_noise.device)\n\n            combined_loss = loss_noise + 0.2 * loss_distogram\n            (combined_loss / config['gradient_accumulation_steps']).backward()\n\n        if (batch_idx + 1) % config['gradient_accumulation_steps'] == 0 or (batch_idx + 1) == len(train_loader):\n            if scaler:\n                scaler.unscale_(optimizer) # Unscale before clipping\n            torch.nn.utils.clip_grad_norm_(diffusion_model.model.parameters(), config['grad_clip'])\n            if scaler:\n                scaler.step(optimizer)\n                scaler.update()\n            else:\n                optimizer.step()\n            optimizer.zero_grad() # Zero gradients after step or end of accumulation\n\n        if (epoch + 1) > config['cos_epoch'] and ( (batch_idx + 1) % config['gradient_accumulation_steps'] == 0 or (batch_idx + 1) == len(train_loader)):\n             scheduler.step() # Step scheduler\n\n        total_epoch_loss_noise += loss_noise.item()\n        total_epoch_loss_distogram += loss_distogram.item()\n        \n        progress_bar.set_postfix({\n            \"Loss_Noise\": f\"{total_epoch_loss_noise / (batch_idx + 1):.4f}\",\n            \"Loss_Dist\": f\"{total_epoch_loss_distogram / (batch_idx + 1):.4f}\",\n            \"LR\": f\"{optimizer.param_groups[0]['lr']:.2e}\"\n        })\n\n    # --- Validation Phase ---\n    diffusion_model.eval() # Set model to evaluation mode\n    val_loss_drmsd_accum = 0\n    val_rmsd_accum = 0\n    val_lddt_accum = 0\n    \n    val_progress_bar = tqdm(val_loader, desc=f\"Epoch {epoch+1} Validation\")\n    with torch.no_grad():\n        for batch in val_progress_bar:\n            src_tokens_val = batch['sequence'].cuda(non_blocking=True) # Shape [1, L]\n            true_xyz_val = batch['xyz'].cuda(non_blocking=True).squeeze(0) # Shape [L, 3]\n            msa_prof_val = batch['msa_profile'].cuda(non_blocking=True) # Shape [1, L, 5]\n\n            # Generate 1 sample for validation\n            # diffusion_model.sample_structures now returns (coords, distogram_logits)\n            predicted_xyz_val_batch, _ = diffusion_model.sample_structures(src_tokens_val, msa_prof_val, num_samples_to_gen=1)\n            predicted_xyz_val = predicted_xyz_val_batch.squeeze(0) # Get the single sample [L,3]\n            \n            # Calculate metrics\n            # Ensure true_xyz_val has no batch dimension for metric functions\n            loss_drmsd_val = drmsd_loss_masked(predicted_xyz_val, true_xyz_val, d_clamp=config['d_clamp'])\n            rmsd_val = kabsch_rmsd_masked(predicted_xyz_val, true_xyz_val)\n            lddt_val = lddt_score_torch(predicted_xyz_val, true_xyz_val)\n\n            if not torch.isnan(loss_drmsd_val): val_loss_drmsd_accum += loss_drmsd_val.item()\n            if not torch.isnan(rmsd_val): val_rmsd_accum += rmsd_val.item()\n            if not torch.isnan(lddt_val): val_lddt_accum += lddt_val.item()\n            \n            val_progress_bar.set_postfix({\n                \"Val_dRMSD\": f\"{val_loss_drmsd_accum / (val_progress_bar.n + 1):.4f}\",\n                \"Val_RMSD\": f\"{val_rmsd_accum / (val_progress_bar.n + 1):.4f}\",\n                \"Val_lDDT\": f\"{val_lddt_accum / (val_progress_bar.n + 1):.4f}\"\n            })\n\n    avg_val_drmsd = val_loss_drmsd_accum / len(val_loader)\n    avg_val_rmsd = val_rmsd_accum / len(val_loader)\n    avg_val_lddt = val_lddt_accum / len(val_loader)\n\n    print(f\"Epoch {epoch+1} Summary: Avg Train Noise Loss: {total_epoch_loss_noise / len(train_loader):.4f}, \"\n          f\"Avg Train Distogram Loss: {total_epoch_loss_distogram / len(train_loader):.4f}\")\n    print(f\"Validation: Avg dRMSD: {avg_val_drmsd:.4f}, Avg RMSD: {avg_val_rmsd:.4f}, Avg lDDT: {avg_val_lddt:.4f}\")\n\n    if avg_val_rmsd < best_val_rmsd:\n        best_val_rmsd = avg_val_rmsd\n        torch.save(diffusion_model.model.state_dict(), \"rnet2_finetuned_msa_best_rmsd.pt\")\n        print(f\"Saved new best model with Val RMSD: {best_val_rmsd:.4f}\")\n\nprint(\"--- Training Complete ---\")\n# To save the final model regardless\ntorch.save(diffusion_model.model.state_dict(), \"rnet2_finetuned_msa_final.pt\")\nprint(\"Saved final model state.\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}