{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":59094,"databundleVersionId":7010844,"sourceType":"competition"},{"sourceId":120352,"sourceType":"modelInstanceVersion","modelInstanceId":101229,"modelId":125424}],"dockerImageVersionId":30776,"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\nimport torch \nimport numpy as np \nimport pandas as pd\nfrom sklearn.preprocessing import OneHotEncoder, StandardScaler\n\nimport os\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","execution":{"iopub.status.busy":"2024-09-27T20:38:40.343503Z","iopub.execute_input":"2024-09-27T20:38:40.344187Z","iopub.status.idle":"2024-09-27T20:38:45.759384Z","shell.execute_reply.started":"2024-09-27T20:38:40.344141Z","shell.execute_reply":"2024-09-27T20:38:45.758385Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f'Using {device}')","metadata":{"execution":{"iopub.status.busy":"2024-09-27T20:38:45.761439Z","iopub.execute_input":"2024-09-27T20:38:45.762014Z","iopub.status.idle":"2024-09-27T20:38:45.804533Z","shell.execute_reply.started":"2024-09-27T20:38:45.761968Z","shell.execute_reply":"2024-09-27T20:38:45.802959Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Installing a package\n!pip install -q rdkit captum git+https://github.com/samoturk/mol2vec;","metadata":{"execution":{"iopub.status.busy":"2024-09-27T20:38:45.806547Z","iopub.execute_input":"2024-09-27T20:38:45.806971Z","iopub.status.idle":"2024-09-27T20:39:17.245264Z","shell.execute_reply.started":"2024-09-27T20:38:45.806924Z","shell.execute_reply":"2024-09-27T20:39:17.244092Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from gensim.models import word2vec\nmodel1 = word2vec.Word2Vec.load('/kaggle/input/m/ankushhv/mol2vec/pytorch/default/1/model_300dim.pkl')","metadata":{"execution":{"iopub.status.busy":"2024-09-27T20:39:17.247933Z","iopub.execute_input":"2024-09-27T20:39:17.248290Z","iopub.status.idle":"2024-09-27T20:39:31.727257Z","shell.execute_reply.started":"2024-09-27T20:39:17.248253Z","shell.execute_reply":"2024-09-27T20:39:31.726412Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_parquet('../input/open-problems-single-cell-perturbations/de_train.parquet')\ndf.tail()","metadata":{"execution":{"iopub.status.busy":"2024-09-27T20:39:31.728666Z","iopub.execute_input":"2024-09-27T20:39:31.728942Z","iopub.status.idle":"2024-09-27T20:39:34.417249Z","shell.execute_reply.started":"2024-09-27T20:39:31.728911Z","shell.execute_reply":"2024-09-27T20:39:34.416255Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## **Preprocessing Classes**","metadata":{}},{"cell_type":"code","source":"from abc import ABC, abstractmethod\nclass BaseEmbedding(ABC):\n    @abstractmethod\n    def preprocess(self, df):\n        pass\n\n    @abstractmethod\n    def get_embedding(self, df):\n        pass\n\n    @property\n    @abstractmethod\n    def embedding_size(self):\n        pass","metadata":{"execution":{"iopub.status.busy":"2024-09-27T20:39:34.418424Z","iopub.execute_input":"2024-09-27T20:39:34.418732Z","iopub.status.idle":"2024-09-27T20:39:34.424503Z","shell.execute_reply.started":"2024-09-27T20:39:34.418699Z","shell.execute_reply":"2024-09-27T20:39:34.423567Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nfrom rdkit import Chem\nfrom rdkit.Chem import Descriptors\nfrom sklearn.preprocessing import OneHotEncoder, StandardScaler\nfrom gensim.models import word2vec\nfrom mol2vec.features import mol2alt_sentence, MolSentence\n\nclass OneHotEmbedding(BaseEmbedding):\n    def __init__(self, columns):\n        self.columns = columns\n        self.encoder = OneHotEncoder(sparse=False)\n        self._embedding_size = None\n\n    def preprocess(self, df):\n        return df[self.columns]\n\n    def get_embedding(self, df, fit=True):\n        if fit:\n            encoded_features = self.encoder.fit_transform(df[self.columns])\n        else:\n            encoded_features = self.encoder.transform(df[self.columns])\n        self._embedding_size = encoded_features.shape[1]\n        return pd.DataFrame(encoded_features, columns=self.encoder.get_feature_names_out(self.columns))\n\n    @property\n    def embedding_size(self):\n        return self._embedding_size\n    \n    \nclass SMILESEmbedding(BaseEmbedding):\n    def __init__(self):\n        self.scaler = StandardScaler()\n\n    def preprocess(self, df):\n        return df\n\n    def get_embedding(self, df, fit=True):\n        smiles_info_list = df['SMILES'].apply(self.extract_smiles_info)\n        smiles_info_df = pd.DataFrame(smiles_info_list.tolist())\n        if fit:\n            smiles_info_df[['Molecular Weight', 'LogP', 'Number of Atoms', 'Number of Bonds']] = self.scaler.fit_transform(\n                smiles_info_df[['Molecular Weight', 'LogP', 'Number of Atoms', 'Number of Bonds']]\n            )\n        else:\n            smiles_info_df[['Molecular Weight', 'LogP', 'Number of Atoms', 'Number of Bonds']] = self.scaler.transform(\n                smiles_info_df[['Molecular Weight', 'LogP', 'Number of Atoms', 'Number of Bonds']]\n            )\n        return smiles_info_df\n\n    @staticmethod\n    def extract_smiles_info(smiles):\n        if smiles is None or smiles == '':\n            return None\n        mol = Chem.MolFromSmiles(smiles)\n        if mol is None:\n            return None\n        info = {\n            'Molecular Weight': Descriptors.MolWt(mol),\n            'LogP': Descriptors.MolLogP(mol),\n            'Number of Atoms': mol.GetNumAtoms(),\n            'Number of Bonds': mol.GetNumBonds()\n        }\n        return info\n\n    @property\n    def embedding_size(self):\n        return 4  # Molecular Weight, LogP, Number of Atoms, Number of Bonds\n\nclass Mol2VecEmbedding(BaseEmbedding):\n    def __init__(self, model_path):\n        self.model = word2vec.Word2Vec.load(model_path)\n        self.keys = set(self.model.wv.key_to_index.keys())\n\n    def preprocess(self, df):\n        df['mol'] = df['SMILES'].apply(lambda x: Chem.MolFromSmiles(x))\n        df['sentence'] = df.apply(lambda x: MolSentence(mol2alt_sentence(x['mol'], 1)), axis=1)\n        return df\n\n    def get_embedding(self, df, fit=True):\n        df['vector'] = df['sentence'].apply(lambda sentence: self.sentence_to_vector(sentence))\n        vector_dim = len(self.model.wv.get_vector(next(iter(self.keys))))\n        vector_columns = [f'vector_{i}' for i in range(vector_dim)]\n        return pd.DataFrame(df['vector'].tolist(), columns=vector_columns)\n\n    def sentence_to_vector(self, sentence, unseen=False, unseen_vec=np.zeros(300)):\n        if unseen:\n            vec = sum([self.model.wv.get_vector(word) if word in self.keys else unseen_vec for word in sentence])\n        else:\n            vec = sum([self.model.wv.get_vector(word) for word in sentence if word in self.keys])\n        return vec\n\n    @property\n    def embedding_size(self):\n        return len(self.model.wv.get_vector(next(iter(self.keys))))","metadata":{"execution":{"iopub.status.busy":"2024-09-27T20:39:34.426348Z","iopub.execute_input":"2024-09-27T20:39:34.426748Z","iopub.status.idle":"2024-09-27T20:39:34.970091Z","shell.execute_reply.started":"2024-09-27T20:39:34.426704Z","shell.execute_reply":"2024-09-27T20:39:34.969251Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Preprocessor:\n    def __init__(self, embeddings):\n        self.embeddings = embeddings\n        self.embedding_indices = {}\n\n    def preprocess(self, df, fit=True):\n        processed_dfs = []\n        current_index = 0\n\n        for embedding in self.embeddings:\n            preprocessed_df = embedding.preprocess(df)\n            embedded_df = embedding.get_embedding(preprocessed_df, fit=fit)\n            processed_dfs.append(embedded_df)\n\n            embedding_name = embedding.__class__.__name__\n            self.embedding_indices[embedding_name] = (current_index, current_index + embedding.embedding_size)\n            current_index += embedding.embedding_size\n\n        return pd.concat(processed_dfs, axis=1)","metadata":{"execution":{"iopub.status.busy":"2024-09-27T20:39:34.971490Z","iopub.execute_input":"2024-09-27T20:39:34.971805Z","iopub.status.idle":"2024-09-27T20:39:34.978830Z","shell.execute_reply.started":"2024-09-27T20:39:34.971772Z","shell.execute_reply":"2024-09-27T20:39:34.977981Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"onehot_embedding = OneHotEmbedding(columns=['cell_type', 'sm_name'])\nsmiles_embedding = SMILESEmbedding()\nmol2vec_embedding = Mol2VecEmbedding('/kaggle/input/m/ankushhv/mol2vec/pytorch/default/1/model_300dim.pkl')","metadata":{"execution":{"iopub.status.busy":"2024-09-27T20:39:34.979968Z","iopub.execute_input":"2024-09-27T20:39:34.980360Z","iopub.status.idle":"2024-09-27T20:39:36.061472Z","shell.execute_reply.started":"2024-09-27T20:39:34.980314Z","shell.execute_reply":"2024-09-27T20:39:36.060590Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preprocessor_train = Preprocessor([onehot_embedding, smiles_embedding, mol2vec_embedding])","metadata":{"execution":{"iopub.status.busy":"2024-09-27T20:39:36.065467Z","iopub.execute_input":"2024-09-27T20:39:36.065802Z","iopub.status.idle":"2024-09-27T20:39:36.070474Z","shell.execute_reply.started":"2024-09-27T20:39:36.065767Z","shell.execute_reply":"2024-09-27T20:39:36.069274Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"target_cols = ['cell_type','sm_name','sm_lincs_id','SMILES','control']\ntargets = df.drop(columns=target_cols)\nprocessed_df_train = preprocessor_train.preprocess(df, fit=True)","metadata":{"execution":{"iopub.status.busy":"2024-09-27T20:39:36.071903Z","iopub.execute_input":"2024-09-27T20:39:36.072282Z","iopub.status.idle":"2024-09-27T20:39:38.261596Z","shell.execute_reply.started":"2024-09-27T20:39:36.072239Z","shell.execute_reply":"2024-09-27T20:39:38.260512Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# processed_df_train.shape\nprint(\"Processed Training DataFrame shape:\", processed_df_train.shape)\nprint(\"Training Embedding indices:\", preprocessor_train.embedding_indices.values())","metadata":{"execution":{"iopub.status.busy":"2024-09-27T20:39:38.262805Z","iopub.execute_input":"2024-09-27T20:39:38.263165Z","iopub.status.idle":"2024-09-27T20:39:38.269072Z","shell.execute_reply.started":"2024-09-27T20:39:38.263128Z","shell.execute_reply":"2024-09-27T20:39:38.268053Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"feature_names = list(processed_df_train.columns)\nprint(len(feature_names))\ntrain_array = processed_df_train.to_numpy()","metadata":{"execution":{"iopub.status.busy":"2024-09-27T20:39:38.270614Z","iopub.execute_input":"2024-09-27T20:39:38.271073Z","iopub.status.idle":"2024-09-27T20:39:38.428985Z","shell.execute_reply.started":"2024-09-27T20:39:38.271004Z","shell.execute_reply":"2024-09-27T20:39:38.427903Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(targets.dtypes)","metadata":{"execution":{"iopub.status.busy":"2024-09-27T20:39:38.430449Z","iopub.execute_input":"2024-09-27T20:39:38.430781Z","iopub.status.idle":"2024-09-27T20:39:38.443368Z","shell.execute_reply.started":"2024-09-27T20:39:38.430747Z","shell.execute_reply":"2024-09-27T20:39:38.442256Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nX = torch.tensor(train_array,dtype=torch.float32)\ny = torch.tensor(targets.values, dtype=torch.float32)\nprint(f\"X shape: {X.shape}, y shape: {y.shape}\")","metadata":{"execution":{"iopub.status.busy":"2024-09-27T20:39:38.444657Z","iopub.execute_input":"2024-09-27T20:39:38.444976Z","iopub.status.idle":"2024-09-27T20:39:38.503990Z","shell.execute_reply.started":"2024-09-27T20:39:38.444943Z","shell.execute_reply":"2024-09-27T20:39:38.503169Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nfrom torch.utils.data import DataLoader, Dataset, TensorDataset\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nX = X.to(device)\ny = y.to(device)\n\ndataset = TensorDataset(X, y)\ntrain_size = int(0.8 * len(dataset))\nval_size = len(dataset) - train_size\ntrain_dataset, val_dataset = torch.utils.data.random_split(dataset, [train_size, val_size])\n\ntrain_loader = DataLoader(train_dataset, batch_size=16, shuffle=True)\nval_loader = DataLoader(val_dataset, batch_size=16)","metadata":{"execution":{"iopub.status.busy":"2024-09-27T20:39:38.505313Z","iopub.execute_input":"2024-09-27T20:39:38.505826Z","iopub.status.idle":"2024-09-27T20:39:38.723740Z","shell.execute_reply.started":"2024-09-27T20:39:38.505787Z","shell.execute_reply":"2024-09-27T20:39:38.722305Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom sklearn.metrics import mean_absolute_error\n\ntorch.manual_seed(69)\nnp.random.seed(69)\n\nclass LSTMExpert(nn.Module):\n    def __init__(self, input_size, output_size):\n        super(LSTMExpert, self).__init__()\n        self.lstm = nn.LSTM(input_size, 128, num_layers=2, batch_first=True)\n        self.linear = nn.Sequential(\n            nn.Linear(128, 512),\n            nn.ReLU(),\n            nn.Dropout(0.3),\n            nn.Linear(512, 512),\n            nn.ReLU(),\n            nn.Dropout(0.3)\n        )\n        self.head1 = nn.Linear(512, output_size)\n        \n    def forward(self, x):\n        if x.dim() == 2:\n            x = x.unsqueeze(1)\n        batch_size, seq_len, _ = x.size()\n        out, (hn, cn) = self.lstm(x)\n        if seq_len == 1:\n            out = out.squeeze(1)\n        else:\n            out = out[:, -1, :]\n        out = self.linear(out)\n        out = self.head1(out)\n        return out\n\nclass NNExpert(nn.Module):\n    def __init__(self, input_size, output_size, dropout_rate=0.5):\n        super(NNExpert, self).__init__()\n        self.network = nn.Sequential(\n            nn.Linear(input_size, 128),\n            nn.ReLU(),\n            nn.Dropout(dropout_rate),\n            nn.Linear(128, 256),\n            nn.ReLU(),\n            nn.Linear(256, 128),\n            nn.ReLU(),\n            nn.Linear(128, output_size),\n            nn.Tanh()\n        )\n\n    def forward(self, x):\n        if x.dim() == 3:\n            x = x.view(x.size(0), -1)\n        output = self.network(x)\n        return output\n\nclass GatingNetwork(nn.Module):\n    def __init__(self, input_size, num_experts):\n        super(GatingNetwork, self).__init__()\n\n        # Define the network using nn.Sequential\n        self.fc1 = nn.Sequential(\n            nn.Linear(input_size, 256),\n            nn.BatchNorm1d(256),\n            nn.ReLU(),\n            nn.Dropout(0.3)\n        )\n\n        self.fc2 = nn.Sequential(\n            nn.Linear(256, 256),\n            nn.BatchNorm1d(256),\n            nn.ReLU(),\n            nn.Dropout(0.3)\n        )\n\n        # Final layer for outputting expert weights\n        self.fc3 = nn.Sequential(\n            nn.Linear(256, num_experts),\n            nn.Softmax(dim=-1)\n        )\n\n    def forward(self, x):\n        identity = self.fc1[0](x)  # Pass through the first Linear layer\n\n        # Residual connection after applying the first and second blocks\n        out = self.fc1(x)          # First block (Linear -> BN -> ReLU -> Dropout)\n        out = self.fc2(out)        # Second block (Linear -> BN -> ReLU -> Dropout)\n\n        # Residual connection: add identity from the first Linear layer\n        out = out + identity\n\n        # Pass through the final layer to get the expert probabilities\n        return self.fc3(out)\n    \n# class GatingNetwork(nn.Module):\n#     def __init__(self, input_size, num_experts, initial_temperature=1.0):\n#         super(GatingNetwork, self).__init__()\n\n#         # First and second blocks using nn.Sequential\n#         self.fc1 = nn.Sequential(\n#             nn.Linear(input_size, 256),\n#             nn.BatchNorm1d(256),\n#             nn.ReLU(),\n#             nn.Dropout(0.3)\n#         )\n\n#         self.fc2 = nn.Sequential(\n#             nn.Linear(256, 256),\n#             nn.BatchNorm1d(256),\n#             nn.ReLU(),\n#             nn.Dropout(0.3)\n#         )\n\n#         # Attention mechanism\n#         self.attention = nn.MultiheadAttention(embed_dim=256, num_heads=4)\n\n#         # Final layer for outputting expert weights (logits)\n#         self.fc3 = nn.Linear(256, num_experts)\n\n#         # Temperature parameter\n#         self.temperature = nn.Parameter(torch.ones(1) * initial_temperature)\n\n#     def forward(self, x):\n#         identity = self.fc1[0](x)  # Pass through the first Linear layer\n\n#         # Pass through the first and second blocks\n#         out = self.fc1(x)\n#         out = self.fc2(out)\n\n#         # Residual connection\n#         out = out + identity\n\n#         # Attention mechanism: Add sequence dimension\n#         out = out.unsqueeze(0)\n#         attn_output, _ = self.attention(out, out, out)\n#         attn_output = attn_output.squeeze(0)\n\n#         # Pass through the final linear layer to get logits\n#         logits = self.fc3(attn_output)\n\n#         # Apply temperature scaling before softmax\n#         return F.softmax(logits / self.temperature, dim=-1)\n    \n# class GatingNetwork(nn.Module):\n#     def __init__(self, input_size, num_experts):\n#         super(GatingNetwork, self).__init__()\n#         self.lstm = nn.LSTM(input_size, 128, 1, batch_first=True)\n#         self.linear = nn.Sequential(\n#             nn.Linear(128, 256),  # Hidden state output from LSTM to fully connected\n#             nn.ReLU(),            # ReLU activation\n#             nn.Dropout(0.3),      # Dropout for regularization\n#             nn.Linear(256, num_experts),  # Second hidden layer\n#             nn.ReLU(),\n#             nn.Dropout(0.3)\n#         )\n#         self.softmax = nn.Softmax(dim=-1)\n\n#     def forward(self, x):\n#         # Check if input is a sequence or not\n#         if x.dim() == 2:\n#             # If input is (batch_size, input_size), add a sequence dimension\n#             x = x.unsqueeze(1)\n        \n#         # Now x shape should be (batch_size, sequence_length, input_size)\n#         lstm_out, _ = self.lstm(x)\n        \n#         # Use the last output of the LSTM\n#         last_output = lstm_out[:, -1, :]\n#         x = self.linear(last_output)\n#         x = self.softmax(x)\n#         return x\n\nclass MixtureOfExperts(nn.Module):\n    def __init__(self, expert_configs, output_size=18211):\n        super(MixtureOfExperts, self).__init__()\n        self.experts = nn.ModuleDict()\n        self.feature_indices = {}\n        total_features = set()\n\n        for expert_name, feature_indices in expert_configs.items():\n            input_size = len(feature_indices)\n            if expert_name.startswith('lstm'):\n                self.experts[expert_name] = LSTMExpert(input_size, output_size)\n            elif expert_name.startswith('nn'):\n                self.experts[expert_name] = NNExpert(input_size, output_size)\n            else:\n                raise ValueError(f\"Unknown expert type for {expert_name}\")\n            \n            self.feature_indices[expert_name] = feature_indices\n            total_features.update(feature_indices)\n\n        # Set up the gating network with the total number of features across all experts\n        self.gating_network = GatingNetwork(len(total_features), len(self.experts))\n        self.total_features = sorted(list(total_features))  # Keep indices sorted\n        self.last_gating_weights = None\n\n    def forward(self, data):\n        # Prepare input for gating network (select total features from the data)\n        gating_input = data[:, self.total_features]\n        gating_weights = self.gating_network(gating_input)\n        self.last_gating_weights = gating_weights.detach()\n\n        expert_outputs = []\n        for expert_name, expert in self.experts.items():\n            # Select features for this expert based on the indices\n            expert_input = data[:, self.feature_indices[expert_name]]\n            expert_output = expert(expert_input)\n            expert_outputs.append(expert_output)\n\n        # Stack and combine expert outputs using gating weights\n        expert_outputs = torch.stack(expert_outputs, dim=1)\n        combined_output = torch.sum(expert_outputs * gating_weights.unsqueeze(-1), dim=1)\n        return combined_output\n\n# Utility functions remain unchanged\ndef RMSE_rowwise_loss(y_pred, y_true):\n    return torch.sqrt(torch.mean((y_pred - y_true)**2, dim=1)).mean()\n\nclass EarlyStopping:\n    def __init__(self, patience=20, min_delta=0, verbose=False):\n        self.patience = patience\n        self.min_delta = min_delta\n        self.verbose = verbose\n        self.counter = 0\n        self.best_score = None\n        self.early_stop = False\n        self.val_loss_min = float('inf')\n\n    def __call__(self, val_loss):\n        score = -val_loss\n\n        if self.best_score is None:\n            self.best_score = score\n            self.val_loss_min = val_loss\n            return False\n        elif score < self.best_score + self.min_delta:\n            self.counter += 1\n            if self.verbose:\n                print(f'EarlyStopping counter: {self.counter} out of {self.patience}')\n            if self.counter >= self.patience:\n                self.early_stop = True\n                return True\n        else:\n            self.best_score = score\n            self.val_loss_min = val_loss\n            self.counter = 0\n        return False","metadata":{"execution":{"iopub.status.busy":"2024-09-27T20:47:40.178102Z","iopub.execute_input":"2024-09-27T20:47:40.178878Z","iopub.status.idle":"2024-09-27T20:47:40.209570Z","shell.execute_reply.started":"2024-09-27T20:47:40.178836Z","shell.execute_reply":"2024-09-27T20:47:40.208523Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tqdm import tqdm\n\ndef train_model(model, train_loader, val_loader, epochs=200, lr=1e-4):\n    optimizer = torch.optim.Adam(model.parameters(), lr=lr, weight_decay=1e-5)\n    scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='min', patience=2, factor=0.5)\n\n    early_stopping = EarlyStopping(patience=10)\n    train_losses = []\n    val_losses = []\n\n    for epoch in tqdm(range(epochs), desc=\"Training Progress\"):\n        model.train()\n        train_loss = 0\n        for batch, (X, y) in enumerate(train_loader):\n            optimizer.zero_grad()\n            output = model(X)\n            loss = RMSE_rowwise_loss(output, y)\n            loss.backward()\n            optimizer.step()\n            train_loss += loss.item()\n\n        train_losses.append(train_loss / len(train_loader))\n\n        model.eval()\n        val_loss = 0\n        with torch.no_grad():\n            for X, y in val_loader:\n                output = model(X)\n                val_loss += RMSE_rowwise_loss(output, y).item()\n\n        val_losses.append(val_loss / len(val_loader))\n        \n        if val_loss < early_stopping.val_loss_min:\n            best_model_state = model.state_dict()\n        \n        if early_stopping(val_loss):\n            print(\"Early stopping triggered\")\n            model.load_state_dict(best_model_state)  # Load the best model\n            break\n\n        if epoch % 10 == 0 or epoch == epochs - 1:  # Also print on the last epoch\n            tqdm.write(f\"Epoch {epoch}: Train Loss: {train_losses[-1]:.4f}, Val Loss: {val_losses[-1]:.4f}\")\n\n        scheduler.step(val_loss)\n\n    plt.figure(figsize=(10, 5))\n    plt.plot(train_losses, label='Train Loss')\n    plt.plot(val_losses, label='Val Loss')\n    plt.xlabel('Epochs')\n    plt.ylabel('Loss')\n    plt.legend()\n    plt.title('Training and Validation Losses')\n    plt.show()\n\n    return train_losses, val_losses","metadata":{"execution":{"iopub.status.busy":"2024-09-27T20:47:40.219556Z","iopub.execute_input":"2024-09-27T20:47:40.219914Z","iopub.status.idle":"2024-09-27T20:47:40.233353Z","shell.execute_reply.started":"2024-09-27T20:47:40.219878Z","shell.execute_reply":"2024-09-27T20:47:40.232311Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.optim.lr_scheduler import CosineAnnealingLR\n\ndef pretrain_expert(expert, train_loader, val_loader, feature_indices, device, epochs=100, lr=1e-4):\n    expert.to(device)\n    optimizer = optim.Adam(expert.parameters(), lr=lr, weight_decay=1e-5)\n    scheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='min', patience=2, factor=0.5)\n    early_stopping = EarlyStopping(patience=10)\n    \n    train_losses = []\n    val_losses = []\n    \n    for epoch in tqdm(range(epochs), desc=f\"Pretraining {type(expert).__name__}\"):\n        expert.train()\n        train_loss = 0\n        for batch, (X, y) in enumerate(train_loader):\n            X, y = X[:, feature_indices].to(device), y.to(device)\n            optimizer.zero_grad()\n            output = expert(X)\n            loss = RMSE_rowwise_loss(output, y)\n            loss.backward()\n            optimizer.step()\n            train_loss += loss.item()\n        \n        train_losses.append(train_loss / len(train_loader))\n        \n        expert.eval()\n        val_loss = 0\n        with torch.no_grad():\n            for X, y in val_loader:\n                X, y = X[:, feature_indices].to(device), y.to(device)\n                output = expert(X)\n                val_loss += RMSE_rowwise_loss(output, y).item()\n        \n        val_losses.append(val_loss / len(val_loader))\n        \n        if val_loss < early_stopping.val_loss_min:\n            best_model_state = expert.state_dict()\n        \n        if early_stopping(val_loss):\n            print(\"Early stopping triggered\")\n            expert.load_state_dict(best_model_state)  # Load the best model\n            break\n        \n        if epoch % 10 == 0 or epoch == epochs - 1:\n            tqdm.write(f\"Epoch {epoch}: Train Loss: {train_losses[-1]:.4f}, Val Loss: {val_losses[-1]:.4f}\")\n        \n        scheduler.step(val_losses[-1])\n    \n    return train_losses, val_losses\n\ndef train_gating_network(moe, train_loader, val_loader, device, epochs=50, lr=1e-3):\n    # Freeze expert parameters\n    for expert in moe.experts.values():\n        for param in expert.parameters():\n            param.requires_grad = False\n    \n    # Ensure gating network parameters are unfrozen\n    for param in moe.gating_network.parameters():\n        param.requires_grad = True\n    \n    # Move model to device\n    moe.to(device)\n    \n    optimizer = optim.Adam(moe.gating_network.parameters(), lr=1e-3, weight_decay=1e-4)\n    scheduler = CosineAnnealingLR(optimizer, T_max=10)\n    early_stopping = EarlyStopping(patience=10)\n    \n    train_losses = []\n    val_losses = []\n    \n    for epoch in tqdm(range(epochs), desc=\"Training Gating Network\"):\n        moe.train()\n        train_loss = 0\n        for batch, (X, y) in enumerate(train_loader):\n            X, y = X.to(device), y.to(device)\n            optimizer.zero_grad()\n            output = moe(X)\n            loss = RMSE_rowwise_loss(output, y)\n            loss.backward()\n            optimizer.step()\n            train_loss += loss.item()\n        \n        train_losses.append(train_loss / len(train_loader))\n        \n        moe.eval()\n        val_loss = 0\n        with torch.no_grad():\n            for X, y in val_loader:\n                X, y = X.to(device), y.to(device)\n                output = moe(X)\n                val_loss += RMSE_rowwise_loss(output, y).item()\n        \n        val_losses.append(val_loss / len(val_loader))\n        \n        if val_loss < early_stopping.val_loss_min:\n            best_model_state = moe.gating_network.state_dict()\n        \n        if early_stopping(val_loss):\n            print(\"Early stopping triggered\")\n            moe.gating_network.load_state_dict(best_model_state)  # Load the best model\n            break\n        \n        if epoch % 5 == 0 or epoch == epochs - 1:\n            tqdm.write(f\"Epoch {epoch}: Train Loss: {train_losses[-1]:.4f}, Val Loss: {val_losses[-1]:.4f}\")\n        \n        scheduler.step(val_losses[-1])\n    \n    # Unfreeze expert parameters after training\n    for expert in moe.experts.values():\n        for param in expert.parameters():\n            param.requires_grad = True\n    \n    return train_losses, val_losses\n\ndef plot_training_curves(train_losses, val_losses, title):\n    plt.figure(figsize=(10, 5))\n    plt.plot(train_losses, label='Train Loss')\n    plt.plot(val_losses, label='Validation Loss')\n    plt.title(title)\n    plt.xlabel('Epoch')\n    plt.ylabel('Loss')\n    plt.legend()\n    plt.show()    ","metadata":{"execution":{"iopub.status.busy":"2024-09-27T20:47:40.297243Z","iopub.execute_input":"2024-09-27T20:47:40.297629Z","iopub.status.idle":"2024-09-27T20:47:40.321293Z","shell.execute_reply.started":"2024-09-27T20:47:40.297592Z","shell.execute_reply":"2024-09-27T20:47:40.320291Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.model_selection import KFold\nfrom sklearn.metrics import mean_squared_error\nimport torch.optim as optim\nimport matplotlib.pyplot as plt\nimport torch\nfrom torch.utils.data import DataLoader, Dataset, Subset\nfrom sklearn.model_selection import KFold\nfrom captum.attr import IntegratedGradients, Saliency, DeepLift, ShapleyValueSampling\n\ndef cross_validate_model(model_class, config, X, y, n_splits=5, epochs=200, lr=0.001):\n    kfold = KFold(n_splits=n_splits, shuffle=True, random_state=42)\n    train_losses_cv = []\n    val_losses_cv = []\n    models = []\n\n    for fold, (train_idx, val_idx) in enumerate(kfold.split(X)):\n        print(f'Fold {fold+1}/{n_splits}')\n\n        # Create model instance for each fold\n        model = model_class(config).to(device)\n        model.float()\n\n        # Create Datasets for the fold\n        train_dataset = Subset(TensorDataset(X, y), train_idx)\n        val_dataset = Subset(TensorDataset(X, y), val_idx)\n\n        # Create DataLoaders for the fold\n        train_loader = DataLoader(train_dataset, batch_size=16, shuffle=True)\n        val_loader = DataLoader(val_dataset, batch_size=16)\n\n        # Train the model for the current fold\n        train_losses, val_losses = train_model(model, train_loader, val_loader, epochs, lr)\n        \n        # Store model and losses for later analysis\n        models.append(model)\n        train_losses_cv.append(train_losses)\n        val_losses_cv.append(val_losses)\n\n    # Plotting average cross-validation losses\n    avg_train_loss = np.mean(train_losses_cv, axis=0)\n    avg_val_loss = np.mean(val_losses_cv, axis=0)\n    \n    plt.figure(figsize=(10, 5))\n    plt.plot(avg_train_loss, label='Avg Train Loss')\n    plt.plot(avg_val_loss, label='Avg Val Loss')\n    plt.xlabel('Epochs')\n    plt.ylabel('Loss')\n    plt.legend()\n    plt.title(f'{n_splits}-Fold Cross-Validation Losses')\n    plt.show()\n\n    return models, train_losses_cv, val_losses_cv\n\ndef compute_feature_importance(model, X, feature_names, k=10, baseline=None, target=0):\n    \"\"\"\n    Compute feature importance using various interpretability methods.\n\n    Parameters:\n    model (nn.Module): Trained PyTorch model.\n    X (torch.Tensor): Input data.\n    feature_names (list): List of feature names.\n    k (int, optional): Number of top features to display. Default is 10.\n    baseline (torch.Tensor, optional): Baseline input for interpretability methods. If None, a zero baseline is used.\n    target (int, optional): Target output index for interpretability methods. Default is 0.\n    methods (tuple, optional): Interpretability methods to use. Can include 'integrated_gradients', 'saliency', 'deeplift', and 'shapley'. Default is all methods.\n    return_all (bool, optional): If True, returns the top k feature names, their importances, and the full feature importance array for each method. Default is False.\n\n    Returns:\n    If return_all is False (default):\n        top_k_features (dict): Dictionary mapping method name to list of top k feature names.\n    If return_all is True:\n        top_k_features (dict), top_k_importances (dict), all_importances (dict)\n    \"\"\"\n    # Ensure the model is on CPU\n    model.cpu()\n    model.eval()  # Set model to evaluation mode\n    \n    # Move input to CPU if necessary\n    X = X.cpu()\n    # Choose baseline: it should have the same shape as X (defaults to zero baseline)\n    if baseline is None:\n        baseline = torch.zeros_like(X)\n    else:\n        baseline = baseline.cpu()  # Ensure baseline is on CPU\n    \n    top_k_features = {}\n    top_k_importances = {}\n    all_importances = {}\n    \n    method = 'integrated_gradients'\n    explainer = IntegratedGradients(model)\n    attributions = explainer.attribute(X, baselines=baseline, target=target)\n    \n#     for method in methods:\n#         # Choose interpretability method\n#         if method == 'integrated_gradients':\n#             explainer = IntegratedGradients(model)\n#             attributions = explainer.attribute(X, baselines=baseline, target=target)\n#         elif method == 'saliency':\n#             explainer = Saliency(model)\n#             attributions = explainer.attribute(X, target=target)\n#         elif method == 'deeplift':\n#             explainer = DeepLift(model)\n#             attributions = explainer.attribute(X, baselines=baseline, target=target)\n#         elif method == 'shapley':\n#             explainer = ShapleyValueSampling(model)\n#             attributions = explainer.attribute(X, baselines=baseline, target=target, show_progress=True)\n#         else:\n#             raise ValueError(f\"Unsupported method: {method}. Use 'integrated_gradients', 'saliency', 'deeplift', or 'shapley'.\")\n        \n        # Convert attributions to numpy for further processing\n    attributions_np = attributions.cpu().detach().numpy()\n    # Average feature importance across all samples\n    avg_importance = attributions_np.mean(axis=0)\n\n    # Get the indices of the top k features\n    top_k_indices = np.argsort(np.abs(avg_importance))[-k:][::-1]\n\n    # Get the names and importances of the top k features\n    top_k_features[method] = [feature_names[i] for i in top_k_indices]\n    top_k_importances[method] = avg_importance[top_k_indices]\n    all_importances[method] = avg_importance\n\n    # Plot the attributions for top k features\n    plt.figure(figsize=(10, 6))\n    plt.barh(np.array(feature_names)[top_k_indices], top_k_importances[method])\n    plt.xlabel(\"Average Importance across samples\")\n    plt.ylabel(\"Top Features\")\n    plt.title(f\"Top {k} Feature Importance ({method})\")\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2024-09-27T20:47:40.323429Z","iopub.execute_input":"2024-09-27T20:47:40.323795Z","iopub.status.idle":"2024-09-27T20:47:40.344133Z","shell.execute_reply.started":"2024-09-27T20:47:40.323759Z","shell.execute_reply":"2024-09-27T20:47:40.343271Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\ndef visualize_expert_weights(model, data_loader, expert_names, num_samples=100):\n    model.eval()\n    all_weights = []\n    all_labels = []\n    \n    with torch.no_grad():\n        for i, (inputs, labels) in enumerate(data_loader):\n            if i >= num_samples:\n                break\n            _ = model(inputs)  # Forward pass\n            weights = model.last_gating_weights\n            all_weights.append(weights.cpu().numpy())\n            all_labels.extend(labels.cpu().numpy())\n    \n    all_weights = np.concatenate(all_weights, axis=0)\n    \n    plt.figure(figsize=(12, 8))\n    sns.heatmap(all_weights, cmap=\"YlOrRd\", annot=False, fmt=\".2f\", cbar_kws={'label': 'Expert Weight'},\n                xticklabels=expert_names)\n    plt.title(\"Expert Weights for Different Input Instances\")\n    plt.xlabel(\"Expert\")\n    plt.ylabel(\"Input Instance\")\n    plt.tight_layout()\n    plt.show()\n    \n    # Bar plot for average expert weights\n    avg_weights = all_weights.mean(axis=0)\n    plt.figure(figsize=(10, 6))\n    plt.bar(expert_names, avg_weights)\n    plt.title(\"Average Expert Weights\")\n    plt.xlabel(\"Expert\")\n    plt.ylabel(\"Average Weight\")\n    plt.xticks(rotation=45, ha=\"right\")\n    plt.tight_layout()\n    plt.show()\n\n#     return all_weights, all_labels","metadata":{"execution":{"iopub.status.busy":"2024-09-27T20:47:40.345914Z","iopub.execute_input":"2024-09-27T20:47:40.346340Z","iopub.status.idle":"2024-09-27T20:47:40.359714Z","shell.execute_reply.started":"2024-09-27T20:47:40.346293Z","shell.execute_reply":"2024-09-27T20:47:40.358786Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport matplotlib.pyplot as plt\nfrom sklearn.manifold import TSNE\nfrom sklearn.cluster import KMeans\nimport seaborn as sns\nimport torch\nfrom sklearn.model_selection import KFold\nfrom torch.utils.data import DataLoader, Dataset, Subset\n\n\ndef analyze_expert_clusters(model, data_loader, expert_names, num_samples=600, n_clusters=5):\n    model.eval()\n    all_inputs = []\n    all_weights = []\n    \n    with torch.no_grad():\n        for inputs, _ in data_loader:\n            if len(all_inputs) * inputs.shape[0] >= num_samples:\n                break\n            _ = model(inputs)\n            weights = model.last_gating_weights\n            all_inputs.append(inputs.cpu().numpy())\n            all_weights.append(weights.cpu().numpy())\n    \n    # Concatenate all batches of inputs and weights\n    all_inputs = np.concatenate(all_inputs, axis=0)\n    all_weights = np.concatenate(all_weights, axis=0)\n    \n    # Ensure we have exactly num_samples\n    num_samples = min(num_samples, all_inputs.shape[0], all_weights.shape[0])\n    all_inputs = all_inputs[:num_samples]\n    all_weights = all_weights[:num_samples]\n    \n    # Clustering\n    kmeans = KMeans(n_clusters=n_clusters, n_init=10, random_state=42)\n    cluster_labels = kmeans.fit_predict(all_inputs)\n    centroids = kmeans.cluster_centers_\n    \n    # Combine input data and centroids for t-SNE\n    combined_data = np.vstack([all_inputs, centroids])\n    \n    # Dimensionality reduction\n    tsne = TSNE(n_components=2, random_state=42)\n    reduced_data = tsne.fit_transform(combined_data)\n    \n    # Split reduced data back into inputs and centroids\n    reduced_inputs = reduced_data[:num_samples]\n    reduced_centroids = reduced_data[num_samples:]\n    \n    # Ensure dominant_experts is also truncated to num_samples\n    dominant_experts = np.argmax(all_weights, axis=1)\n    dominant_experts = dominant_experts[:num_samples]\n    \n    # Plotting\n    plt.figure(figsize=(15, 10))\n    \n    # Scatter plot\n    scatter = plt.scatter(reduced_inputs[:, 0], reduced_inputs[:, 1], \n                          c=dominant_experts, cmap='viridis', \n                          alpha=0.6, s=50)\n    \n    # Add cluster centroids\n    plt.scatter(reduced_centroids[:, 0], reduced_centroids[:, 1], \n                marker='x', s=200, linewidths=3, color='r', label='Cluster Centroids')\n    \n    plt.colorbar(scatter, label='Dominant Expert')\n    plt.title('t-SNE Visualization of Input Data\\nColored by Dominant Expert with Cluster Centroids')\n    plt.xlabel('t-SNE Component 1')\n    plt.ylabel('t-SNE Component 2')\n    plt.legend()\n    plt.tight_layout()\n    plt.show()\n    \n    # Heatmap of expert dominance per cluster\n    expert_cluster_counts = np.zeros((n_clusters, len(expert_names)))\n    for cluster in range(n_clusters):\n        cluster_mask = cluster_labels == cluster\n        expert_cluster_counts[cluster] = np.sum(all_weights[cluster_mask], axis=0)\n    \n    expert_cluster_percentages = expert_cluster_counts / expert_cluster_counts.sum(axis=1, keepdims=True)\n    \n    plt.figure(figsize=(12, 8))\n    sns.heatmap(expert_cluster_percentages, annot=True, fmt='.2f', cmap='YlOrRd',\n                xticklabels=expert_names)\n    plt.title('Expert Dominance per Cluster')\n    plt.xlabel('Expert')\n    plt.ylabel('Cluster')\n    plt.tight_layout()\n    plt.show()\n#     return reduced_inputs, dominant_experts, cluster_labels","metadata":{"execution":{"iopub.status.busy":"2024-09-27T20:47:40.376492Z","iopub.execute_input":"2024-09-27T20:47:40.376811Z","iopub.status.idle":"2024-09-27T20:47:40.394339Z","shell.execute_reply.started":"2024-09-27T20:47:40.376777Z","shell.execute_reply":"2024-09-27T20:47:40.393287Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch.optim as optim\nfrom torch.autograd import Variable\n\nexpert_config = {\n    \"lstm_naive\": list(range(152)),\n    \"lstm_mol2vec\": list(range(152)) + list(range(152,452)),\n    \"lstm_rdkit\": list(range(152)) + list(range(452,456)),\n}\n\noutput_size = 18211\nepochs = 300\nfolds = 3\n\n# model = MixtureOfExperts(expert_config, output_size)\n# model.float()\n# model.to(device)\n\n# train_losses, val_losses = train_model(model, train_loader, val_loader, epochs)\n# models, train_losses_cv, val_losses_cv = cross_validate_model(MixtureOfExperts, expert_config, X, y, folds, epochs)\nmoe = MixtureOfExperts(expert_config, output_size)\nmoe.to(device)\n\nfor expert_name, expert in moe.experts.items():\n    print(f\"Pretraining {expert_name}\")\n    feature_indices = expert_config[expert_name]\n    train_losses, val_losses = pretrain_expert(expert, train_loader, val_loader, feature_indices, device, epochs=epochs)\n    plot_training_curves(train_losses, val_losses, f\"{expert_name} Pretraining\")\n    \n# Train gating network\nprint(\"Training gating network\")\ntrain_losses, val_losses = train_gating_network(moe, train_loader, val_loader, device, epochs=300)\nplot_training_curves(train_losses, val_losses, \"Gating Network Training\")\n\n# Visualize final gating weights\nplt.figure(figsize=(10, 6))\nplt.bar(range(len(moe.experts)), moe.last_gating_weights.cpu().mean(dim=0))\nplt.title(\"Average Gating Weights\")\nplt.xlabel(\"Expert\")\nplt.ylabel(\"Weight\")\nplt.xticks(range(len(moe.experts)), list(moe.experts.keys()), rotation=45)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-09-27T20:47:40.514102Z","iopub.execute_input":"2024-09-27T20:47:40.514474Z","iopub.status.idle":"2024-09-27T20:48:05.829597Z","shell.execute_reply.started":"2024-09-27T20:47:40.514440Z","shell.execute_reply":"2024-09-27T20:48:05.828553Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"expert_names = list(expert_config.keys())\n# model = models[-1]\nmodel = moe\nnum_samples = 50\nclusters = 6\n# Call the function with the expert names\n\n\nvisualize_expert_weights(model, train_loader, expert_names, num_samples)\nanalyze_expert_clusters(model, train_loader, expert_names, num_samples, clusters)\n\n# compute_feature_importance(model, X, feature_names, k=10)","metadata":{"execution":{"iopub.status.busy":"2024-09-27T20:48:05.831426Z","iopub.execute_input":"2024-09-27T20:48:05.831753Z","iopub.status.idle":"2024-09-27T20:48:08.359572Z","shell.execute_reply.started":"2024-09-27T20:48:05.831719Z","shell.execute_reply":"2024-09-27T20:48:08.358438Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"id_map = pd.read_csv(\"/kaggle/input/open-problems-single-cell-perturbations/id_map.csv\")","metadata":{"execution":{"iopub.status.busy":"2024-09-27T20:48:08.361347Z","iopub.execute_input":"2024-09-27T20:48:08.361645Z","iopub.status.idle":"2024-09-27T20:48:08.368270Z","shell.execute_reply.started":"2024-09-27T20:48:08.361613Z","shell.execute_reply":"2024-09-27T20:48:08.367281Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sm_name_to_smiles = df.set_index('sm_name')['SMILES'].to_dict()\nlen(sm_name_to_smiles)","metadata":{"execution":{"iopub.status.busy":"2024-09-27T20:48:08.370901Z","iopub.execute_input":"2024-09-27T20:48:08.371421Z","iopub.status.idle":"2024-09-27T20:48:08.420456Z","shell.execute_reply.started":"2024-09-27T20:48:08.371385Z","shell.execute_reply":"2024-09-27T20:48:08.419512Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"id_map['SMILES'] = id_map['sm_name'].map(sm_name_to_smiles)\nid_map.isna().sum()","metadata":{"execution":{"iopub.status.busy":"2024-09-27T20:48:08.421720Z","iopub.execute_input":"2024-09-27T20:48:08.422123Z","iopub.status.idle":"2024-09-27T20:48:08.434671Z","shell.execute_reply.started":"2024-09-27T20:48:08.422078Z","shell.execute_reply":"2024-09-27T20:48:08.433601Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_test = pd.DataFrame(id_map, columns=['cell_type', 'sm_name', 'SMILES'])\nprocessed_df_test = preprocessor_train.preprocess(df_test, fit=False)\nprint(\"Processed Testing DataFrame shape:\", processed_df_test.shape)\nprint(\"Testing Embedding indices:\", preprocessor_train.embedding_indices)","metadata":{"execution":{"iopub.status.busy":"2024-09-27T20:48:08.435760Z","iopub.execute_input":"2024-09-27T20:48:08.436085Z","iopub.status.idle":"2024-09-27T20:48:08.982517Z","shell.execute_reply.started":"2024-09-27T20:48:08.436018Z","shell.execute_reply":"2024-09-27T20:48:08.981538Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_array =  processed_df_test.to_numpy()","metadata":{"execution":{"iopub.status.busy":"2024-09-27T20:48:08.983745Z","iopub.execute_input":"2024-09-27T20:48:08.984094Z","iopub.status.idle":"2024-09-27T20:48:08.988811Z","shell.execute_reply.started":"2024-09-27T20:48:08.984054Z","shell.execute_reply":"2024-09-27T20:48:08.987856Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.to(device)\nmodel.eval()\ntarget_pred = model(torch.Tensor(test_array).to(device))","metadata":{"execution":{"iopub.status.busy":"2024-09-27T20:48:08.989972Z","iopub.execute_input":"2024-09-27T20:48:08.990323Z","iopub.status.idle":"2024-09-27T20:48:09.013368Z","shell.execute_reply.started":"2024-09-27T20:48:08.990290Z","shell.execute_reply":"2024-09-27T20:48:09.012518Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"res = target_pred.cpu()","metadata":{"execution":{"iopub.status.busy":"2024-09-27T20:48:09.014448Z","iopub.execute_input":"2024-09-27T20:48:09.014726Z","iopub.status.idle":"2024-09-27T20:48:09.027144Z","shell.execute_reply.started":"2024-09-27T20:48:09.014695Z","shell.execute_reply":"2024-09-27T20:48:09.026294Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample_submission = pd.read_csv(\"/kaggle/input/open-problems-single-cell-perturbations/sample_submission.csv\")","metadata":{"execution":{"iopub.status.busy":"2024-09-27T20:48:09.030101Z","iopub.execute_input":"2024-09-27T20:48:09.030434Z","iopub.status.idle":"2024-09-27T20:48:11.734688Z","shell.execute_reply.started":"2024-09-27T20:48:09.030397Z","shell.execute_reply":"2024-09-27T20:48:11.733646Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample_columns = sample_submission.columns\nsample_columns= sample_columns[1:]\nsubmission_df = pd.DataFrame(res.detach().numpy(), columns=sample_columns)","metadata":{"execution":{"iopub.status.busy":"2024-09-27T20:48:11.735947Z","iopub.execute_input":"2024-09-27T20:48:11.736295Z","iopub.status.idle":"2024-09-27T20:48:11.741800Z","shell.execute_reply.started":"2024-09-27T20:48:11.736261Z","shell.execute_reply":"2024-09-27T20:48:11.740804Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_df","metadata":{"execution":{"iopub.status.busy":"2024-09-27T20:48:11.742968Z","iopub.execute_input":"2024-09-27T20:48:11.743298Z","iopub.status.idle":"2024-09-27T20:48:11.832420Z","shell.execute_reply.started":"2024-09-27T20:48:11.743265Z","shell.execute_reply":"2024-09-27T20:48:11.831521Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_df.insert(0, 'id', range(255))","metadata":{"execution":{"iopub.status.busy":"2024-09-27T20:48:11.834088Z","iopub.execute_input":"2024-09-27T20:48:11.834643Z","iopub.status.idle":"2024-09-27T20:48:11.839805Z","shell.execute_reply.started":"2024-09-27T20:48:11.834597Z","shell.execute_reply":"2024-09-27T20:48:11.838932Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_df.to_csv(\"submission.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2024-09-27T20:48:11.840786Z","iopub.execute_input":"2024-09-27T20:48:11.841092Z","iopub.status.idle":"2024-09-27T20:48:20.724943Z","shell.execute_reply.started":"2024-09-27T20:48:11.841053Z","shell.execute_reply":"2024-09-27T20:48:20.723929Z"},"trusted":true},"execution_count":null,"outputs":[]}]}