{"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":119858,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":100806,"modelId":124987}],"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\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-24T19:47:12.164971Z","iopub.execute_input":"2024-09-24T19:47:12.165350Z","iopub.status.idle":"2024-09-24T19:47:17.462727Z","shell.execute_reply.started":"2024-09-24T19:47:12.165313Z","shell.execute_reply":"2024-09-24T19:47:17.461804Z"},"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-24T19:47:22.818357Z","iopub.execute_input":"2024-09-24T19:47:22.819236Z","iopub.status.idle":"2024-09-24T19:47:22.856186Z","shell.execute_reply.started":"2024-09-24T19:47:22.819197Z","shell.execute_reply":"2024-09-24T19:47:22.855266Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Installing a package\n!pip install git+https://github.com/samoturk/mol2vec;","metadata":{"execution":{"iopub.status.busy":"2024-09-24T19:47:23.314452Z","iopub.execute_input":"2024-09-24T19:47:23.315205Z","iopub.status.idle":"2024-09-24T19:47:52.553354Z","shell.execute_reply.started":"2024-09-24T19:47:23.315164Z","shell.execute_reply":"2024-09-24T19:47:52.552255Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install rdkit","metadata":{"execution":{"iopub.status.busy":"2024-09-24T19:47:52.555359Z","iopub.execute_input":"2024-09-24T19:47:52.555739Z","iopub.status.idle":"2024-09-24T19:48:06.045568Z","shell.execute_reply.started":"2024-09-24T19:47:52.555704Z","shell.execute_reply":"2024-09-24T19:48:06.044414Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from gensim.models import word2vec\nmodel1 = word2vec.Word2Vec.load('/kaggle/input/mol2vec/pytorch/default/1/model_300dim.pkl')","metadata":{"execution":{"iopub.status.busy":"2024-09-24T19:48:06.047019Z","iopub.execute_input":"2024-09-24T19:48:06.047382Z","iopub.status.idle":"2024-09-24T19:48:19.347109Z","shell.execute_reply.started":"2024-09-24T19:48:06.047328Z","shell.execute_reply":"2024-09-24T19:48:19.345968Z"},"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()\n","metadata":{"execution":{"iopub.status.busy":"2024-09-24T19:48:19.349510Z","iopub.execute_input":"2024-09-24T19:48:19.349830Z","iopub.status.idle":"2024-09-24T19:48:21.897698Z","shell.execute_reply.started":"2024-09-24T19:48:19.349796Z","shell.execute_reply":"2024-09-24T19:48:21.896273Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"feature_cols =['cell_type','sm_name']\ntarget_cols = ['cell_type','sm_name','sm_lincs_id','SMILES','control']\ntargets = df.drop(columns=target_cols)\nfeatures = pd.DataFrame(df,columns=feature_cols)\nfeatures","metadata":{"execution":{"iopub.status.busy":"2024-09-24T19:48:21.899321Z","iopub.execute_input":"2024-09-24T19:48:21.899843Z","iopub.status.idle":"2024-09-24T19:48:21.956830Z","shell.execute_reply.started":"2024-09-24T19:48:21.899783Z","shell.execute_reply":"2024-09-24T19:48:21.955083Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"encoder = OneHotEncoder(sparse=False)\nencoded_features = encoder.fit_transform(df[['cell_type', 'sm_name']])\nencoded_df = pd.DataFrame(encoded_features, columns=encoder.get_feature_names_out(['cell_type', 'sm_name']))","metadata":{"execution":{"iopub.status.busy":"2024-09-24T19:48:21.958074Z","iopub.execute_input":"2024-09-24T19:48:21.958385Z","iopub.status.idle":"2024-09-24T19:48:21.973797Z","shell.execute_reply.started":"2024-09-24T19:48:21.958352Z","shell.execute_reply":"2024-09-24T19:48:21.972536Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\nfrom rdkit import Chem\nfrom rdkit.Chem import Descriptors, rdMolDescriptors, QED\nfrom rdkit.Chem.rdMolDescriptors import CalcTPSA, CalcNumRotatableBonds, CalcNumHBA, CalcNumHBD, CalcFractionCSP3\nfrom rdkit.Chem import BRICS, Recap\n\n\ndef 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#     info['Molecular Formula'] = Chem.rdMolDescriptors.CalcMolFormula(mol)\n    info['Molecular Weight'] = Descriptors.MolWt(mol)\n    info['LogP'] = Descriptors.MolLogP(mol)\n    info['Number of Atoms'] = mol.GetNumAtoms()\n    info['Number of Bonds'] = mol.GetNumBonds()\n    return info","metadata":{"execution":{"iopub.status.busy":"2024-09-24T19:48:21.975460Z","iopub.execute_input":"2024-09-24T19:48:21.975962Z","iopub.status.idle":"2024-09-24T19:48:22.366042Z","shell.execute_reply.started":"2024-09-24T19:48:21.975906Z","shell.execute_reply":"2024-09-24T19:48:22.364975Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"smiles_info_list = df['SMILES'].apply(extract_smiles_info)\nsmiles_info_df = pd.DataFrame(smiles_info_list.tolist())","metadata":{"execution":{"iopub.status.busy":"2024-09-24T19:48:22.367351Z","iopub.execute_input":"2024-09-24T19:48:22.367763Z","iopub.status.idle":"2024-09-24T19:48:22.940658Z","shell.execute_reply.started":"2024-09-24T19:48:22.367706Z","shell.execute_reply":"2024-09-24T19:48:22.939737Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"scaler = StandardScaler()\nsmiles_info_df[['Molecular Weight', 'LogP', 'Number of Atoms', 'Number of Bonds']] = scaler.fit_transform(smiles_info_df[['Molecular Weight', 'LogP', 'Number of Atoms', 'Number of Bonds']])","metadata":{"execution":{"iopub.status.busy":"2024-09-24T19:48:22.941843Z","iopub.execute_input":"2024-09-24T19:48:22.942161Z","iopub.status.idle":"2024-09-24T19:48:22.953777Z","shell.execute_reply.started":"2024-09-24T19:48:22.942120Z","shell.execute_reply":"2024-09-24T19:48:22.952800Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Importing Chem module\nfrom rdkit import Chem ","metadata":{"execution":{"iopub.status.busy":"2024-09-24T19:48:22.958543Z","iopub.execute_input":"2024-09-24T19:48:22.959029Z","iopub.status.idle":"2024-09-24T19:48:22.999666Z","shell.execute_reply.started":"2024-09-24T19:48:22.958973Z","shell.execute_reply":"2024-09-24T19:48:22.998598Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df['mol'] = df['SMILES'].apply(lambda x: Chem.MolFromSmiles(x))","metadata":{"execution":{"iopub.status.busy":"2024-09-24T19:48:23.000748Z","iopub.execute_input":"2024-09-24T19:48:23.001025Z","iopub.status.idle":"2024-09-24T19:48:23.204970Z","shell.execute_reply.started":"2024-09-24T19:48:23.000994Z","shell.execute_reply":"2024-09-24T19:48:23.204212Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from mol2vec.features import mol2alt_sentence, mol2sentence, MolSentence, DfVec, sentences2vec\nfrom gensim.models import word2vec\n\ndf['sentence'] = df.apply(lambda x: MolSentence(mol2alt_sentence(x['mol'], 1)), axis=1)","metadata":{"execution":{"iopub.status.busy":"2024-09-24T19:48:23.206235Z","iopub.execute_input":"2024-09-24T19:48:23.206588Z","iopub.status.idle":"2024-09-24T19:48:24.537203Z","shell.execute_reply.started":"2024-09-24T19:48:23.206554Z","shell.execute_reply":"2024-09-24T19:48:24.536392Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\n\n# Assuming 'model' is your pre-trained Word2Vec model and 'keys' is the set of valid words in your model's vocabulary\nkeys = set(model1.wv.key_to_index.keys())\n\ndef sentence_to_vector(sentence, model, keys, unseen=False, unseen_vec=np.zeros(300)):\n    if unseen:\n        vec = sum([\n            model.wv.get_vector(word) if word in keys else unseen_vec\n            for word in sentence\n        ])\n    else:\n        vec = sum([\n            model.wv.get_vector(word)\n            for word in sentence\n            if word in keys\n        ])\n    return vec\n\n# Apply the function to each row in your DataFrame\ndf['vector'] = df['sentence'].apply(lambda sentence: sentence_to_vector(sentence, model1, keys))\n\n# If you want to create individual columns for each dimension of the vector\nvector_dim = len(model1.wv.get_vector(next(iter(keys))))  # assuming all vectors have the same dimension\nvector_columns = [f'vector_{i}' for i in range(vector_dim)]\n\n# Split the vector into multiple columns\ndf[vector_columns] = pd.DataFrame(df['vector'].tolist(), index=df.index)\n\n# Now, df will have the sentence vectors spread across new columns","metadata":{"execution":{"iopub.status.busy":"2024-09-24T19:48:24.538349Z","iopub.execute_input":"2024-09-24T19:48:24.539002Z","iopub.status.idle":"2024-09-24T19:48:25.322400Z","shell.execute_reply.started":"2024-09-24T19:48:24.538967Z","shell.execute_reply":"2024-09-24T19:48:25.321423Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.head()","metadata":{"execution":{"iopub.status.busy":"2024-09-24T19:48:51.771452Z","iopub.execute_input":"2024-09-24T19:48:51.771888Z","iopub.status.idle":"2024-09-24T19:48:51.835255Z","shell.execute_reply.started":"2024-09-24T19:48:51.771851Z","shell.execute_reply":"2024-09-24T19:48:51.834369Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mol2vec_df = df[vector_columns]","metadata":{"execution":{"iopub.status.busy":"2024-09-24T19:48:52.385534Z","iopub.execute_input":"2024-09-24T19:48:52.385919Z","iopub.status.idle":"2024-09-24T19:48:52.403262Z","shell.execute_reply.started":"2024-09-24T19:48:52.385882Z","shell.execute_reply":"2024-09-24T19:48:52.402156Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = pd.concat([encoded_df,mol2vec_df,smiles_info_df], axis=1)","metadata":{"execution":{"iopub.status.busy":"2024-09-24T19:48:53.161847Z","iopub.execute_input":"2024-09-24T19:48:53.162608Z","iopub.status.idle":"2024-09-24T19:48:53.175205Z","shell.execute_reply.started":"2024-09-24T19:48:53.162549Z","shell.execute_reply":"2024-09-24T19:48:53.174121Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# train_df.drop(columns = ['Molecular Formula'])","metadata":{"execution":{"iopub.status.busy":"2024-09-24T19:48:53.858985Z","iopub.execute_input":"2024-09-24T19:48:53.859679Z","iopub.status.idle":"2024-09-24T19:48:53.863737Z","shell.execute_reply.started":"2024-09-24T19:48:53.859637Z","shell.execute_reply":"2024-09-24T19:48:53.862809Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_array = train_df.to_numpy()","metadata":{"execution":{"iopub.status.busy":"2024-09-24T19:48:54.549991Z","iopub.execute_input":"2024-09-24T19:48:54.550389Z","iopub.status.idle":"2024-09-24T19:48:54.555352Z","shell.execute_reply.started":"2024-09-24T19:48:54.550353Z","shell.execute_reply":"2024-09-24T19:48:54.554268Z"},"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)","metadata":{"execution":{"iopub.status.busy":"2024-09-24T19:48:57.851057Z","iopub.execute_input":"2024-09-24T19:48:57.851454Z","iopub.status.idle":"2024-09-24T19:48:57.897850Z","shell.execute_reply.started":"2024-09-24T19:48:57.851420Z","shell.execute_reply":"2024-09-24T19:48:57.896971Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"X.shape\n# y.shape","metadata":{"execution":{"iopub.status.busy":"2024-09-24T19:48:58.599877Z","iopub.execute_input":"2024-09-24T19:48:58.600259Z","iopub.status.idle":"2024-09-24T19:48:58.606366Z","shell.execute_reply.started":"2024-09-24T19:48:58.600222Z","shell.execute_reply":"2024-09-24T19:48:58.605532Z"},"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-24T19:49:02.988207Z","iopub.execute_input":"2024-09-24T19:49:02.988617Z","iopub.status.idle":"2024-09-24T19:49:03.171858Z","shell.execute_reply.started":"2024-09-24T19:49:02.988580Z","shell.execute_reply":"2024-09-24T19:49:03.170924Z"},"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(123)\nnp.random.seed(123)\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        self.network = nn.Sequential(\n            nn.Linear(input_size, 64),\n            nn.ReLU(),\n            nn.Linear(64, 128),\n            nn.ReLU(),\n            nn.Linear(128, num_experts),\n            nn.Softmax(dim=-1)\n        )\n\n    def forward(self, data):\n        return self.network(data)\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):\n        self.patience = patience\n        self.min_delta = min_delta\n        self.best_loss = None\n        self.counter = 0\n\n    def __call__(self, val_loss):\n        if self.best_loss is None:\n            self.best_loss = val_loss\n        elif val_loss > self.best_loss - self.min_delta:\n            self.counter += 1\n            if self.counter >= self.patience:\n                return True\n        else:\n            self.best_loss = val_loss\n            self.counter = 0\n        return False","metadata":{"execution":{"iopub.status.busy":"2024-09-24T19:49:04.077353Z","iopub.execute_input":"2024-09-24T19:49:04.078266Z","iopub.status.idle":"2024-09-24T19:49:04.183400Z","shell.execute_reply.started":"2024-09-24T19:49:04.078227Z","shell.execute_reply":"2024-09-24T19:49:04.182655Z"},"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#         if early_stopping(val_loss / len(val_loader)):\n#              print(\"Early stopping triggered\")\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-24T19:57:33.224948Z","iopub.execute_input":"2024-09-24T19:57:33.225863Z","iopub.status.idle":"2024-09-24T19:57:33.237617Z","shell.execute_reply.started":"2024-09-24T19:57:33.225822Z","shell.execute_reply":"2024-09-24T19:57:33.236665Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def 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    # Ensure all loss lists have the same length by padding with np.nan\n    max_epochs = min(len(losses) for losses in train_losses_cv)  # Use min if you want to truncate to shortest\n    train_losses_cv = [losses[:max_epochs] for losses in train_losses_cv]  # Truncate to minimum length\n    val_losses_cv = [losses[:max_epochs] for losses in val_losses_cv]\n\n    # Convert to numpy arrays for averaging\n    avg_train_loss = np.mean(np.array(train_losses_cv), axis=0)\n    avg_val_loss = np.mean(np.array(val_losses_cv), axis=0)\n    \n    # Plotting average cross-validation losses\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","metadata":{"execution":{"iopub.status.busy":"2024-09-24T19:57:34.357771Z","iopub.execute_input":"2024-09-24T19:57:34.358452Z","iopub.status.idle":"2024-09-24T19:57:34.370174Z","shell.execute_reply.started":"2024-09-24T19:57:34.358411Z","shell.execute_reply":"2024-09-24T19:57:34.369299Z"},"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-24T19:57:34.571811Z","iopub.execute_input":"2024-09-24T19:57:34.572163Z","iopub.status.idle":"2024-09-24T19:57:34.582426Z","shell.execute_reply.started":"2024-09-24T19:57:34.572128Z","shell.execute_reply":"2024-09-24T19:57:34.581537Z"},"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=1000, 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-24T19:57:34.747528Z","iopub.execute_input":"2024-09-24T19:57:34.747894Z","iopub.status.idle":"2024-09-24T19:57:34.765793Z","shell.execute_reply.started":"2024-09-24T19:57:34.747860Z","shell.execute_reply":"2024-09-24T19:57:34.764901Z"},"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_expert_1\": list(range(152)),\n    \"lstm_expert_2\": list(range(152)) + list(range(152,452)),\n    \"lstm_expert_3\": list(range(152)) + list(range(452,456)),\n}\n\noutput_size = 18211\nepochs = 100\n\nmodel = MixtureOfExperts(expert_config, output_size)\nmodel.float()\nmodel.to(device)\n\n# train_losses, val_losses = train_model(model, train_loader, val_loader, epochs)\ncross_validate_model(MixtureOfExperts,expert_config, X, y, n_splits=5, epochs=100, lr=0.001)","metadata":{"execution":{"iopub.status.busy":"2024-09-24T19:57:34.940848Z","iopub.execute_input":"2024-09-24T19:57:34.941550Z","iopub.status.idle":"2024-09-24T20:01:32.060510Z","shell.execute_reply.started":"2024-09-24T19:57:34.941514Z","shell.execute_reply":"2024-09-24T20:01:32.059546Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"expert_names = list(expert_config.keys())\n\nnum_samples = 100\nclusters = 5\n# Call the function with the expert names\nvisualize_expert_weights(model, train_loader, expert_names, num_samples)\nanalyze_expert_clusters(model, train_loader, expert_names, num_samples, clusters)","metadata":{"execution":{"iopub.status.busy":"2024-09-24T20:01:32.062485Z","iopub.execute_input":"2024-09-24T20:01:32.062964Z","iopub.status.idle":"2024-09-24T20:01:34.495600Z","shell.execute_reply.started":"2024-09-24T20:01:32.062919Z","shell.execute_reply":"2024-09-24T20:01:34.494634Z"},"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-24T20:04:05.743042Z","iopub.execute_input":"2024-09-24T20:04:05.743940Z","iopub.status.idle":"2024-09-24T20:04:05.774143Z","shell.execute_reply.started":"2024-09-24T20:04:05.743897Z","shell.execute_reply":"2024-09-24T20:04:05.773342Z"},"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-24T20:04:06.426902Z","iopub.execute_input":"2024-09-24T20:04:06.427780Z","iopub.status.idle":"2024-09-24T20:04:06.477627Z","shell.execute_reply.started":"2024-09-24T20:04:06.427740Z","shell.execute_reply":"2024-09-24T20:04:06.476739Z"},"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-24T20:04:07.064505Z","iopub.execute_input":"2024-09-24T20:04:07.065364Z","iopub.status.idle":"2024-09-24T20:04:07.075201Z","shell.execute_reply.started":"2024-09-24T20:04:07.065326Z","shell.execute_reply":"2024-09-24T20:04:07.074193Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"testdata = pd.DataFrame(id_map, columns=['cell_type', 'sm_name', 'SMILES'])\ntestdata","metadata":{"execution":{"iopub.status.busy":"2024-09-24T20:04:07.682406Z","iopub.execute_input":"2024-09-24T20:04:07.683139Z","iopub.status.idle":"2024-09-24T20:04:07.707690Z","shell.execute_reply.started":"2024-09-24T20:04:07.683099Z","shell.execute_reply":"2024-09-24T20:04:07.706682Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"testdata['mol'] = testdata['SMILES'].apply(lambda x: Chem.MolFromSmiles(x))","metadata":{"execution":{"iopub.status.busy":"2024-09-24T20:04:08.209198Z","iopub.execute_input":"2024-09-24T20:04:08.210101Z","iopub.status.idle":"2024-09-24T20:04:08.298694Z","shell.execute_reply.started":"2024-09-24T20:04:08.210059Z","shell.execute_reply":"2024-09-24T20:04:08.297662Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"testdata['sentence'] = testdata.apply(lambda x: MolSentence(mol2alt_sentence(x['mol'], 1)), axis=1)","metadata":{"execution":{"iopub.status.busy":"2024-09-24T20:04:08.672945Z","iopub.execute_input":"2024-09-24T20:04:08.673873Z","iopub.status.idle":"2024-09-24T20:04:08.778923Z","shell.execute_reply.started":"2024-09-24T20:04:08.673831Z","shell.execute_reply":"2024-09-24T20:04:08.777863Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Apply the function to each row in your DataFrame\ntestdata['vector'] = testdata['sentence'].apply(lambda sentence: sentence_to_vector(sentence, model1, keys))\n\n# If you want to create individual columns for each dimension of the vector\nvector_dim = len(model1.wv.get_vector(next(iter(keys))))  # assuming all vectors have the same dimension\nvector_columns = [f'vector_{i}' for i in range(vector_dim)]\n\n# Split the vector into multiple columns\ntestdata[vector_columns] = pd.DataFrame(testdata['vector'].tolist(), index=testdata.index)","metadata":{"execution":{"iopub.status.busy":"2024-09-24T20:04:08.868164Z","iopub.execute_input":"2024-09-24T20:04:08.868979Z","iopub.status.idle":"2024-09-24T20:04:09.081743Z","shell.execute_reply.started":"2024-09-24T20:04:08.868941Z","shell.execute_reply":"2024-09-24T20:04:09.080837Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"testdata","metadata":{"execution":{"iopub.status.busy":"2024-09-24T20:04:09.138917Z","iopub.execute_input":"2024-09-24T20:04:09.139732Z","iopub.status.idle":"2024-09-24T20:04:09.239950Z","shell.execute_reply.started":"2024-09-24T20:04:09.139693Z","shell.execute_reply":"2024-09-24T20:04:09.239034Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_smiles_info = testdata['SMILES'].apply(extract_smiles_info)\ntest_smiles_info_df = pd.DataFrame(test_smiles_info.tolist())\ntest_smiles_info_df[['Molecular Weight', 'LogP', 'Number of Atoms', 'Number of Bonds']] = scaler.transform(test_smiles_info_df[['Molecular Weight', 'LogP', 'Number of Atoms', 'Number of Bonds']])","metadata":{"execution":{"iopub.status.busy":"2024-09-24T20:04:09.252846Z","iopub.execute_input":"2024-09-24T20:04:09.253153Z","iopub.status.idle":"2024-09-24T20:04:09.483298Z","shell.execute_reply.started":"2024-09-24T20:04:09.253120Z","shell.execute_reply":"2024-09-24T20:04:09.482332Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_mol2vec = testdata[vector_columns]","metadata":{"execution":{"iopub.status.busy":"2024-09-24T20:04:09.485087Z","iopub.execute_input":"2024-09-24T20:04:09.485821Z","iopub.status.idle":"2024-09-24T20:04:09.499820Z","shell.execute_reply.started":"2024-09-24T20:04:09.485775Z","shell.execute_reply":"2024-09-24T20:04:09.498932Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_encoded_features = encoder.transform(testdata[['cell_type', 'sm_name']])\ntest_encoded_features_df = pd.DataFrame(test_encoded_features, columns=encoder.get_feature_names_out(['cell_type', 'sm_name']))","metadata":{"execution":{"iopub.status.busy":"2024-09-24T20:04:09.692634Z","iopub.execute_input":"2024-09-24T20:04:09.693231Z","iopub.status.idle":"2024-09-24T20:04:09.701126Z","shell.execute_reply.started":"2024-09-24T20:04:09.693189Z","shell.execute_reply":"2024-09-24T20:04:09.700214Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df = pd.concat([test_encoded_features_df,test_mol2vec,test_smiles_info_df], axis=1)","metadata":{"execution":{"iopub.status.busy":"2024-09-24T20:04:09.733841Z","iopub.execute_input":"2024-09-24T20:04:09.734139Z","iopub.status.idle":"2024-09-24T20:04:09.745868Z","shell.execute_reply.started":"2024-09-24T20:04:09.734107Z","shell.execute_reply":"2024-09-24T20:04:09.744932Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df","metadata":{"execution":{"iopub.status.busy":"2024-09-24T20:04:10.010005Z","iopub.execute_input":"2024-09-24T20:04:10.010871Z","iopub.status.idle":"2024-09-24T20:04:10.101925Z","shell.execute_reply.started":"2024-09-24T20:04:10.010830Z","shell.execute_reply":"2024-09-24T20:04:10.101043Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_array = test_df.to_numpy()\n","metadata":{"execution":{"iopub.status.busy":"2024-09-24T20:04:10.103544Z","iopub.execute_input":"2024-09-24T20:04:10.103853Z","iopub.status.idle":"2024-09-24T20:04:10.108506Z","shell.execute_reply.started":"2024-09-24T20:04:10.103820Z","shell.execute_reply":"2024-09-24T20:04:10.107420Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"target_pred = model(torch.Tensor(test_array).to(device))","metadata":{"execution":{"iopub.status.busy":"2024-09-24T20:04:10.251801Z","iopub.execute_input":"2024-09-24T20:04:10.252667Z","iopub.status.idle":"2024-09-24T20:04:10.278693Z","shell.execute_reply.started":"2024-09-24T20:04:10.252628Z","shell.execute_reply":"2024-09-24T20:04:10.277843Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"res = target_pred.cpu()","metadata":{"execution":{"iopub.status.busy":"2024-09-24T20:04:10.461264Z","iopub.execute_input":"2024-09-24T20:04:10.461997Z","iopub.status.idle":"2024-09-24T20:04:10.478842Z","shell.execute_reply.started":"2024-09-24T20:04:10.461959Z","shell.execute_reply":"2024-09-24T20:04:10.478052Z"},"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-24T20:04:10.591956Z","iopub.execute_input":"2024-09-24T20:04:10.592774Z","iopub.status.idle":"2024-09-24T20:04:13.227475Z","shell.execute_reply.started":"2024-09-24T20:04:10.592736Z","shell.execute_reply":"2024-09-24T20:04:13.226604Z"},"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-24T20:04:13.229102Z","iopub.execute_input":"2024-09-24T20:04:13.229405Z","iopub.status.idle":"2024-09-24T20:04:13.234538Z","shell.execute_reply.started":"2024-09-24T20:04:13.229372Z","shell.execute_reply":"2024-09-24T20:04:13.233638Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_df","metadata":{"execution":{"iopub.status.busy":"2024-09-24T20:04:13.235686Z","iopub.execute_input":"2024-09-24T20:04:13.236119Z","iopub.status.idle":"2024-09-24T20:04:13.321181Z","shell.execute_reply.started":"2024-09-24T20:04:13.236074Z","shell.execute_reply":"2024-09-24T20:04:13.320212Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_df.insert(0, 'id', range(255))","metadata":{"execution":{"iopub.status.busy":"2024-09-24T20:04:13.322875Z","iopub.execute_input":"2024-09-24T20:04:13.323167Z","iopub.status.idle":"2024-09-24T20:04:13.328113Z","shell.execute_reply.started":"2024-09-24T20:04:13.323136Z","shell.execute_reply":"2024-09-24T20:04:13.327127Z"},"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-24T20:04:13.329225Z","iopub.execute_input":"2024-09-24T20:04:13.329522Z","iopub.status.idle":"2024-09-24T20:04:22.156374Z","shell.execute_reply.started":"2024-09-24T20:04:13.329491Z","shell.execute_reply":"2024-09-24T20:04:22.155538Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}