{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":87793,"databundleVersionId":11403143,"sourceType":"competition"}],"dockerImageVersionId":30918,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## Foreword\nThis is the first Kaggle competition I'm joining, so I hope with this notebook I can help anyone else starting out, or anyone that prefers using Torch instead of TF. Happy scoreboard climbing!\n## Checking contents of dataset","metadata":{}},{"cell_type":"code","source":"#Get path directories to double make sure everything is there\n    #I've commented the printing out not to flood the outputs\nimport os\n\"\"\"\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\"\"\"","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-03-15T22:03:31.791447Z","iopub.execute_input":"2025-03-15T22:03:31.791896Z","iopub.status.idle":"2025-03-15T22:03:31.801240Z","shell.execute_reply.started":"2025-03-15T22:03:31.791858Z","shell.execute_reply":"2025-03-15T22:03:31.799463Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Imports, Torch specific","metadata":{}},{"cell_type":"code","source":"#Add more imports as needed\nimport numpy as np \nimport pandas as pd \nfrom collections import Counter\nimport torch\nfrom torch.utils.data import Dataset, DataLoader\nimport torch.nn as nn\nimport torch.optim as optim\nimport torch.nn.functional as F\nfrom tqdm import tqdm  \n\n#Set device, will auto-pick between GPU or CPU accordingly\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(device)\n#Set seed to desired value\ntorch.manual_seed(46)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-15T22:03:33.631938Z","iopub.execute_input":"2025-03-15T22:03:33.632341Z","iopub.status.idle":"2025-03-15T22:03:36.176699Z","shell.execute_reply.started":"2025-03-15T22:03:33.632307Z","shell.execute_reply":"2025-03-15T22:03:36.175621Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Data Loading and Prep","metadata":{}},{"cell_type":"code","source":"#Load the CSV files as Pandas dataframes\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\")\nvalid_sequences_df = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/validation_sequences.csv\")\nvalid_labels_df = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/validation_labels.csv\")\ntest_sequences_df = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/test_sequences.csv\")\n\nprint(\"Train sequence rows\", train_sequences_df.shape[0])\nprint(\"Train label rows\", train_labels_df.shape[0])\nprint(\"Valid sequence rows\", valid_sequences_df.shape[0])\nprint(\"Valid label rows\", valid_labels_df.shape[0])\nprint(\"Test sequence rows\", test_sequences_df.shape[0])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-15T22:03:38.269049Z","iopub.execute_input":"2025-03-15T22:03:38.269574Z","iopub.status.idle":"2025-03-15T22:03:38.738064Z","shell.execute_reply.started":"2025-03-15T22:03:38.269543Z","shell.execute_reply":"2025-03-15T22:03:38.736915Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Filtering based on temporal cutoff:\n    #Based on competition description, advised if you are using the current validation_sequences csv (as of date of writing March 15 2025, subject to change)\n\"\"\"\ntrain_sequences_df[\"temporal_cutoff\"] = pd.to_datetime(train_sequences_df[\"temporal_cutoff\"])\ntrain_sequences_df = train_sequences_df[train_sequences_df[\"temporal_cutoff\"] < \"2022-05-27\"]\n\nprint(\"Train sequence rows\", train_sequences_df.shape[0])\n\"\"\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-15T22:03:48.888334Z","iopub.execute_input":"2025-03-15T22:03:48.888746Z","iopub.status.idle":"2025-03-15T22:03:48.895346Z","shell.execute_reply.started":"2025-03-15T22:03:48.888716Z","shell.execute_reply":"2025-03-15T22:03:48.894215Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#train_labels_df has A lot of NaN values in the coords\nnan_check = train_labels_df[['x_1', 'y_1', 'z_1']].isna().sum()\n\nprint(\"NaN values in each column:\")\nprint(nan_check)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-15T22:03:49.340331Z","iopub.execute_input":"2025-03-15T22:03:49.340728Z","iopub.status.idle":"2025-03-15T22:03:49.370668Z","shell.execute_reply.started":"2025-03-15T22:03:49.340698Z","shell.execute_reply":"2025-03-15T22:03:49.369527Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Any series that has NaN in the coords is removed here\n\n# Step 1: Extract unique target identifiers from NaN rows\nnan_identifiers = train_labels_df[train_labels_df[['x_1', 'y_1', 'z_1']].isna().any(axis=1)]['ID']\nnan_identifiers = nan_identifiers.str.rsplit(\"_\", n=1).str[0].unique().tolist()\n\nprint(f\"Identifiers with NaNs: {nan_identifiers}\")\n\n# Step 2: Remove rows from train_sequences_df where target_id is in the nan_identifiers list\ntrain_sequences_df = train_sequences_df[~train_sequences_df['target_id'].isin(nan_identifiers)].reset_index(drop=True)\nprint(f\"Remaining sequences in train_sequences_df: {len(train_sequences_df)}\")\n\n# Step 3: Remove rows from train_labels_df where ID starts with any identifier in nan_identifiers\n# We use str.startswith to check if the 'ID' starts with any of the values in nan_identifiers\ntrain_labels_df = train_labels_df[~train_labels_df['ID'].apply(lambda x: any(x.startswith(identifier) for identifier in nan_identifiers))].reset_index(drop=True)\nprint(f\"Remaining labels in train_labels_df: {len(train_labels_df)}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-15T22:03:51.312003Z","iopub.execute_input":"2025-03-15T22:03:51.312358Z","iopub.status.idle":"2025-03-15T22:03:55.306155Z","shell.execute_reply.started":"2025-03-15T22:03:51.312331Z","shell.execute_reply":"2025-03-15T22:03:55.305043Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Checking again for labels, no NaN values here\n\n# Step 1: Extract unique target identifiers from NaN rows in validation data\nnan_identifiers_valid = valid_labels_df[valid_labels_df[['x_1', 'y_1', 'z_1']].isna().any(axis=1)]['ID']\nnan_identifiers_valid = nan_identifiers_valid.str.rsplit(\"_\", n=1).str[0].unique().tolist()\n\nprint(f\"Identifiers with NaNs in validation data: {nan_identifiers_valid}\")\n\n# Step 2: Remove rows from valid_sequences_df where target_id is in the nan_identifiers_valid list\nvalid_sequences_df = valid_sequences_df[~valid_sequences_df['target_id'].isin(nan_identifiers_valid)].reset_index(drop=True)\nprint(f\"Remaining sequences in valid_sequences_df: {len(valid_sequences_df)}\")\n\n# Step 3: Remove rows from valid_labels_df where ID starts with any identifier in nan_identifiers_valid\nvalid_labels_df = valid_labels_df[~valid_labels_df['ID'].apply(lambda x: any(x.startswith(identifier) for identifier in nan_identifiers_valid))].reset_index(drop=True)\nprint(f\"Remaining labels in valid_labels_df: {len(valid_labels_df)}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-15T22:03:55.307584Z","iopub.execute_input":"2025-03-15T22:03:55.307948Z","iopub.status.idle":"2025-03-15T22:03:55.321885Z","shell.execute_reply.started":"2025-03-15T22:03:55.307914Z","shell.execute_reply":"2025-03-15T22:03:55.320748Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Dataset class\nI've tried to make this class as modular as possible, including options whether to normalize coords or not (recommended) and whether to one-hot encode the labels or map them to sequential values for use with an embedding layer. I've stuck with the standard function set of a Torch dataset.","metadata":{}},{"cell_type":"code","source":"class RNADataset(Dataset):\n    def __init__(self, sequences: pd.DataFrame, labels: pd.DataFrame, normalize_coords: bool = True, one_hot=False):\n        self.sequences = sequences\n        self.labels = labels\n        self.normalize_coords = normalize_coords\n        self.one_hot = one_hot\n\n        #One hot encoding if true, integer mapping for false (for use with embedding layers etc)\n            #Note that I've set it up so that you can use zeroes for padding \n        if self.one_hot:\n            self.nucleotide_map = {'A': [1, 0, 0, 0],\n                                   'U': [0, 1, 0, 0],\n                                   'G': [0, 0, 1, 0],\n                                   'C': [0, 0, 0, 1]}\n        else:\n            self.nucleotide_map = {'A': 1, 'U': 2, 'G': 3, 'C': 4}\n            #self.nucleotide_map = {'A': 0, 'U': 1, 'G': 2, 'C': 3} #Alternate option\n        \n        #Here I'm manually filtering some sequences that have Xs or -, they're very few so they don't interfere too much\n        self.sequences = self.sequences[self.sequences['sequence'].apply(lambda seq: all(nuc not in ['X', '-'] for nuc in seq))]\n        \n        #Normalization happens on global min maxes for consistent scaling\n        if self.normalize_coords:\n            self.coord_min = torch.tensor(labels[['x_1', 'y_1', 'z_1']].min().values, dtype=torch.float32)\n            self.coord_max = torch.tensor(labels[['x_1', 'y_1', 'z_1']].max().values, dtype=torch.float32)\n\n    def __len__(self):\n        return len(self.sequences)\n\n    def __getitem__(self, idx):\n        seq_row = self.sequences.iloc[idx]\n        target_id = seq_row['target_id'] #eg ABCD\n        sequence = seq_row['sequence'] #GCAU sequence\n\n        if self.one_hot:\n            #One-hot encode the sequence\n            sequence = torch.tensor([self.nucleotide_map[nuc] for nuc in sequence if nuc in self.nucleotide_map], dtype=torch.float32)\n        else:\n            #Convert the sequence to a list of integer indices as per the mapping in __init__\n            sequence = torch.tensor([self.nucleotide_map[nuc] for nuc in sequence if nuc in self.nucleotide_map], dtype=torch.long)\n        \n        #Get corresponding label rows\n        label_rows = self.labels[self.labels['ID'].str.startswith(target_id + \"_\")]\n        \n        if label_rows.empty:\n            raise ValueError(f\"No matching labels found for target_id: {target_id}\")\n        \n        #Convert coordinates to tensor\n        coords = torch.tensor(label_rows[['x_1', 'y_1', 'z_1']].values, dtype=torch.float32)\n        \n        #Normalize coordinates if set to true\n        if self.normalize_coords:\n            coords = (coords - self.coord_min) / (self.coord_max - self.coord_min)\n        \n        return sequence, coords\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-15T22:05:05.786725Z","iopub.execute_input":"2025-03-15T22:05:05.787066Z","iopub.status.idle":"2025-03-15T22:05:05.797530Z","shell.execute_reply.started":"2025-03-15T22:05:05.787041Z","shell.execute_reply":"2025-03-15T22:05:05.796425Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Custom collate function  \nDynamic padding to max length of sequence in batch, enables > 1 batch size when training <br>\n- Find max length in batch <br>\n- Create zeroes tensors of that len <br>\n- FIll those tensors with the values of each sequence <br>\n- Mask to make sure zeroes don't contribute to loss function calculations <br>\n\nCollate function is called when instantiating the Dataloader","metadata":{}},{"cell_type":"code","source":"def collate_fn(batch):\n    sequences, coords = zip(*batch)  # Unzip batch into sequences and coordinates\n    \n    # Find max length in this batch\n    max_len = max(seq.shape[0] for seq in sequences)\n\n    # Pad sequences and coords to max length\n    padded_sequences = torch.zeros((len(sequences), max_len, 4))  # (batch_size, max_seq_len, one_hot_dim)\n    padded_coords = torch.zeros((len(coords), max_len, 3))  # (batch_size, max_seq_len, coord_dim)\n    mask = torch.zeros((len(sequences), max_len))  # Mask for valid positions\n\n    for i, (seq, coord) in enumerate(zip(sequences, coords)):\n        seq_len = seq.shape[0]\n        padded_sequences[i, :seq_len, :] = seq\n        padded_coords[i, :seq_len, :] = coord\n        mask[i, :seq_len] = 1  # Mark valid positions\n\n    return padded_sequences, padded_coords, mask","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-15T22:05:08.460193Z","iopub.execute_input":"2025-03-15T22:05:08.460595Z","iopub.status.idle":"2025-03-15T22:05:08.467437Z","shell.execute_reply.started":"2025-03-15T22:05:08.460564Z","shell.execute_reply":"2025-03-15T22:05:08.466163Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Template Model\nSimple placeholder LSTM model to make sure everything works before working on the model (importing extra data, inference etc)","metadata":{}},{"cell_type":"code","source":"#Simple placeholder LSTM model to make sure everything works before working on the model (importing extra data, inference etc)\nclass SimpleLSTMModel(nn.Module):\n    def __init__(self, input_size=4, hidden_size=64, num_layers=1, output_size=3, dropout=0.1, bidirectional=False):\n        super(SimpleLSTMModel, self).__init__()\n        \n        # LSTM Layer with optional dropout and bidirectionality\n        self.lstm = nn.LSTM(input_size, hidden_size, num_layers, batch_first=True, dropout=dropout, bidirectional=bidirectional)\n        \n        # Fully connected output layer for each timestep\n        # If bidirectional, multiply hidden_size by 2\n        self.fc = nn.Linear(hidden_size * (2 if bidirectional else 1), output_size)\n        \n        # Apply weight initialization\n        self.apply(self._initialize_weights)\n\n    def forward(self, x):\n        \"\"\"\n        x: (batch_size, seq_length, input_size=4)\n        \"\"\"\n        # Get all hidden states from the LSTM\n        lstm_out, _ = self.lstm(x)  # lstm_out: (batch_size, seq_length, hidden_size * num_directions)\n        \n        # If bidirectional, concatenate the forward and backward hidden states\n        if self.lstm.bidirectional:\n            lstm_out = lstm_out[:, :, :self.lstm.hidden_size] + lstm_out[:, :, self.lstm.hidden_size:]  # sum of forward and backward hidden states\n\n        # Apply the fully connected layer to each timestep in the sequence\n        out = self.fc(lstm_out)  # Output shape: (batch_size, seq_length, 3)\n\n        return out  # Output shape: (batch_size, seq_length, 3)\n\n    def _initialize_weights(self, module):\n        if isinstance(module, nn.LSTM):\n            # Xavier initialization for LSTM weights\n            for name, param in module.named_parameters():\n                if 'weight_ih' in name:\n                    nn.init.xavier_uniform_(param)\n                elif 'weight_hh' in name:\n                    nn.init.xavier_uniform_(param)\n                elif 'bias' in name:\n                    nn.init.zeros_(param)\n\n        elif isinstance(module, nn.Linear):\n            # Xavier initialization for Linear layers\n            nn.init.xavier_uniform_(module.weight)\n            if module.bias is not None:\n                nn.init.zeros_(module.bias)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-15T22:05:10.049154Z","iopub.execute_input":"2025-03-15T22:05:10.049531Z","iopub.status.idle":"2025-03-15T22:05:10.059789Z","shell.execute_reply.started":"2025-03-15T22:05:10.049501Z","shell.execute_reply":"2025-03-15T22:05:10.058201Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training Hyperparams","metadata":{}},{"cell_type":"code","source":"#Input output dims\ninput_size = 4  # G, C, A, U\noutput_size = 3  # x_1 y_1 z_1 per nucleotide\n\n#for throwaway model \nhidden_size = 64  # Hidden layer size\nmodel = SimpleLSTMModel(input_size, hidden_size, output_size)\nmodel.to(device)\n\n# Loss \n#criterion = nn.MSELoss()  #default option \ncriterion = torch.nn.SmoothL1Loss() #more mathematically stable opposed to MSE\nlr = 0.0001\n\n#AdamW to have some regularization in the optim \noptimizer = optim.AdamW(model.parameters(), lr=lr, weight_decay = 0.0001)\n\n#modify as needed\nbatch_size = 8","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-15T22:05:12.164073Z","iopub.execute_input":"2025-03-15T22:05:12.164423Z","iopub.status.idle":"2025-03-15T22:05:12.175241Z","shell.execute_reply.started":"2025-03-15T22:05:12.164396Z","shell.execute_reply":"2025-03-15T22:05:12.173940Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Dataset Classes, setting batch size = 1 for valid to avoid padding\ntrain_set = RNADataset(sequences = train_sequences_df, labels = train_labels_df, normalize_coords = True, one_hot = True)\nvalid_set = RNADataset(sequences = valid_sequences_df, labels = valid_labels_df, normalize_coords = True, one_hot = True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-15T22:05:19.190994Z","iopub.execute_input":"2025-03-15T22:05:19.191559Z","iopub.status.idle":"2025-03-15T22:05:19.236447Z","shell.execute_reply.started":"2025-03-15T22:05:19.191479Z","shell.execute_reply":"2025-03-15T22:05:19.235228Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#DataLoaders \ntrain_dataloader = DataLoader(train_set, batch_size=batch_size, collate_fn = collate_fn, shuffle=True)\nprint(len(train_dataloader))\nvalid_dataloader = DataLoader(valid_set, batch_size = 1, collate_fn = collate_fn, shuffle = True)\nprint(len(valid_dataloader))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-15T22:05:20.696623Z","iopub.execute_input":"2025-03-15T22:05:20.697081Z","iopub.status.idle":"2025-03-15T22:05:20.704817Z","shell.execute_reply.started":"2025-03-15T22:05:20.697043Z","shell.execute_reply":"2025-03-15T22:05:20.703280Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#LR Scheduler, set to your liking\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\n\n# Scheduler setup\nT_max = len(train_dataloader)  # One full cycle is one epoch\neta_min = 1e-6                 # Minimum learning rate\n\nscheduler = CosineAnnealingLR(\n    optimizer=optimizer,\n    T_max=T_max,\n    eta_min=eta_min\n)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-15T22:05:21.515583Z","iopub.execute_input":"2025-03-15T22:05:21.515933Z","iopub.status.idle":"2025-03-15T22:05:21.521208Z","shell.execute_reply.started":"2025-03-15T22:05:21.515907Z","shell.execute_reply":"2025-03-15T22:05:21.520037Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Epochs to train for\nnum_epochs = 10","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-15T22:05:23.596857Z","iopub.execute_input":"2025-03-15T22:05:23.597417Z","iopub.status.idle":"2025-03-15T22:05:23.602912Z","shell.execute_reply.started":"2025-03-15T22:05:23.597357Z","shell.execute_reply":"2025-03-15T22:05:23.601555Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Train Loop","metadata":{}},{"cell_type":"code","source":"for epoch in range(num_epochs):\n    model.train()  # Set the model to training mode\n    epoch_loss = 0.0\n    val_loss = 0.0  # Initialize validation loss tracker\n\n    # Train Loop\n    with tqdm(train_dataloader, desc=f\"Epoch {epoch+1}/{num_epochs}\", unit=\"batch\") as pbar:\n        for batch_idx, (sequence, coordinates, mask) in enumerate(pbar):  # Unpack mask\n            # Move data to the appropriate device\n            sequence, coordinates, mask = sequence.to(device), coordinates.to(device), mask.to(device)\n\n            optimizer.zero_grad()  # Zero gradients for each batch    \n\n            # Forward pass\n            predictions = model(sequence.float())  # Model predicts 3D coordinates\n\n            # Apply mask to ignore padding in loss calculation\n            loss = (criterion(predictions, coordinates.float()) * mask.unsqueeze(-1)).sum() / mask.sum()\n\n            epoch_loss += loss.item()\n\n            # Backward pass and optimization\n            loss.backward()  # Backpropagate the gradients\n            torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)  # Clip gradients\n            optimizer.step()  # Update the weights\n\n            #scores = tm_score(coordinates, predictions)\n            #print(f\"Batch TM-scores: {scores}\")\n            \n            # Update progress bar description with the current loss\n            pbar.set_postfix(loss=loss.item())\n\n    # Print average training loss for the epoch\n    print(f\"Training Loss: {epoch_loss / len(train_dataloader):.4f}\")\n\n    # Validation Loop\n    model.eval()  # Set the model to evaluation mode\n    with torch.no_grad():  # No gradient calculation during evaluation\n        for batch_idx, (sequence, coordinates, mask) in enumerate(valid_dataloader):  # Unpack mask\n            # Move data to the appropriate device\n            sequence, coordinates, mask = sequence.to(device), coordinates.to(device), mask.to(device)\n\n            # Forward pass\n            predictions = model(sequence.float())\n\n            # Apply mask to ignore padding in loss calculation\n            loss = (criterion(predictions, coordinates.float()) * mask.unsqueeze(-1)).sum() / mask.sum()\n            val_loss += loss.item()\n\n    # Print average validation loss for the epoch\n    print(f\"Validation Loss: {val_loss / len(valid_dataloader):.4f}\")\n\n    #Change path to save model accordingly     \n    model_path = '/kaggle/working/lstm_{}'.format(epoch) #CHANGE ACCORDINGLY\n\n    \n    torch.save({\n        'model_dict': model.state_dict(),\n        'optimizer_dict': optimizer.state_dict(),\n        'scheduler_dict': scheduler.state_dict(),\n        #'scaler_state_dict': scaler.state_dict() #amp\n    }, model_path)\n\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-15T22:05:25.042437Z","iopub.execute_input":"2025-03-15T22:05:25.042846Z","iopub.status.idle":"2025-03-15T22:11:02.866988Z","shell.execute_reply.started":"2025-03-15T22:05:25.042816Z","shell.execute_reply":"2025-03-15T22:11:02.865455Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Inference + submission.csv creation \nNote that the submission.csv requires 5 coordinate sets per nucleotide in the sequence. The current iteration just copies the same results over and over since we set eval and the template LSTM isn't stochastic on inference. You will have to modify the code to make 5 unique predictions each time if your model is stochastic on inference","metadata":{}},{"cell_type":"code","source":"#Load last trained model, change accordingly\ninf_model_path = \"/kaggle/working/lstm_9\"\n\n#Use the same params as training model to instantiate inf model, otherwise the dims won't match\ninf_model = SimpleLSTMModel(input_size, hidden_size, output_size)\n\ncheckpoint= torch.load(inf_model_path)\nprint(checkpoint.keys())\n\ninf_model.load_state_dict(checkpoint['model_dict'])\n\ninf_model.eval()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-15T22:11:34.796208Z","iopub.execute_input":"2025-03-15T22:11:34.796630Z","iopub.status.idle":"2025-03-15T22:11:34.816886Z","shell.execute_reply.started":"2025-03-15T22:11:34.796598Z","shell.execute_reply":"2025-03-15T22:11:34.815551Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Separate Dataset class to handle inference\nclass RNAInferenceDataset(Dataset):\n    def __init__(self, sequences: pd.DataFrame, normalize_coords: bool = True, one_hot=False):\n        self.sequences = sequences\n        self.normalize_coords = normalize_coords\n        self.one_hot = one_hot\n\n        if self.one_hot:\n            self.nucleotide_map = {'A': [1, 0, 0, 0],\n                                   'U': [0, 1, 0, 0],\n                                   'G': [0, 0, 1, 0],\n                                   'C': [0, 0, 0, 1]}\n        else:\n            # Mapping nucleotides to integer indices (for embedding layer)\n            self.nucleotide_map = {'A': 1, 'U': 2, 'G': 3, 'C': 4}\n            #self.nucleotide_map = {'A': 0, 'U': 1, 'G': 2, 'C': 3}\n\n    def __len__(self):\n        return len(self.sequences)\n\n    def __getitem__(self, idx):\n        seq_row = self.sequences.iloc[idx]\n        target_id = seq_row['target_id']\n        sequence_raw = seq_row['sequence']\n\n        if self.one_hot:\n            # One-hot encode the sequence\n            sequence = torch.tensor([self.nucleotide_map[nuc] for nuc in sequence_raw if nuc in self.nucleotide_map], dtype=torch.float32)\n        else:\n            # Convert the sequence to a list of integer indices for embedding\n            sequence = torch.tensor([self.nucleotide_map[nuc] for nuc in sequence_raw if nuc in self.nucleotide_map], dtype=torch.long)\n\n        \n        return target_id, sequence, sequence_raw","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-15T22:14:36.491351Z","iopub.execute_input":"2025-03-15T22:14:36.491748Z","iopub.status.idle":"2025-03-15T22:14:36.501781Z","shell.execute_reply.started":"2025-03-15T22:14:36.491717Z","shell.execute_reply":"2025-03-15T22:14:36.500166Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def run_inference(model, df, device, batch_size=1):  # Set batch_size to 1\n    dataset = RNAInferenceDataset(sequences=df, one_hot=True)  # or one_hot=False depending on your setup\n    dataloader = DataLoader(dataset, batch_size=batch_size, shuffle=False)\n    \n    results = []\n    \n    with torch.no_grad():\n        for target_id, sequence, sequence_raw in dataloader:\n            sequence = sequence.to(device)\n            \n            preds = model(sequence)  \n            preds = preds.cpu().numpy()\n\n            #\n            for i, pred in enumerate(preds[0]):  #Iterate through each 3D coordinate set\n                results.append({\n                    \"ID\": f\"{target_id[0]}_{i+1}\",\n                    \"resname\": sequence_raw[0][i],  \n                    \"resid\": i+1,\n                    \"x_1\": pred[0],\n                    \"y_1\": pred[1],\n                    \"z_1\": pred[2],\n                    #not the most elegant solution but gets the job done if your inference doens't change on multiple runs\n                    \"x_2\": pred[0],\n                    \"y_2\": pred[1],\n                    \"z_2\": pred[2],\n                    \"x_3\": pred[0],\n                    \"y_3\": pred[1],\n                    \"z_3\": pred[2],\n                    \"x_4\": pred[0],\n                    \"y_4\": pred[1],\n                    \"z_4\": pred[2],\n                    \"x_5\": pred[0],\n                    \"y_5\": pred[1],\n                    \"z_5\": pred[2]\n                })\n    \n    return pd.DataFrame(results)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-15T22:14:57.963868Z","iopub.execute_input":"2025-03-15T22:14:57.964200Z","iopub.status.idle":"2025-03-15T22:14:57.973112Z","shell.execute_reply.started":"2025-03-15T22:14:57.964175Z","shell.execute_reply":"2025-03-15T22:14:57.971591Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submission_df = run_inference(inf_model, test_sequences_df, device)\n#verify format is competition appropriate\nprint(submission_df.head(5))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-15T22:14:59.218987Z","iopub.execute_input":"2025-03-15T22:14:59.219387Z","iopub.status.idle":"2025-03-15T22:14:59.335687Z","shell.execute_reply.started":"2025-03-15T22:14:59.219346Z","shell.execute_reply":"2025-03-15T22:14:59.334457Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Finally, write the csv\nsubmission_df.to_csv(\"submission.csv\", index=False)\nprint(\"Submission file saved as submission.csv\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-15T22:15:04.626404Z","iopub.execute_input":"2025-03-15T22:15:04.626778Z","iopub.status.idle":"2025-03-15T22:15:04.683248Z","shell.execute_reply.started":"2025-03-15T22:15:04.626752Z","shell.execute_reply":"2025-03-15T22:15:04.682391Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Afterword\nAnd you're done! Hopefully this notebook will help you get set up quickly. If everything runs and submission.csv is written in the correct format, hitting Submit for the comp should get you on the leaderboard. You can change the model block to whatever model you feel like using, as long as it's using Torch and the in/out channels remain the same. Good luck in the contest!","metadata":{}}]}