{"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-28T15:56:35.010746Z","iopub.execute_input":"2024-09-28T15:56:35.011514Z","iopub.status.idle":"2024-09-28T15:56:35.019948Z","shell.execute_reply.started":"2024-09-28T15:56:35.011471Z","shell.execute_reply":"2024-09-28T15:56:35.019074Z"},"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-28T15:56:35.029374Z","iopub.execute_input":"2024-09-28T15:56:35.029674Z","iopub.status.idle":"2024-09-28T15:56:35.034730Z","shell.execute_reply.started":"2024-09-28T15:56:35.029641Z","shell.execute_reply":"2024-09-28T15:56:35.033886Z"},"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-28T15:56:35.051206Z","iopub.execute_input":"2024-09-28T15:56:35.051709Z","iopub.status.idle":"2024-09-28T15:56:57.270753Z","shell.execute_reply.started":"2024-09-28T15:56:35.051677Z","shell.execute_reply":"2024-09-28T15:56:57.269435Z"},"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-28T15:56:57.272941Z","iopub.execute_input":"2024-09-28T15:56:57.273313Z","iopub.status.idle":"2024-09-28T15:56:58.399203Z","shell.execute_reply.started":"2024-09-28T15:56:57.273270Z","shell.execute_reply":"2024-09-28T15:56:58.398156Z"},"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-28T15:56:58.400354Z","iopub.execute_input":"2024-09-28T15:56:58.400680Z","iopub.status.idle":"2024-09-28T15:56:59.845740Z","shell.execute_reply.started":"2024-09-28T15:56:58.400645Z","shell.execute_reply":"2024-09-28T15:56:59.844564Z"},"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-28T15:56:59.848718Z","iopub.execute_input":"2024-09-28T15:56:59.849131Z","iopub.status.idle":"2024-09-28T15:56:59.855596Z","shell.execute_reply.started":"2024-09-28T15:56:59.849086Z","shell.execute_reply":"2024-09-28T15:56:59.854750Z"},"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\nfrom rdkit.Chem import AllChem\nfrom rdkit import DataStructs\nfrom rdkit.Chem import Descriptors, rdMolDescriptors, QED\nfrom rdkit.Chem.rdMolDescriptors import CalcTPSA, CalcNumRotatableBonds, CalcNumHBA, CalcNumHBD, CalcFractionCSP3\nfrom rdkit.Chem import BRICS, Recap\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import DataLoader, TensorDataset\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    \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 = self.scaler.fit_transform(smiles_info_df)\n        else:\n            smiles_info_df = self.scaler.transform(smiles_info_df)\n        return pd.DataFrame(smiles_info_df, columns=[\n            'Molecular Weight', 'LogP', 'TPSA', 'Number of Atoms', 'Number of Bonds',\n            'Number of Rotatable Bonds', 'Number of Hydrogen Bond Acceptors', 'Number of Hydrogen Bond Donors',\n            'Number of Rings', 'Number of Aromatic Rings', 'Number of Stereocenters',\n            'Fraction of sp3 Carbons', 'Balaban J Index', 'Bertz CT', 'QED Score'\n        ])\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            'TPSA': CalcTPSA(mol),\n            'Number of Atoms': mol.GetNumAtoms(),\n            'Number of Bonds': mol.GetNumBonds(),\n            'Number of Rotatable Bonds': CalcNumRotatableBonds(mol),\n            'Number of Hydrogen Bond Acceptors': CalcNumHBA(mol),\n            'Number of Hydrogen Bond Donors': CalcNumHBD(mol),\n            'Number of Rings': Descriptors.RingCount(mol),\n            'Number of Aromatic Rings': rdMolDescriptors.CalcNumAromaticRings(mol),\n            'Number of Stereocenters': len(Chem.FindMolChiralCenters(mol, includeUnassigned=True)),\n            'Fraction of sp3 Carbons': CalcFractionCSP3(mol),\n            'Balaban J Index': Descriptors.BalabanJ(mol),\n            'Bertz CT': Descriptors.BertzCT(mol),\n            'QED Score': QED.qed(mol)\n        }\n        return info\n\n    @property\n    def embedding_size(self):\n        return 15  \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))))\n    \n\nclass Autoencoder(nn.Module):\n    def __init__(self, input_size, hidden_size):\n        super(Autoencoder, self).__init__()\n        # Encoder\n        self.encoder = nn.Sequential(\n            nn.Linear(input_size, hidden_size),\n            nn.ReLU()\n        )\n        # Decoder\n        self.decoder = nn.Sequential(\n            nn.Linear(hidden_size, input_size),\n            nn.Sigmoid()\n        )\n\n    def forward(self, x):\n        encoded = self.encoder(x)\n        decoded = self.decoder(encoded)\n        return encoded, decoded      \n    \nclass MorganFingerPrintEmbedding(BaseEmbedding):\n    def __init__(self, hidden_size=128):\n        self.hidden_size = hidden_size\n        self.autoencoder = None  # This will store the trained autoencoder\n\n    def preprocess(self, df):\n        return df\n\n    def get_embedding(self, df, fit=True):\n        # Extract Morgan fingerprints\n        morgan_fp_list = df['SMILES'].apply(self.extract_morgan_fingerprint)\n        morgan_fp_array = np.stack(morgan_fp_list)\n\n        input_size = morgan_fp_array.shape[1]\n\n        if fit:\n            # Train autoencoder\n            self.autoencoder = self.train_autoencoder(morgan_fp_array, input_size, self.hidden_size)\n        \n        # Get compressed fingerprints\n        with torch.no_grad():\n            compressed_fp = self.autoencoder.encoder(torch.tensor(morgan_fp_array, dtype=torch.float32)).numpy()\n\n        compressed_fp_df = pd.DataFrame(compressed_fp, columns=[f'CompressedFP_{i}' for i in range(self.hidden_size)])\n        return compressed_fp_df\n\n    def extract_morgan_fingerprint(self, smiles, radius=2, nBits=2048):\n        if smiles is None or smiles == '':\n            return None\n        mol = Chem.MolFromSmiles(smiles)\n        if mol is None:\n            return None\n        \n        morgan_gen = AllChem.GetMorganGenerator(radius=radius, fpSize=nBits)\n        fp = morgan_gen.GetFingerprint(mol)\n        \n        fp_array = np.zeros((nBits,))\n        DataStructs.ConvertToNumpyArray(fp, fp_array)\n        \n        return fp_array\n\n    def train_autoencoder(self, data, input_size, hidden_size, num_epochs=70, batch_size=32, learning_rate=0.001):\n        autoencoder = Autoencoder(input_size=input_size, hidden_size=hidden_size)\n        criterion = nn.MSELoss()\n        optimizer = torch.optim.Adam(autoencoder.parameters(), lr=learning_rate)\n\n        dataset = TensorDataset(torch.tensor(data, dtype=torch.float32))\n        dataloader = DataLoader(dataset, batch_size=batch_size, shuffle=True)\n\n        for epoch in range(num_epochs):\n            for batch in dataloader:\n                inputs = batch[0]\n                _, reconstructed = autoencoder(inputs)\n\n                loss = criterion(reconstructed, inputs)\n                optimizer.zero_grad()\n                loss.backward()\n                optimizer.step()\n\n            if (epoch + 1) % 10 == 0:\n                print(f'Epoch [{epoch + 1}/{num_epochs}], Loss: {loss.item():.4f}')\n\n        return autoencoder\n\n    @property\n    def embedding_size(self):\n        return self.hidden_size","metadata":{"execution":{"iopub.status.busy":"2024-09-28T15:56:59.857066Z","iopub.execute_input":"2024-09-28T15:56:59.857623Z","iopub.status.idle":"2024-09-28T15:56:59.899616Z","shell.execute_reply.started":"2024-09-28T15:56:59.857581Z","shell.execute_reply":"2024-09-28T15:56:59.898709Z"},"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)\n\n    def generate_expert_config(self, expert_specs):\n        \"\"\"\n        expert_specs: A dictionary where keys are expert names and values are lists of embedding names to combine.\n        \n        Example format:\n        {\n            \"lstm_expert_1\": [\"OneHotEmbedding\", \"MorganFingerPrintEmbedding\"],\n            \"lstm_expert_2\": [\"OneHotEmbedding\", \"SMILESEmbedding\"],\n            \"lstm_expert_3\": [\"SMILESEmbedding\", \"Mol2VecEmbedding\"],\n        }\n        \"\"\"\n        expert_config = {}\n        \n        for expert, embedding_names in expert_specs.items():\n            indices = []\n            for embedding_name in embedding_names:\n                start, end = self.embedding_indices[embedding_name]\n                indices.extend(range(start, end))\n            expert_config[expert] = indices\n        \n        return expert_config","metadata":{"execution":{"iopub.status.busy":"2024-09-28T15:56:59.900741Z","iopub.execute_input":"2024-09-28T15:56:59.901092Z","iopub.status.idle":"2024-09-28T15:56:59.913119Z","shell.execute_reply.started":"2024-09-28T15:56:59.901058Z","shell.execute_reply":"2024-09-28T15:56:59.912179Z"},"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')\nmorgan_embedding = MorganFingerPrintEmbedding(hidden_size=128)","metadata":{"execution":{"iopub.status.busy":"2024-09-28T15:56:59.914271Z","iopub.execute_input":"2024-09-28T15:56:59.914562Z","iopub.status.idle":"2024-09-28T15:57:01.055506Z","shell.execute_reply.started":"2024-09-28T15:56:59.914530Z","shell.execute_reply":"2024-09-28T15:57:01.054718Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preprocessor_train = Preprocessor([onehot_embedding, smiles_embedding, mol2vec_embedding,morgan_embedding])","metadata":{"execution":{"iopub.status.busy":"2024-09-28T15:57:01.056591Z","iopub.execute_input":"2024-09-28T15:57:01.056888Z","iopub.status.idle":"2024-09-28T15:57:01.064435Z","shell.execute_reply.started":"2024-09-28T15:57:01.056855Z","shell.execute_reply":"2024-09-28T15:57:01.063698Z"},"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-28T15:57:01.065654Z","iopub.execute_input":"2024-09-28T15:57:01.066577Z","iopub.status.idle":"2024-09-28T15:57:12.122210Z","shell.execute_reply.started":"2024-09-28T15:57:01.066524Z","shell.execute_reply":"2024-09-28T15:57:12.121241Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preprocessor_train.embedding_indices","metadata":{"execution":{"iopub.status.busy":"2024-09-28T15:57:12.125727Z","iopub.execute_input":"2024-09-28T15:57:12.126158Z","iopub.status.idle":"2024-09-28T15:57:12.132816Z","shell.execute_reply.started":"2024-09-28T15:57:12.126119Z","shell.execute_reply":"2024-09-28T15:57:12.131673Z"},"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)","metadata":{"execution":{"iopub.status.busy":"2024-09-28T15:57:12.134099Z","iopub.execute_input":"2024-09-28T15:57:12.134477Z","iopub.status.idle":"2024-09-28T15:57:12.144000Z","shell.execute_reply.started":"2024-09-28T15:57:12.134433Z","shell.execute_reply":"2024-09-28T15:57:12.143001Z"},"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-28T15:57:12.145289Z","iopub.execute_input":"2024-09-28T15:57:12.145593Z","iopub.status.idle":"2024-09-28T15:57:12.155862Z","shell.execute_reply.started":"2024-09-28T15:57:12.145549Z","shell.execute_reply":"2024-09-28T15:57:12.154879Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(targets.dtypes)","metadata":{"execution":{"iopub.status.busy":"2024-09-28T15:57:12.157061Z","iopub.execute_input":"2024-09-28T15:57:12.157460Z","iopub.status.idle":"2024-09-28T15:57:12.168044Z","shell.execute_reply.started":"2024-09-28T15:57:12.157407Z","shell.execute_reply":"2024-09-28T15:57:12.167012Z"},"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-28T15:57:12.169317Z","iopub.execute_input":"2024-09-28T15:57:12.169672Z","iopub.status.idle":"2024-09-28T15:57:12.196936Z","shell.execute_reply.started":"2024-09-28T15:57:12.169635Z","shell.execute_reply":"2024-09-28T15:57:12.196012Z"},"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=32, shuffle=True)\nval_loader = DataLoader(val_dataset, batch_size=32)","metadata":{"execution":{"iopub.status.busy":"2024-09-28T15:57:12.198421Z","iopub.execute_input":"2024-09-28T15:57:12.198828Z","iopub.status.idle":"2024-09-28T15:57:12.220485Z","shell.execute_reply.started":"2024-09-28T15:57:12.198781Z","shell.execute_reply":"2024-09-28T15:57:12.219716Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import random\nseed = 42\nrandom.seed(seed)\nos.environ['PYTHONHASHSEED'] = str(seed)\nos.environ['TOKENIZERS_PARALLELISM'] = 'true'\nnp.random.seed(seed)\ntorch.manual_seed(seed)\nif torch.cuda.is_available():\n    torch.cuda.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n    print('-----Seed Set!-----')","metadata":{"execution":{"iopub.status.busy":"2024-09-28T15:57:12.221850Z","iopub.execute_input":"2024-09-28T15:57:12.222304Z","iopub.status.idle":"2024-09-28T15:57:12.230667Z","shell.execute_reply.started":"2024-09-28T15:57:12.222260Z","shell.execute_reply":"2024-09-28T15:57:12.229704Z"},"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\nclass LSTMExpert(nn.Module):\n    def __init__(self, input_size, output_size, hidden_size=256, num_layers=2, dropout=0.3):\n        super(LSTMExpert, self).__init__()\n        self.input_size = input_size\n        self.hidden_size = hidden_size\n        self.num_layers = num_layers\n        \n        # Bidirectional LSTM layers\n        self.lstm = nn.LSTM(input_size, hidden_size, num_layers=num_layers, \n                            batch_first=True, bidirectional=True, dropout=dropout)\n        \n        # Attention mechanism\n        self.attention = nn.MultiheadAttention(hidden_size * 2, num_heads=4)\n        \n        # Fully connected layers with residual connections\n        self.fc1 = nn.Linear(hidden_size * 2, hidden_size * 2)\n        self.fc2 = nn.Linear(hidden_size * 2, hidden_size * 2)\n        self.fc3 = nn.Linear(hidden_size * 2, output_size)\n        \n        # Layer normalization\n        self.layer_norm1 = nn.LayerNorm(hidden_size * 2)\n        self.layer_norm2 = nn.LayerNorm(hidden_size * 2)\n        self.layer_norm3 = nn.LayerNorm(hidden_size * 2)\n        \n        self.dropout = nn.Dropout(dropout)\n        \n    def forward(self, x):\n        # Ensure input is 3D\n        if x.dim() == 2:\n            x = x.unsqueeze(1)\n        \n        # LSTM layers\n        lstm_out, _ = self.lstm(x)\n        \n        # Apply attention\n        lstm_out = lstm_out.permute(1, 0, 2)  # Change to (seq_len, batch, feature)\n        attn_out, _ = self.attention(lstm_out, lstm_out, lstm_out)\n        attn_out = attn_out.permute(1, 0, 2)  # Change back to (batch, seq_len, feature)\n        \n        # Take the last output or apply global average pooling\n        if attn_out.size(1) == 1:\n            out = attn_out.squeeze(1)\n        else:\n            out = F.adaptive_avg_pool1d(attn_out.transpose(1, 2), 1).squeeze(-1)\n        \n        # Fully connected layers with residual connections\n        out = self.layer_norm1(out)\n        residual = out\n        \n        out = self.fc1(out)\n        out = F.relu(out)\n        out = self.dropout(out)\n        out = self.layer_norm2(out)\n        out = out + residual\n        \n        residual = out\n        out = self.fc2(out)\n        out = F.relu(out)\n        out = self.dropout(out)\n        out = self.layer_norm3(out)\n        out = out + residual\n        \n        out = self.fc3(out)\n        \n        return out\n\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, hidden_size=512):\n        super(GatingNetwork, self).__init__()\n        self.input_size = input_size\n        self.num_experts = num_experts\n        \n        self.feature_extractor = nn.Sequential(\n            nn.Linear(input_size, hidden_size),\n            nn.LayerNorm(hidden_size),\n            nn.ReLU(),\n            nn.Dropout(0.3),\n            nn.Linear(hidden_size, hidden_size),\n            nn.LayerNorm(hidden_size),\n            nn.ReLU(),\n            nn.Dropout(0.3)\n        )\n        \n        self.attention = nn.MultiheadAttention(hidden_size, num_heads=4)\n        \n        self.output_layer = nn.Linear(hidden_size, num_experts)\n        \n    def forward(self, x):\n        features = self.feature_extractor(x)\n        features = features.unsqueeze(0)  # Add sequence dimension for attention\n        attn_output, _ = self.attention(features, features, features)\n        attn_output = attn_output.squeeze(0)\n        \n        logits = self.output_layer(attn_output)\n        return F.softmax(logits, dim=-1)\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, return_gating_logits=False):\n        # Prepare input for gating network\n        gating_input = data[:, self.total_features]\n        gating_logits = self.gating_network(gating_input)\n        gating_weights = F.softmax(gating_logits, dim=-1)\n        self.last_gating_weights = gating_weights.detach()\n\n        expert_outputs = []\n        for expert_name, expert in self.experts.items():\n            expert_input = data[:, self.feature_indices[expert_name]]\n            expert_output = expert(expert_input)\n            expert_outputs.append(expert_output)\n\n        expert_outputs = torch.stack(expert_outputs, dim=1)\n        combined_output = torch.sum(expert_outputs * gating_weights.unsqueeze(-1), dim=1)\n\n        if return_gating_logits:\n            return combined_output, gating_logits\n        else:\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-28T16:10:25.955389Z","iopub.execute_input":"2024-09-28T16:10:25.955780Z","iopub.status.idle":"2024-09-28T16:10:25.993129Z","shell.execute_reply.started":"2024-09-28T16:10:25.955725Z","shell.execute_reply":"2024-09-28T16:10:25.992099Z"},"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-28T16:10:25.995146Z","iopub.execute_input":"2024-09-28T16:10:25.995583Z","iopub.status.idle":"2024-09-28T16:10:26.011351Z","shell.execute_reply.started":"2024-09-28T16:10:25.995536Z","shell.execute_reply":"2024-09-28T16:10:26.010514Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.optim.lr_scheduler import CosineAnnealingLR\nfrom torch.nn.utils import clip_grad_norm_\n\ndef noisy_top_k_gating(logits, noise_epsilon=0.1, training=True):\n    if training:\n        noise = torch.randn_like(logits) * noise_epsilon\n        noisy_logits = logits + noise\n    else:\n        noisy_logits = logits\n    return F.softmax(noisy_logits, dim=-1)\n\ndef custom_loss(y_pred, y_true, gating_logits, diversity_weight=0.1, noise_epsilon=0.1, training=True):\n    # Main loss (RMSE)\n    main_loss = RMSE_rowwise_loss(y_pred, y_true)\n    \n    # Apply noisy top-k gating\n    noisy_gating_weights = noisy_top_k_gating(gating_logits, noise_epsilon, training)\n    \n    # Diversity loss (encourage using different experts)\n    avg_expert_usage = noisy_gating_weights.mean(dim=0)\n    diversity_loss = -torch.sum(avg_expert_usage * torch.log(avg_expert_usage + 1e-10))\n    \n    # Combine losses\n    total_loss = main_loss - diversity_weight * diversity_loss\n    \n    return total_loss, noisy_gating_weights\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, noise_epsilon=1e-2):\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=lr, weight_decay=1e-4)\n    scheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='min', factor=0.5, patience=5, verbose=True)\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            \n            # Forward pass\n            output, gating_logits = moe(X, return_gating_logits=True)\n            \n            # Custom loss with noisy top-k gating\n            loss, noisy_gating_weights = custom_loss(output, y, gating_logits, \n                                                     noise_epsilon=noise_epsilon, \n                                                     training=True)\n            \n            loss.backward()\n            clip_grad_norm_(moe.parameters(), max_norm=1.0)\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, gating_logits = moe(X, return_gating_logits=True)\n                loss, _ = custom_loss(output, y, gating_logits, \n                                      noise_epsilon=noise_epsilon, \n                                      training=False)\n                val_loss += loss.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        if epoch % 10 == 0:\n            plot_expert_utilization(moe.last_gating_weights)\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()  \n    \ndef plot_expert_utilization(gating_weights):\n    avg_expert_usage = gating_weights.mean(dim=0).cpu().numpy()\n    plt.figure(figsize=(10, 5))\n    plt.bar(range(len(avg_expert_usage)), avg_expert_usage)\n    plt.title(\"Average Expert Utilization\")\n    plt.xlabel(\"Expert\")\n    plt.ylabel(\"Average Usage\")\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2024-09-28T16:10:26.180558Z","iopub.execute_input":"2024-09-28T16:10:26.181125Z","iopub.status.idle":"2024-09-28T16:10:26.213663Z","shell.execute_reply.started":"2024-09-28T16:10:26.181082Z","shell.execute_reply":"2024-09-28T16:10:26.212544Z"},"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-28T16:10:26.215562Z","iopub.execute_input":"2024-09-28T16:10:26.216253Z","iopub.status.idle":"2024-09-28T16:10:26.229781Z","shell.execute_reply.started":"2024-09-28T16:10:26.216199Z","shell.execute_reply":"2024-09-28T16:10:26.228885Z"},"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-28T16:10:26.230887Z","iopub.execute_input":"2024-09-28T16:10:26.231854Z","iopub.status.idle":"2024-09-28T16:10:26.244254Z","shell.execute_reply.started":"2024-09-28T16:10:26.231817Z","shell.execute_reply":"2024-09-28T16:10:26.243418Z"},"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-28T16:10:26.245519Z","iopub.execute_input":"2024-09-28T16:10:26.245866Z","iopub.status.idle":"2024-09-28T16:10:26.263922Z","shell.execute_reply.started":"2024-09-28T16:10:26.245831Z","shell.execute_reply":"2024-09-28T16:10:26.262997Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch.optim as optim\nfrom torch.autograd import Variable\n\n# expert_config = {\n#     \"lstm_expert_1\": list(range(152)) + list(range(467,595)) ,\n#     \"lstm_expert_2\": list(range(152)) + list(range(152,452)),\n#     \"lstm_expert_3\": list(range(152)) + list(range(452,467)),\n# }\nexpert_specs = {\n    \"lstm_expert_1\": [\"OneHotEmbedding\", \"MorganFingerPrintEmbedding\"],\n    \"lstm_expert_2\": [\"OneHotEmbedding\", \"Mol2VecEmbedding\"],\n    \"lstm_expert_3\": [\"OneHotEmbedding\", \"SMILESEmbedding\"],\n}\n\nexpert_config = preprocessor_train.generate_expert_config(expert_specs)\n\n# print(expert_config)\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\n# Pretrain experts (if needed)\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 with improvements\nprint(\"Training gating network\")\ntrain_losses, val_losses = train_gating_network(moe, train_loader, val_loader, device, epochs=300, lr=1e-4)\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-28T16:10:26.267399Z","iopub.execute_input":"2024-09-28T16:10:26.267699Z","iopub.status.idle":"2024-09-28T16:10:49.060619Z","shell.execute_reply.started":"2024-09-28T16:10:26.267667Z","shell.execute_reply":"2024-09-28T16:10:49.059670Z"},"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","metadata":{"execution":{"iopub.status.busy":"2024-09-28T16:10:49.062046Z","iopub.execute_input":"2024-09-28T16:10:49.062451Z","iopub.status.idle":"2024-09-28T16:10:51.487674Z","shell.execute_reply.started":"2024-09-28T16:10:49.062401Z","shell.execute_reply":"2024-09-28T16:10:51.486653Z"},"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-28T16:10:51.488869Z","iopub.execute_input":"2024-09-28T16:10:51.489181Z","iopub.status.idle":"2024-09-28T16:10:51.496812Z","shell.execute_reply.started":"2024-09-28T16:10:51.489146Z","shell.execute_reply":"2024-09-28T16:10:51.495952Z"},"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-28T16:10:51.498066Z","iopub.execute_input":"2024-09-28T16:10:51.498468Z","iopub.status.idle":"2024-09-28T16:10:51.547712Z","shell.execute_reply.started":"2024-09-28T16:10:51.498414Z","shell.execute_reply":"2024-09-28T16:10:51.546754Z"},"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-28T16:10:51.549035Z","iopub.execute_input":"2024-09-28T16:10:51.549360Z","iopub.status.idle":"2024-09-28T16:10:51.559635Z","shell.execute_reply.started":"2024-09-28T16:10:51.549326Z","shell.execute_reply":"2024-09-28T16:10:51.558547Z"},"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-28T16:10:51.560895Z","iopub.execute_input":"2024-09-28T16:10:51.561336Z","iopub.status.idle":"2024-09-28T16:10:53.782352Z","shell.execute_reply.started":"2024-09-28T16:10:51.561282Z","shell.execute_reply":"2024-09-28T16:10:53.781458Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_array =  processed_df_test.to_numpy()","metadata":{"execution":{"iopub.status.busy":"2024-09-28T16:10:53.783652Z","iopub.execute_input":"2024-09-28T16:10:53.784064Z","iopub.status.idle":"2024-09-28T16:10:53.789073Z","shell.execute_reply.started":"2024-09-28T16:10:53.784018Z","shell.execute_reply":"2024-09-28T16:10:53.788221Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.to(device)\ntarget_pred = model(torch.Tensor(test_array).to(device))","metadata":{"execution":{"iopub.status.busy":"2024-09-28T16:10:53.790505Z","iopub.execute_input":"2024-09-28T16:10:53.790896Z","iopub.status.idle":"2024-09-28T16:10:53.822581Z","shell.execute_reply.started":"2024-09-28T16:10:53.790836Z","shell.execute_reply":"2024-09-28T16:10:53.821739Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"res = target_pred.cpu()","metadata":{"execution":{"iopub.status.busy":"2024-09-28T16:10:53.826328Z","iopub.execute_input":"2024-09-28T16:10:53.826719Z","iopub.status.idle":"2024-09-28T16:10:53.843044Z","shell.execute_reply.started":"2024-09-28T16:10:53.826684Z","shell.execute_reply":"2024-09-28T16:10:53.842214Z"},"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-28T16:10:53.844479Z","iopub.execute_input":"2024-09-28T16:10:53.845097Z","iopub.status.idle":"2024-09-28T16:10:56.383983Z","shell.execute_reply.started":"2024-09-28T16:10:53.845046Z","shell.execute_reply":"2024-09-28T16:10:56.383008Z"},"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-28T16:10:56.385249Z","iopub.execute_input":"2024-09-28T16:10:56.385666Z","iopub.status.idle":"2024-09-28T16:10:56.392578Z","shell.execute_reply.started":"2024-09-28T16:10:56.385615Z","shell.execute_reply":"2024-09-28T16:10:56.391469Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_df","metadata":{"execution":{"iopub.status.busy":"2024-09-28T16:10:56.393879Z","iopub.execute_input":"2024-09-28T16:10:56.394247Z","iopub.status.idle":"2024-09-28T16:10:56.487366Z","shell.execute_reply.started":"2024-09-28T16:10:56.394214Z","shell.execute_reply":"2024-09-28T16:10:56.486341Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_df.insert(0, 'id', range(255))","metadata":{"execution":{"iopub.status.busy":"2024-09-28T16:10:56.488619Z","iopub.execute_input":"2024-09-28T16:10:56.489016Z","iopub.status.idle":"2024-09-28T16:10:56.494744Z","shell.execute_reply.started":"2024-09-28T16:10:56.488938Z","shell.execute_reply":"2024-09-28T16:10:56.493779Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_df.to_csv(\"submission1.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2024-09-28T16:10:56.496085Z","iopub.execute_input":"2024-09-28T16:10:56.496732Z","iopub.status.idle":"2024-09-28T16:11:04.994499Z","shell.execute_reply.started":"2024-09-28T16:10:56.496693Z","shell.execute_reply":"2024-09-28T16:11:04.993685Z"},"trusted":true},"execution_count":null,"outputs":[]}]}