{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","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"}],"dockerImageVersionId":30699,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# **OPxMOE**","metadata":{}},{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-09-24T20:24:24.193251Z","iopub.execute_input":"2024-09-24T20:24:24.193672Z","iopub.status.idle":"2024-09-24T20:24:24.201306Z","shell.execute_reply.started":"2024-09-24T20:24:24.193638Z","shell.execute_reply":"2024-09-24T20:24:24.200313Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install -qU captum","metadata":{"execution":{"iopub.status.busy":"2024-09-24T20:24:24.206750Z","iopub.execute_input":"2024-09-24T20:24:24.207114Z","iopub.status.idle":"2024-09-24T20:24:36.696510Z","shell.execute_reply.started":"2024-09-24T20:24:24.207087Z","shell.execute_reply":"2024-09-24T20:24:36.695217Z"},"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-24T20:24:36.698690Z","iopub.execute_input":"2024-09-24T20:24:36.698980Z","iopub.status.idle":"2024-09-24T20:24:37.926534Z","shell.execute_reply.started":"2024-09-24T20:24:36.698952Z","shell.execute_reply":"2024-09-24T20:24:37.925427Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## **Data Preprocessing**","metadata":{}},{"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-24T20:24:37.927735Z","iopub.execute_input":"2024-09-24T20:24:37.928016Z","iopub.status.idle":"2024-09-24T20:24:37.973479Z","shell.execute_reply.started":"2024-09-24T20:24:37.927992Z","shell.execute_reply":"2024-09-24T20:24:37.972563Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.preprocessing import OneHotEncoder\none_hot = OneHotEncoder()\nfeatures_array = one_hot.fit_transform(features)","metadata":{"execution":{"iopub.status.busy":"2024-09-24T20:24:37.975645Z","iopub.execute_input":"2024-09-24T20:24:37.975950Z","iopub.status.idle":"2024-09-24T20:24:37.984280Z","shell.execute_reply.started":"2024-09-24T20:24:37.975924Z","shell.execute_reply":"2024-09-24T20:24:37.983369Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"features_one_hot=features_array\nfeature_names = one_hot.get_feature_names_out()\nprint(len(feature_names))\nprint(features_one_hot.toarray().shape)","metadata":{"execution":{"iopub.status.busy":"2024-09-24T20:24:37.985452Z","iopub.execute_input":"2024-09-24T20:24:37.985790Z","iopub.status.idle":"2024-09-24T20:24:37.994950Z","shell.execute_reply.started":"2024-09-24T20:24:37.985755Z","shell.execute_reply":"2024-09-24T20:24:37.994050Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### **Getting Feature indices**","metadata":{}},{"cell_type":"code","source":"features_to_find = feature_names\nindices = [list(feature_names).index(feature) for feature in features_to_find]\n\nprint(\"Indices for selected features:\", indices)","metadata":{"execution":{"iopub.status.busy":"2024-09-24T20:24:37.995927Z","iopub.execute_input":"2024-09-24T20:24:37.996190Z","iopub.status.idle":"2024-09-24T20:24:38.004863Z","shell.execute_reply.started":"2024-09-24T20:24:37.996154Z","shell.execute_reply":"2024-09-24T20:24:38.003997Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nX = torch.tensor(features_one_hot.toarray(), dtype=torch.float32)\ny = torch.tensor(targets.values, dtype=torch.float32)","metadata":{"execution":{"iopub.status.busy":"2024-09-24T20:24:38.006056Z","iopub.execute_input":"2024-09-24T20:24:38.006465Z","iopub.status.idle":"2024-09-24T20:24:38.032557Z","shell.execute_reply.started":"2024-09-24T20:24:38.006438Z","shell.execute_reply":"2024-09-24T20:24:38.031735Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"featurespace = features_one_hot.toarray()\ntargetsspace = targets.values\ntargetsspace.shape","metadata":{"execution":{"iopub.status.busy":"2024-09-24T20:24:38.033725Z","iopub.execute_input":"2024-09-24T20:24:38.034004Z","iopub.status.idle":"2024-09-24T20:24:38.040815Z","shell.execute_reply.started":"2024-09-24T20:24:38.033981Z","shell.execute_reply":"2024-09-24T20:24:38.039917Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(f\"X shape: {X.shape}, y shape: {y.shape}, indices shape: {len(indices)}\")","metadata":{"execution":{"iopub.status.busy":"2024-09-24T20:24:38.043864Z","iopub.execute_input":"2024-09-24T20:24:38.044189Z","iopub.status.idle":"2024-09-24T20:24:38.049032Z","shell.execute_reply.started":"2024-09-24T20:24:38.044149Z","shell.execute_reply":"2024-09-24T20:24:38.048223Z"},"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-24T20:24:38.052732Z","iopub.execute_input":"2024-09-24T20:24:38.052977Z","iopub.status.idle":"2024-09-24T20:24:38.073860Z","shell.execute_reply.started":"2024-09-24T20:24:38.052956Z","shell.execute_reply":"2024-09-24T20:24:38.073207Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Mixture of Experts Model**","metadata":{}},{"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(3)\nnp.random.seed(3)\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-24T20:58:10.867559Z","iopub.execute_input":"2024-09-24T20:58:10.868220Z","iopub.status.idle":"2024-09-24T20:58:10.896675Z","shell.execute_reply.started":"2024-09-24T20:58:10.868188Z","shell.execute_reply":"2024-09-24T20:58:10.895407Z"},"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 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-24T20:58:14.214278Z","iopub.execute_input":"2024-09-24T20:58:14.214793Z","iopub.status.idle":"2024-09-24T20:58:14.227162Z","shell.execute_reply.started":"2024-09-24T20:58:14.214755Z","shell.execute_reply":"2024-09-24T20:58:14.226024Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.model_selection import KFold\nfrom sklearn.metrics import mean_squared_error\nimport numpy as np\nimport torch.optim as optim\nimport matplotlib.pyplot as plt\n\nimport torch\nfrom torch.utils.data import DataLoader, Dataset, Subset\nfrom sklearn.model_selection import KFold\nfrom captum.attr import IntegratedGradients\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\n\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom captum.attr import IntegratedGradients, Saliency, DeepLift, ShapleyValueSampling\n\ndef compute_feature_importance(model, X, feature_names, k=10, baseline=None, target=0, methods=('integrated_gradients', 'saliency', 'deeplift')):\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    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-24T20:58:40.345981Z","iopub.execute_input":"2024-09-24T20:58:40.346870Z","iopub.status.idle":"2024-09-24T20:58:40.369396Z","shell.execute_reply.started":"2024-09-24T20:58:40.346836Z","shell.execute_reply":"2024-09-24T20:58:40.368361Z"},"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-24T20:58:43.546614Z","iopub.execute_input":"2024-09-24T20:58:43.546972Z","iopub.status.idle":"2024-09-24T20:58:43.557733Z","shell.execute_reply.started":"2024-09-24T20:58:43.546943Z","shell.execute_reply":"2024-09-24T20:58:43.556555Z"},"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\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-24T20:58:45.754912Z","iopub.execute_input":"2024-09-24T20:58:45.755565Z","iopub.status.idle":"2024-09-24T20:58:45.772922Z","shell.execute_reply.started":"2024-09-24T20:58:45.755532Z","shell.execute_reply":"2024-09-24T20:58:45.771996Z"},"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\": indices[:50],\n    \"lstm_expert_2\": indices[50:100],\n    \"lstm_expert_3\": indices[100:],\n}\n\noutput_size = 18211\nepochs = 50\nfolds = 5\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\nmodels, train_losses_cv, val_losses_cv = cross_validate_model(MixtureOfExperts, expert_config, X, y, folds, epochs)","metadata":{"execution":{"iopub.status.busy":"2024-09-24T21:00:08.238951Z","iopub.execute_input":"2024-09-24T21:00:08.239383Z","iopub.status.idle":"2024-09-24T21:01:52.558657Z","shell.execute_reply.started":"2024-09-24T21:00:08.239352Z","shell.execute_reply":"2024-09-24T21:01:52.557702Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"expert_names = list(expert_config.keys())\nmodel = models[-1]\nnum_samples = 100\nclusters = 5\n# Call the function with the expert names\n\nvisualize_expert_weights(model, train_loader, expert_names, num_samples)\nanalyze_expert_clusters(model, train_loader, expert_names, num_samples, clusters)\n\ncompute_feature_importance(model, X, feature_names, k=10)","metadata":{"execution":{"iopub.status.busy":"2024-09-24T21:02:00.261190Z","iopub.execute_input":"2024-09-24T21:02:00.261874Z","iopub.status.idle":"2024-09-24T21:02:41.540287Z","shell.execute_reply.started":"2024-09-24T21:02:00.261842Z","shell.execute_reply":"2024-09-24T21:02:41.539137Z"},"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-24T21:02:55.836053Z","iopub.execute_input":"2024-09-24T21:02:55.836798Z","iopub.status.idle":"2024-09-24T21:02:55.851923Z","shell.execute_reply.started":"2024-09-24T21:02:55.836764Z","shell.execute_reply":"2024-09-24T21:02:55.851161Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"testdata = pd.DataFrame(id_map,columns=feature_cols) \n#testdata.head()","metadata":{"execution":{"iopub.status.busy":"2024-09-24T21:02:58.265786Z","iopub.execute_input":"2024-09-24T21:02:58.266297Z","iopub.status.idle":"2024-09-24T21:02:58.276990Z","shell.execute_reply.started":"2024-09-24T21:02:58.266253Z","shell.execute_reply":"2024-09-24T21:02:58.275928Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"one_hot_test = one_hot.transform(testdata)\nprint(features_one_hot.toarray().shape,one_hot_test.toarray().shape)","metadata":{"execution":{"iopub.status.busy":"2024-09-24T21:03:01.125652Z","iopub.execute_input":"2024-09-24T21:03:01.126533Z","iopub.status.idle":"2024-09-24T21:03:01.137189Z","shell.execute_reply.started":"2024-09-24T21:03:01.126486Z","shell.execute_reply":"2024-09-24T21:03:01.136098Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.to(device)\nmodel.eval()\ntarget_pred = model(torch.Tensor(one_hot_test.toarray()).to(device))","metadata":{"execution":{"iopub.status.busy":"2024-09-24T21:03:04.085556Z","iopub.execute_input":"2024-09-24T21:03:04.086186Z","iopub.status.idle":"2024-09-24T21:03:04.182984Z","shell.execute_reply.started":"2024-09-24T21:03:04.086145Z","shell.execute_reply":"2024-09-24T21:03:04.182148Z"},"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-24T21:03:06.684110Z","iopub.execute_input":"2024-09-24T21:03:06.685034Z","iopub.status.idle":"2024-09-24T21:03:09.316781Z","shell.execute_reply.started":"2024-09-24T21:03:06.684999Z","shell.execute_reply":"2024-09-24T21:03:09.315900Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample_columns = sample_submission.columns\nsample_columns= sample_columns[1:]\nsubmission_df = pd.DataFrame(target_pred.detach().cpu().numpy(), columns=sample_columns)","metadata":{"execution":{"iopub.status.busy":"2024-09-24T21:03:09.318585Z","iopub.execute_input":"2024-09-24T21:03:09.319033Z","iopub.status.idle":"2024-09-24T21:03:09.328349Z","shell.execute_reply.started":"2024-09-24T21:03:09.318997Z","shell.execute_reply":"2024-09-24T21:03:09.327436Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_df.insert(0, 'id', range(255))","metadata":{"execution":{"iopub.status.busy":"2024-09-24T21:03:11.000908Z","iopub.execute_input":"2024-09-24T21:03:11.001620Z","iopub.status.idle":"2024-09-24T21:03:11.010309Z","shell.execute_reply.started":"2024-09-24T21:03:11.001587Z","shell.execute_reply":"2024-09-24T21:03:11.009307Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample_submission","metadata":{"execution":{"iopub.status.busy":"2024-09-24T21:03:12.825306Z","iopub.execute_input":"2024-09-24T21:03:12.826004Z","iopub.status.idle":"2024-09-24T21:03:12.872452Z","shell.execute_reply.started":"2024-09-24T21:03:12.825972Z","shell.execute_reply":"2024-09-24T21:03:12.871575Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_df","metadata":{"execution":{"iopub.status.busy":"2024-09-24T21:03:14.214559Z","iopub.execute_input":"2024-09-24T21:03:14.215234Z","iopub.status.idle":"2024-09-24T21:03:14.246968Z","shell.execute_reply.started":"2024-09-24T21:03:14.215196Z","shell.execute_reply":"2024-09-24T21:03:14.246069Z"},"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-24T21:03:17.129533Z","iopub.execute_input":"2024-09-24T21:03:17.130152Z","iopub.status.idle":"2024-09-24T21:03:25.670402Z","shell.execute_reply.started":"2024-09-24T21:03:17.130118Z","shell.execute_reply":"2024-09-24T21:03:25.669359Z"},"trusted":true},"execution_count":null,"outputs":[]}]}