{"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":"markdown","source":"# **OPxMOE V2**","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\nimport torch \nimport numpy as np\nimport pandas as pd\nimport os\nimport random\nimport kagglehub\nfrom kaggle_secrets import UserSecretsClient\n\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n        \ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f'\\nUsing {device}')\n\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\nprint('-----Seed Set!-----')\n\nuser_secrets = UserSecretsClient()\nos.environ[\"KAGGLE_USERNAME\"] = user_secrets.get_secret(\"Kaggle_Username\")\nos.environ[\"KAGGLE_KEY\"] = user_secrets.get_secret(\"Kaggle_Key\")\nkagglehub.login()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2025-01-19T18:40:40.474940Z","iopub.execute_input":"2025-01-19T18:40:40.475817Z","iopub.status.idle":"2025-01-19T18:40:42.693253Z","shell.execute_reply.started":"2025-01-19T18:40:40.475763Z","shell.execute_reply":"2025-01-19T18:40:42.692221Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **Installing Requirements**","metadata":{}},{"cell_type":"code","source":"#Installing a package\n!pip install -q umap-learn rdkit captum git+https://github.com/samoturk/mol2vec;","metadata":{"execution":{"iopub.status.busy":"2025-01-19T18:40:42.695174Z","iopub.execute_input":"2025-01-19T18:40:42.695655Z","iopub.status.idle":"2025-01-19T18:41:09.743188Z","shell.execute_reply.started":"2025-01-19T18:40:42.695615Z","shell.execute_reply":"2025-01-19T18:41:09.742237Z"},"trusted":true},"outputs":[],"execution_count":null},{"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":"2025-01-19T18:41:09.744436Z","iopub.execute_input":"2025-01-19T18:41:09.744737Z","iopub.status.idle":"2025-01-19T18:41:18.777408Z","shell.execute_reply.started":"2025-01-19T18:41:09.744709Z","shell.execute_reply":"2025-01-19T18:41:18.776718Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = pd.read_parquet('../input/open-problems-single-cell-perturbations/de_train.parquet')\ndf.tail()","metadata":{"execution":{"iopub.status.busy":"2025-01-19T18:41:18.779123Z","iopub.execute_input":"2025-01-19T18:41:18.779393Z","iopub.status.idle":"2025-01-19T18:41:20.577674Z","shell.execute_reply.started":"2025-01-19T18:41:18.779366Z","shell.execute_reply":"2025-01-19T18:41:20.576815Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **Preprocessing Classes**","metadata":{}},{"cell_type":"code","source":"from abc import ABC, abstractmethod\n\nclass BaseEmbedding(ABC):\n    \"\"\"Abstract base class for embeddings with save/load functionality.\"\"\"\n    \n    def __init__(self):\n        self.embedding_dict = {}\n        self.metadata = {}\n    \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\n    \n    def compute_and_store_embeddings(self, df, entity_column):\n        \"\"\"\n        Compute and store embeddings for unique entities in the specified column.\n        \n        Args:\n            df (pd.DataFrame): Input DataFrame\n            entity_column (str): Column name containing entities (e.g., 'cell_type' or 'sm_name')\n        \"\"\"\n        unique_entities = df[entity_column].unique()\n        for entity in unique_entities:\n            entity_df = df[df[entity_column] == entity].copy()\n            embedding = self.get_embedding(entity_df, fit=False)\n            # Store the mean embedding if there are multiple rows\n            self.embedding_dict[entity] = embedding.mean().values\n            \n    def get_entity_embedding(self, entity_name):\n        \"\"\"\n        Retrieve embedding for a specific entity.\n        \n        Args:\n            entity_name (str): Name of the entity\n            \n        Returns:\n            np.ndarray: Embedding vector for the entity\n        \"\"\"\n        if entity_name in self.embedding_dict:\n            return self.embedding_dict[entity_name]\n        else:\n            return np.zeros(self.embedding_size)\n    \n    def save_embeddings(self, filepath, metadata=None):\n        \"\"\"\n        Save embeddings and metadata to disk.\n        \n        Args:\n            filepath (str): Path to save the embeddings\n            metadata (dict, optional): Additional metadata to save\n        \"\"\"\n        directory = os.path.dirname(filepath)\n        if directory and not os.path.exists(directory):\n            os.makedirs(directory)\n            \n        # Convert numpy arrays to lists for JSON serialization\n        serializable_dict = {\n            k: v.tolist() if isinstance(v, (np.ndarray, torch.Tensor)) else v \n            for k, v in self.embedding_dict.items()\n        }\n        \n        # Prepare save data\n        save_data = {\n            'embeddings': serializable_dict,\n            'metadata': metadata or self.metadata,\n            'embedding_size': self.embedding_size,\n            'embedding_type': self.__class__.__name__\n        }\n        \n        # Save as pickle if the filepath ends with .pkl, otherwise save as JSON\n        if filepath.endswith('.pkl'):\n            with open(filepath, 'wb') as f:\n                pickle.dump(save_data, f)\n        else:\n            with open(filepath, 'w') as f:\n                json.dump(save_data, f)\n                \n    def load_embeddings(self, filepath):\n        \"\"\"\n        Load embeddings and metadata from disk.\n        \n        Args:\n            filepath (str): Path to load the embeddings from\n            \n        Returns:\n            bool: True if loading was successful\n        \"\"\"\n        try:\n            # Load pickle if the filepath ends with .pkl, otherwise load JSON\n            if filepath.endswith('.pkl'):\n                with open(filepath, 'rb') as f:\n                    save_data = pickle.load(f)\n            else:\n                with open(filepath, 'r') as f:\n                    save_data = json.load(f)\n            \n            # Convert lists back to numpy arrays\n            self.embedding_dict = {\n                k: np.array(v) if isinstance(v, list) else v \n                for k, v in save_data['embeddings'].items()\n            }\n            \n            self.metadata = save_data.get('metadata', {})\n            \n            # Verify embedding size matches\n            if save_data['embedding_size'] != self.embedding_size:\n                raise ValueError(\n                    f\"Loaded embedding size ({save_data['embedding_size']}) \"\n                    f\"doesn't match current embedding size ({self.embedding_size})\"\n                )\n            \n            return True\n            \n        except Exception as e:\n            print(f\"Error loading embeddings: {str(e)}\")\n            return False","metadata":{"execution":{"iopub.status.busy":"2025-01-19T18:41:20.579252Z","iopub.execute_input":"2025-01-19T18:41:20.579784Z","iopub.status.idle":"2025-01-19T18:41:20.592865Z","shell.execute_reply.started":"2025-01-19T18:41:20.579743Z","shell.execute_reply":"2025-01-19T18:41:20.591951Z"},"trusted":true},"outputs":[],"execution_count":null},{"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\nfrom sklearn.base import BaseEstimator, TransformerMixin\n\nclass OneHotEmbedding(BaseEmbedding):\n    def __init__(self, columns):\n        super().__init__()\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        \n        self._embedding_size = encoded_features.shape[1]\n        encoded_df = pd.DataFrame(encoded_features, columns=self.encoder.get_feature_names_out(self.columns))\n        \n        # Store embeddings using base class functionality\n        for idx, row in df[self.columns].iterrows():\n            key = tuple(row.values)\n            self.embedding_dict[key] = encoded_df.loc[idx].values\n            \n        return encoded_df\n\n    @property\n    def embedding_size(self):\n        return self._embedding_size\n\n    \nclass SMILESEmbedding(BaseEmbedding):\n    def __init__(self):\n        super().__init__()\n        self.scaler = StandardScaler()\n\n    def preprocess(self, df):\n        return df\n\n    def create_molecule_embedding_dict(self, df):\n        \"\"\"\n        Create a dictionary of molecule embeddings.\n        \n        Args:\n            df (pd.DataFrame): DataFrame containing SMILES column\n        \"\"\"\n        self.compute_and_store_embeddings(df, 'SMILES')\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\n            \n        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        result_df = pd.DataFrame(smiles_info_df, columns=columns)\n        \n        # Store embeddings using base class functionality\n        for idx, smiles in enumerate(df['SMILES']):\n            self.embedding_dict[smiles] = result_df.loc[idx].values\n            \n        return result_df\n\n    @staticmethod\n    def extract_smiles_info(smiles):\n        if smiles is None or smiles == '':\n            return None\n        mol = Chem.MolFromSmiles(smiles)\n        if mol is None:\n            return None\n        info = {\n            'Molecular Weight': Descriptors.MolWt(mol),\n            'LogP': Descriptors.MolLogP(mol),\n            '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        super().__init__()\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        result_df = pd.DataFrame(df['vector'].tolist(), columns=vector_columns)\n        \n        # Store embeddings using base class functionality\n        for idx, smiles in enumerate(df['SMILES']):\n            self.embedding_dict[smiles] = result_df.loc[idx].values\n            \n        return result_df\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    def create_molecule_embedding_dict(self, df):\n        \"\"\"\n        Create a dictionary of molecule embeddings.\n        \n        Args:\n            df (pd.DataFrame): DataFrame containing SMILES column\n        \"\"\"\n        self.compute_and_store_embeddings(df, 'SMILES')\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.Sigmoid(),\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 TargetEmbedding(BaseEmbedding, nn.Module):\n    \"\"\"Target embedding class using autoencoder to learn compressed representations of medians.\"\"\"\n    \n    def __init__(self, df_train, emb_size=256, hidden_size=1024):\n        \"\"\"\n        Initialize the autoencoder target embedding class.\n        \n        Args:\n            df_train (pd.DataFrame): Training DataFrame for storing medians\n            emb_size (int): Size of the latent embedding for each component\n            hidden_size (int): Size of the hidden layer in encoder/decoder\n        \"\"\"\n        \n        BaseEmbedding.__init__(self)\n        nn.Module.__init__(self)\n        \n        self.emb_size = emb_size\n        self._embedding_size = emb_size * 2  # combined size of both embeddings\n        self.input_size = 18211  # Original dimension of median values\n        \n        # Encoder networks for cell type and small molecule embeddings\n        self.cell_type_encoder = nn.Sequential(\n            nn.Linear(self.input_size, hidden_size),\n            nn.GELU(),\n            nn.Linear(hidden_size, hidden_size),\n            nn.GELU(),\n            nn.Linear(hidden_size, emb_size)\n        )\n        \n        self.sm_encoder = nn.Sequential(\n            nn.Linear(self.input_size, hidden_size),\n            nn.GELU(),\n            nn.Linear(hidden_size, hidden_size),\n            nn.GELU(),\n            nn.Linear(hidden_size, emb_size)\n        )\n        \n        # Decoder networks for reconstruction\n        self.cell_type_decoder = nn.Sequential(\n            nn.Linear(emb_size, hidden_size),\n            nn.GELU(),\n            nn.Linear(hidden_size, hidden_size),\n            nn.GELU(),\n            nn.Linear(hidden_size, self.input_size)\n        )\n        \n        self.sm_decoder = nn.Sequential(\n            nn.Linear(emb_size, hidden_size),\n            nn.GELU(),\n            nn.Linear(hidden_size, hidden_size),\n            nn.GELU(),\n            nn.Linear(hidden_size, self.input_size)\n        )\n        \n        # Initialize dictionaries\n        self.cell_type_dict = {}\n        self.sm_dict = {}\n        \n        # Process training data\n        df_train = self.preprocess(df_train)\n        \n        # Store training data and compute medians\n        self.de_cell_type_train = df_train.iloc[:, [0] + list(range(5, df_train.shape[1]))]\n        self.de_sm_name_train = df_train.iloc[:, [1] + list(range(5, df_train.shape[1]))]\n        \n        # Rename columns for consistency\n        self.de_cell_type_train.columns = ['cell_type' if i == 0 else col \n                                         for i, col in enumerate(self.de_cell_type_train.columns)]\n        self.de_sm_name_train.columns = ['sm_name' if i == 0 else col \n                                         for i, col in enumerate(self.de_sm_name_train.columns)]\n        \n        # Compute medians from training data\n        self.cell_type_medians = self.de_cell_type_train.groupby('cell_type').median()\n        self.sm_name_medians = self.de_sm_name_train.groupby('sm_name').median()\n        \n        # Convert medians to tensors\n        self.cell_type_tensors = {k: torch.tensor(v).float() \n                                 for k, v in zip(self.cell_type_medians.index, \n                                               self.cell_type_medians.values)}\n        self.sm_tensors = {k: torch.tensor(v).float() \n                          for k, v in zip(self.sm_name_medians.index, \n                                        self.sm_name_medians.values)}\n        \n        self.__class__.name = \"TargetEmbedding\"\n        self.is_fitted = False\n    \n    def preprocess(self, df):\n        \"\"\"\n        Preprocess the input DataFrame.\n        \n        Args:\n            df (pd.DataFrame): Input DataFrame\n            \n        Returns:\n            pd.DataFrame: Preprocessed DataFrame\n        \"\"\"\n        # Copy the DataFrame to avoid modifying the original\n        df = df.copy()\n        \n        # Ensure required columns exist\n        required_cols = ['cell_type', 'sm_name']\n        missing_cols = [col for col in required_cols if col not in df.columns]\n        if missing_cols:\n            raise ValueError(f\"Missing required columns: {missing_cols}\")\n            \n        # Handle any missing values in gene expression columns\n        numeric_cols = df.select_dtypes(include=['float64', 'int64']).columns\n        df[numeric_cols] = df[numeric_cols].fillna(0)\n        \n        return df\n    \n    def train_autoencoder(self, num_epochs=100, batch_size=32, learning_rate=1e-3):\n        \"\"\"\n        Train the autoencoder to learn compressed representations of the medians.\n        \n        Args:\n            num_epochs (int): Number of training epochs\n            batch_size (int): Batch size for training\n            learning_rate (float): Learning rate for optimizer\n        \"\"\"\n        self.train()  # Set to training mode\n        optimizer = torch.optim.Adam(self.parameters(), lr=learning_rate)\n        criterion = nn.MSELoss()\n        \n        # Convert median dictionaries to tensors for training\n        cell_type_data = torch.stack(list(self.cell_type_tensors.values()))\n        sm_data = torch.stack(list(self.sm_tensors.values()))\n        \n        for epoch in range(num_epochs):\n            # Train cell type autoencoder\n            cell_type_encoded = self.cell_type_encoder(cell_type_data)\n            cell_type_decoded = self.cell_type_decoder(cell_type_encoded)\n            cell_type_loss = criterion(cell_type_decoded, cell_type_data)\n            \n            # Train small molecule autoencoder\n            sm_encoded = self.sm_encoder(sm_data)\n            sm_decoded = self.sm_decoder(sm_encoded)\n            sm_loss = criterion(sm_decoded, sm_data)\n            \n            # Combined loss\n            total_loss = cell_type_loss + sm_loss\n            \n            # Backpropagation\n            optimizer.zero_grad()\n            total_loss.backward()\n            optimizer.step()\n            \n            if (epoch + 1) % 10 == 0:\n                print(f'Epoch [{epoch+1}/{num_epochs}], '\n                      f'Loss: {total_loss.item():.4f}, '\n                      f'CT Loss: {cell_type_loss.item():.4f}, '\n                      f'SM Loss: {sm_loss.item():.4f}')\n        \n        self.eval()  # Set to evaluation mode\n        # After training, update the dictionaries with encoded values\n        with torch.no_grad():\n            for k, v in self.cell_type_tensors.items():\n                self.cell_type_dict[k] = self.cell_type_encoder(v.unsqueeze(0)).squeeze(0)\n            \n            for k, v in self.sm_tensors.items():\n                self.sm_dict[k] = self.sm_encoder(v.unsqueeze(0)).squeeze(0)\n        \n        self.is_fitted = True\n    \n    def get_embedding(self, df, fit=False):\n        \"\"\"\n        Get embeddings for the input data using the trained autoencoder.\n        \n        Args:\n            df (pd.DataFrame): Input DataFrame with cell_type and sm_name columns\n            fit (bool): If True, train the autoencoder before getting embeddings\n            \n        Returns:\n            pd.DataFrame: DataFrame containing the embedded features\n        \"\"\"\n        # Preprocess input data\n        df = self.preprocess(df)\n        \n        # Train autoencoder if fit=True and not already fitted\n        if fit and not self.is_fitted:\n            self.train_autoencoder()\n        \n        self.eval()  # Set to evaluation mode\n        with torch.no_grad():\n            cell_type_tensors = []\n            sm_tensors = []\n            \n            for _, row in df.iterrows():\n                # Get cell type embedding\n                ct = row['cell_type']\n                if ct in self.cell_type_dict:\n                    cell_type_tensors.append(self.cell_type_dict[ct])\n                else:\n                    cell_type_tensors.append(torch.zeros(self.emb_size))\n                \n                # Get small molecule embedding\n                sm = row['sm_name']\n                if sm in self.sm_dict:\n                    sm_tensors.append(self.sm_dict[sm])\n                else:\n                    sm_tensors.append(torch.zeros(self.emb_size))\n                \n                # Update the parent class embedding dictionary for both entities\n                # Store cell type embedding\n                if ct in self.cell_type_dict:\n                    self.embedding_dict[f\"cell_type_{ct}\"] = self.cell_type_dict[ct].numpy()\n                \n                # Store small molecule embedding\n                if sm in self.sm_dict:\n                    self.embedding_dict[f\"sm_name_{sm}\"] = self.sm_dict[sm].numpy()\n            \n            # Convert to tensors\n            cell_type_tensor = torch.stack(cell_type_tensors)\n            sm_tensor = torch.stack(sm_tensors)\n            \n            # Combine embeddings\n            combined_embedding = torch.cat([cell_type_tensor, sm_tensor], dim=1)\n        \n        # Create column names for the embedding DataFrame\n        ct_cols = [f'target_ct_emb_{i}' for i in range(self.emb_size)]\n        sm_cols = [f'target_sm_emb_{i}' for i in range(self.emb_size)]\n        all_cols = ct_cols + sm_cols\n        \n        return pd.DataFrame(combined_embedding.detach().numpy(), columns=all_cols, index=df.index)\n        \n    @property\n    def embedding_size(self):\n        \"\"\"Returns the size of the combined embedding.\"\"\"\n        return self._embedding_size\n    \nclass MorganFingerPrintEmbedding(BaseEmbedding):\n    def __init__(self, hidden_size=128):\n        super().__init__()\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        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            self.autoencoder = self.train_autoencoder(morgan_fp_array, input_size, self.hidden_size)\n        \n        with torch.no_grad():\n            compressed_fp = self.autoencoder.encoder(torch.tensor(morgan_fp_array, dtype=torch.float32)).numpy()\n\n        result_df = pd.DataFrame(compressed_fp, columns=[f'CompressedFP_{i}' for i in range(self.hidden_size)])\n        \n        # Store embeddings using base class functionality\n        for idx, smiles in enumerate(df['SMILES']):\n            self.embedding_dict[smiles] = result_df.loc[idx].values\n            \n        return result_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=50, 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":"2025-01-19T18:41:20.594408Z","iopub.execute_input":"2025-01-19T18:41:20.594668Z","iopub.status.idle":"2025-01-19T18:41:21.130442Z","shell.execute_reply.started":"2025-01-19T18:41:20.594629Z","shell.execute_reply":"2025-01-19T18:41:21.129768Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Preprocessor:\n    def __init__(self, embeddings):\n        self.embeddings = embeddings\n        self.embedding_indices = {}\n        self.latest_processed_data = None\n        self.sm_name_to_smiles = {}  \n        self.unique_cell_types = set()\n        self.unique_sm_names = set()\n        \n    def preprocess(self, df, fit=True):\n        \"\"\"\n        Process the input DataFrame through all embeddings and combine the results.\n        \n        Args:\n            df (pd.DataFrame): Input DataFrame\n            fit (bool): Whether to fit the embeddings\n            \n        Returns:\n            pd.DataFrame: Combined embedded features\n        \"\"\"\n\n        if fit:\n            if 'cell_type' in df.columns:\n                self.unique_cell_types.update(df['cell_type'].unique())\n            if 'sm_name' in df.columns:\n                self.unique_sm_names.update(df['sm_name'].unique())\n        \n        if fit and 'sm_name' in df.columns and 'SMILES' in df.columns:\n            sm_name_smiles_map = df[['sm_name', 'SMILES']].drop_duplicates()\n            self.sm_name_to_smiles.update(\n                dict(zip(sm_name_smiles_map['sm_name'], sm_name_smiles_map['SMILES']))\n            )\n        \n        processed_dfs = []\n        current_index = 0\n        \n        for embedding in self.embeddings:\n            # For SMILES-based embeddings, ensure SMILES column exists\n            if embedding.__class__.__name__ in ['SMILESEmbedding', 'MorganFingerPrintEmbedding', 'Mol2VecEmbedding']:\n                if 'SMILES' not in df.columns:\n                    # Create a copy of the DataFrame to avoid modifying the original\n                    df_with_smiles = df.copy()\n                    # Add SMILES column using the mapping\n                    df_with_smiles['SMILES'] = df_with_smiles['sm_name'].map(self.sm_name_to_smiles)\n                    preprocessed_df = embedding.preprocess(df_with_smiles)\n                    embedded_df = embedding.get_embedding(preprocessed_df, fit=fit)\n                else:\n                    preprocessed_df = embedding.preprocess(df)\n                    embedded_df = embedding.get_embedding(preprocessed_df, fit=fit)\n            else:\n                preprocessed_df = embedding.preprocess(df)\n                embedded_df = embedding.get_embedding(preprocessed_df, fit=fit)\n            \n            processed_dfs.append(embedded_df)\n            \n            embedding_name = embedding.__class__.__name__\n            self.embedding_indices[embedding_name] = {\n                'start': current_index,\n                'end': current_index + embedding.embedding_size,\n                'columns': list(embedded_df.columns)  # Store column names\n            }\n            current_index += embedding.embedding_size\n        \n        self.latest_processed_data = pd.concat(processed_dfs, axis=1)\n        return self.latest_processed_data\n    \n    def get_entity_embeddings(self, cell_type=None, sm_name=None):\n        \"\"\"\n        Get combined embeddings for a specific cell type and/or small molecule,\n        maintaining the same order as in preprocess.\n        \n        Args:\n            cell_type (str, optional): Cell type name\n            sm_name (str, optional): Small molecule name\n            \n        Returns:\n            pd.DataFrame: DataFrame with embeddings in the same order as preprocess\n            dict: Dictionary containing individual embeddings from each embedding type\n        \"\"\"\n        if cell_type is None and sm_name is None:\n            raise ValueError(\"At least one of cell_type or sm_name must be provided\")\n        \n        individual_embeddings = {}\n        all_values = []\n        all_columns = []\n        \n        # Process embeddings in the same order as in preprocess\n        for embedding in self.embeddings:\n            embedding_name = embedding.__class__.__name__\n            embedding_info = self.embedding_indices[embedding_name]\n            \n            # Skip embeddings that don't have the get_entity_embedding method\n            if not hasattr(embedding, 'get_entity_embedding'):\n                # Fill with zeros to maintain consistent shape\n                zeros = np.zeros(embedding.embedding_size)\n                all_values.append(zeros)\n                all_columns.extend(embedding_info['columns'])\n                continue\n                \n            entity_embedding = None\n            \n            # Handle target embeddings\n            if hasattr(embedding, 'cell_type_dict') and cell_type is not None:\n                key = f\"cell_type_{cell_type}\"\n                ct_embedding = embedding.get_entity_embedding(key)\n                individual_embeddings[f\"{embedding_name}_cell_type\"] = ct_embedding\n                entity_embedding = ct_embedding\n                \n            if hasattr(embedding, 'sm_dict') and sm_name is not None:\n                key = f\"sm_name_{sm_name}\"\n                sm_embedding = embedding.get_entity_embedding(key)\n                individual_embeddings[f\"{embedding_name}_sm\"] = sm_embedding\n                entity_embedding = sm_embedding if entity_embedding is None else np.concatenate([entity_embedding, sm_embedding])\n            \n            # Handle SMILES-based embeddings\n            elif sm_name is not None and embedding_name in ['SMILESEmbedding', 'MorganFingerPrintEmbedding', 'Mol2VecEmbedding']:\n                if sm_name in self.sm_name_to_smiles:\n                    smiles = self.sm_name_to_smiles[sm_name]\n                    entity_embedding = embedding.get_entity_embedding(smiles)\n                    individual_embeddings[embedding_name] = entity_embedding\n                else:\n                    entity_embedding = np.zeros(embedding.embedding_size)\n            \n            # If no embedding was found, use zeros to maintain shape\n            if entity_embedding is None:\n                entity_embedding = np.zeros(embedding.embedding_size)\n            \n            all_values.append(entity_embedding)\n            all_columns.extend(embedding_info['columns'])\n        \n        # Combine all embeddings maintaining order\n        combined_vector = np.concatenate(all_values)\n        \n        # Create DataFrame with proper column names\n        result_df = pd.DataFrame([combined_vector], columns=all_columns)\n        \n        return result_df, individual_embeddings\n    \n    def generate_expert_config(self, expert_specs):\n        \"\"\"\n        Generate expert configuration based on embedding combinations.\n        \n        Args:\n            expert_specs (dict): Dictionary mapping expert names to lists of embedding names\n            \n        Returns:\n            dict: Expert configuration with corresponding feature indices\n        \"\"\"\n        expert_config = {}\n        for expert, embedding_names in expert_specs.items():\n            indices = []\n            for embedding_name in embedding_names:\n                if embedding_name in self.embedding_indices:\n                    indices.extend(range(\n                        self.embedding_indices[embedding_name]['start'],\n                        self.embedding_indices[embedding_name]['end']\n                    ))\n                else:\n                    raise KeyError(f\"Embedding '{embedding_name}' not found in preprocessor\")\n            expert_config[expert] = indices\n        return expert_config\n    \n    def get_embedding_info(self):\n        \"\"\"\n        Get information about available embeddings and their ranges.\n        \n        Returns:\n            dict: Dictionary containing embedding information including column names\n        \"\"\"\n        return {\n            embedding_name: {\n                'size': self.embedding_indices[embedding_name]['end'] - \n                       self.embedding_indices[embedding_name]['start'],\n                'start_index': self.embedding_indices[embedding_name]['start'],\n                'end_index': self.embedding_indices[embedding_name]['end'],\n                'columns': self.embedding_indices[embedding_name]['columns']\n            }\n            for embedding_name in self.embedding_indices\n        }\n    \n    def get_smiles_mapping(self):\n        \"\"\"\n        Get the dictionary mapping sm_names to SMILES strings.\n        \n        Returns:\n            dict: Dictionary containing sm_name to SMILES mapping\n        \"\"\"\n        return self.sm_name_to_smiles.copy()","metadata":{"execution":{"iopub.status.busy":"2025-01-19T18:41:21.131525Z","iopub.execute_input":"2025-01-19T18:41:21.131789Z","iopub.status.idle":"2025-01-19T18:41:21.148966Z","shell.execute_reply.started":"2025-01-19T18:41:21.131763Z","shell.execute_reply":"2025-01-19T18:41:21.148026Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# onehot_embedding = OneHotEmbedding(columns=['cell_type', 'sm_name'])\ntarget_embedding = TargetEmbedding(df)\nsmiles_embedding = SMILESEmbedding()\nmol2vec_embedding = Mol2VecEmbedding('/kaggle/input/m/ankushhv/mol2vec/pytorch/default/1/model_300dim.pkl')\nmorgan_embedding = MorganFingerPrintEmbedding(hidden_size=256)","metadata":{"execution":{"iopub.status.busy":"2025-01-19T18:41:21.149782Z","iopub.execute_input":"2025-01-19T18:41:21.150055Z","iopub.status.idle":"2025-01-19T18:41:33.178291Z","shell.execute_reply.started":"2025-01-19T18:41:21.150031Z","shell.execute_reply":"2025-01-19T18:41:33.177131Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"preprocessor_train = Preprocessor([target_embedding, smiles_embedding, mol2vec_embedding,morgan_embedding])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-19T18:41:33.179567Z","iopub.execute_input":"2025-01-19T18:41:33.179816Z","iopub.status.idle":"2025-01-19T18:41:33.183758Z","shell.execute_reply.started":"2025-01-19T18:41:33.179790Z","shell.execute_reply":"2025-01-19T18:41:33.182953Z"}},"outputs":[],"execution_count":null},{"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":"2025-01-19T18:41:33.186287Z","iopub.execute_input":"2025-01-19T18:41:33.186543Z","iopub.status.idle":"2025-01-19T18:43:16.343633Z","shell.execute_reply.started":"2025-01-19T18:41:33.186518Z","shell.execute_reply":"2025-01-19T18:43:16.342690Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"preprocessor_train.embedding_indices.keys()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-19T18:43:16.344799Z","iopub.execute_input":"2025-01-19T18:43:16.345245Z","iopub.status.idle":"2025-01-19T18:43:16.351124Z","shell.execute_reply.started":"2025-01-19T18:43:16.345217Z","shell.execute_reply":"2025-01-19T18:43:16.350247Z"}},"outputs":[],"execution_count":null},{"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)\nfeature_names = list(processed_df_train.columns)\nprint(len(feature_names))\ntrain_array = processed_df_train.to_numpy()\nprint(targets.dtypes)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-19T18:43:16.352387Z","iopub.execute_input":"2025-01-19T18:43:16.352777Z","iopub.status.idle":"2025-01-19T18:43:16.367298Z","shell.execute_reply.started":"2025-01-19T18:43:16.352738Z","shell.execute_reply":"2025-01-19T18:43:16.366372Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **Preparing Dataloaders**","metadata":{}},{"cell_type":"code","source":"import torch\nfrom torch.utils.data import DataLoader, Dataset, TensorDataset\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}\")\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)\n\ndata_save_directory = 'data'\nos.makedirs(data_save_directory, exist_ok=True)\n\ntorch.save(X, os.path.join(data_save_directory, 'X_tensor.pt'))\ntorch.save(y, os.path.join(data_save_directory, 'y_tensor.pt'))\n\nprint(f\"Tensors saved in {data_save_directory} directory\")\n\n# X_loaded = torch.load(os.path.join(data_save_directory, 'X_tensor.pt'))\n# y_loaded = torch.load(os.path.join(data_save_directory, 'y_tensor.pt'))\n\n# print(\"Tensors loaded successfully\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-19T18:43:16.368356Z","iopub.execute_input":"2025-01-19T18:43:16.368623Z","iopub.status.idle":"2025-01-19T18:43:16.678658Z","shell.execute_reply.started":"2025-01-19T18:43:16.368595Z","shell.execute_reply":"2025-01-19T18:43:16.677725Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **Model Architecture**","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\nfrom sklearn.feature_selection import SelectKBest, f_regression\n\nclass LSTMExpert(nn.Module):\n    def __init__(self, input_size=152, output_size=18211):\n        super(LSTMExpert, self).__init__()\n        self.name = 'LSTMExpert'\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        out, _ = self.lstm(x)\n        out = out[:, -1, :]\n        out = self.linear(out)\n        out = self.head1(out)\n        return out\n\nclass GatingNetwork(nn.Module):\n    def __init__(self, input_size, num_experts, hidden_size=512, initial_temp=2.0, \n                 warmup_epochs=10, sparsity_factor=0.1, top_k=2):\n        super(GatingNetwork, self).__init__()\n        self.input_size = input_size\n        self.num_experts = num_experts\n        self.top_k = top_k\n        self.sparsity_factor = sparsity_factor\n        \n        # Feature extraction with residual connections\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        # Expert-specific feature extractors\n        self.expert_specific_layers = nn.ModuleList([\n            nn.Linear(hidden_size, hidden_size // 2) \n            for _ in range(num_experts)\n        ])\n        \n        # Multi-head attention\n        self.attention = nn.MultiheadAttention(hidden_size, num_heads=4)\n        \n        # Output layer with larger initial weights\n        self.output_layer = nn.Linear(hidden_size, num_experts)\n        with torch.no_grad():\n            self.output_layer.weight.data *= 2.0\n        \n        # Learnable temperature parameter with constraints\n        self.log_temperature = nn.Parameter(torch.log(torch.ones(1) * initial_temp))\n        self.min_temperature = 0.1\n        self.max_temperature = 10.0\n        \n        # Warmup parameters\n        self.warmup_epochs = warmup_epochs\n        self.current_epoch = 0\n        \n        # Load balancing parameters\n        self.importance_scores = nn.Parameter(torch.ones(num_experts))\n        \n    @property\n    def temperature(self):\n        return torch.clamp(torch.exp(self.log_temperature), \n                         self.min_temperature, \n                         self.max_temperature)\n    \n    def compute_load_balancing_loss(self, gates):\n        # Compute the fraction of data assigned to each expert\n        expert_usage = gates.mean(dim=0)\n        # Ideal uniform distribution\n        uniform_usage = torch.ones_like(expert_usage) / self.num_experts\n        # KL divergence from uniform\n        load_balancing_loss = F.kl_div(\n            expert_usage.log(), uniform_usage, \n            reduction='batchmean'\n        )\n        return load_balancing_loss\n    \n    def forward(self, x, expert_performances=None):\n        batch_size = x.size(0)\n        \n        # Extract base features\n        features = self.feature_extractor(x)\n        \n        # Add sequence dimension for attention\n        features = features.unsqueeze(0)\n        \n        # Apply attention\n        attn_output, _ = self.attention(features, features, features)\n        attn_output = attn_output.squeeze(0)\n        \n        # Apply expert-specific processing\n        expert_features = []\n        for expert_layer in self.expert_specific_layers:\n            expert_feat = expert_layer(attn_output)\n            expert_features.append(expert_feat)\n        \n        expert_features = torch.stack(expert_features, dim=1)\n        \n        # Generate logits with importance scores\n        logits = self.output_layer(attn_output) * self.importance_scores\n        \n        # Apply temperature scaling\n        logits = logits / self.temperature\n        \n        # During warmup, return equal weights for all experts\n        if self.current_epoch < self.warmup_epochs:\n            return torch.ones_like(logits) / self.num_experts\n        \n        # After warmup, incorporate expert performances if provided\n        if expert_performances is not None:\n            # Normalize expert performances\n            normalized_performances = F.softmax(expert_performances, dim=-1)\n            logits = logits * normalized_performances\n        \n        # Apply top-k gating if specified\n        if self.top_k is not None:\n            # Get top-k values and set others to negative infinity\n            top_k_values, _ = torch.topk(logits, k=self.top_k, dim=-1)\n            kth_values = top_k_values[:, -1].unsqueeze(-1)\n            logits = torch.where(logits < kth_values, \n                               torch.full_like(logits, float('-inf')), \n                               logits)\n        \n        # Apply softmax with optional sparsity encouragement\n        gates = F.softmax(logits, dim=-1)\n        \n        if self.sparsity_factor > 0:\n            # Add sparsity regularization\n            gates = gates * (1 + self.sparsity_factor * (gates > gates.mean(dim=1, keepdim=True)).float())\n            gates = gates / gates.sum(dim=1, keepdim=True)  # Renormalize\n        \n        return gates\n    \n    def get_regularization_loss(self, gates):\n        # Combine different regularization terms\n        load_balance_loss = self.compute_load_balancing_loss(gates)\n        sparsity_loss = -torch.mean(torch.sum(gates * torch.log(gates + 1e-10), dim=1))\n        \n        return load_balance_loss + self.sparsity_factor * sparsity_loss\n    \n    def increment_epoch(self):\n        self.current_epoch += 1\n\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, y_true=None, return_info=False):\n        # Ensure all components are on the same device as the input data\n        device = data.device\n        self.to(device)\n\n        # Prepare input for gating network\n        gating_input = data[:, self.total_features]\n        \n        expert_outputs = []\n        expert_performances = []\n\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            if y_true is not None:\n                # Calculate instance-wise performance for each expert\n                expert_performance = 1 / (torch.sum((expert_output - y_true) ** 2, dim=1) + 1e-5)\n                expert_performances.append(expert_performance)\n\n        expert_outputs = torch.stack(expert_outputs, dim=1)\n        \n        if y_true is not None:\n            expert_performances = torch.stack(expert_performances, dim=1)\n            gating_weights = self.gating_network(gating_input, expert_performances)\n        else:\n            gating_weights = self.gating_network(gating_input)\n\n        self.last_gating_weights = gating_weights.detach()\n\n        combined_output = torch.sum(expert_outputs * gating_weights.unsqueeze(-1), dim=1)\n\n        if return_info:\n            return combined_output, expert_outputs, gating_weights\n        else:\n            return combined_output\n\n    def save(self, directory):\n        \"\"\"\n        Save the model to the specified directory.\n    \n        Parameters:\n            directory (str): The directory to save the model.\n        \"\"\"\n        os.makedirs(directory, exist_ok=True)\n    \n        # Save model state\n        model_state_path = os.path.join(directory, 'model_state.pth')\n        torch.save(self.state_dict(), model_state_path)\n    \n        # Save expert configurations and total features\n        config_path = os.path.join(directory, 'model_config.pth')\n        torch.save({\n            'expert_configs': self.feature_indices,\n            'total_features': self.total_features\n        }, config_path)\n    \n        print(f\"Model saved to directory {directory}\")\n\n    @classmethod\n    def load(cls, directory, device=None):\n        \"\"\"\n        Load the model from the specified directory.\n    \n        Parameters:\n            directory (str): The directory containing the model files.\n            device (torch.device or str, optional): The device to load the model onto.\n    \n        Returns:\n            MixtureOfExperts: The loaded model instance.\n        \"\"\"\n        # Paths to model state and configuration\n        model_state_path = os.path.join(directory, 'model_state.pth')\n        config_path = os.path.join(directory, 'model_config.pth')\n    \n        # Load configuration\n        config = torch.load(config_path, map_location=device)\n        expert_configs = config['expert_configs']\n        total_features = config['total_features']\n    \n        # Reconstruct the model\n        loaded_model = cls(expert_configs=expert_configs)\n        loaded_model.total_features = total_features\n    \n        # Load model state\n        model_state = torch.load(model_state_path, map_location=device)\n        loaded_model.load_state_dict(model_state)\n    \n        if device:\n            loaded_model.to(device)\n    \n        print(f\"Model loaded from directory {directory}\")\n        return loaded_model\n        \n\ndef RMSE_rowwise_loss(y_pred, y_true):\n    return torch.sqrt(torch.mean((y_pred - y_true)**2, dim=1)).mean()\n\ndef RMSE_rowwise_loss_numpy(y_pred, y_true):\n    \"\"\"Calculate row-wise RMSE loss using numpy arrays\"\"\"\n    # Ensure inputs are numpy arrays\n    if torch.is_tensor(y_pred):\n        y_pred = y_pred.detach().cpu().numpy()\n    if torch.is_tensor(y_true):\n        y_true = y_true.detach().cpu().numpy()\n    \n    # Convert to numpy arrays if they're not already\n    y_pred = np.array(y_pred)\n    y_true = np.array(y_true)\n    \n    # Calculate RMSE\n    row_mse = np.mean((y_pred - y_true)**2, axis=1)\n    row_rmse = np.sqrt(row_mse)\n    return float(np.mean(row_rmse))\n\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":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-19T18:43:16.679969Z","iopub.execute_input":"2025-01-19T18:43:16.680245Z","iopub.status.idle":"2025-01-19T18:43:16.957543Z","shell.execute_reply.started":"2025-01-19T18:43:16.680218Z","shell.execute_reply":"2025-01-19T18:43:16.956662Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **Training Logic**","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.optim as optim\nfrom tqdm import tqdm\nimport numpy as np\n\ndef compute_loss(output, target, gating_weights, model, reg_weight):\n        # Compute main loss with gradient clipping\n        main_loss = RMSE_rowwise_loss(output, target)\n        \n        # Check for NaN in main loss\n        if torch.isnan(main_loss):\n            print(\"Warning: NaN detected in main loss\")\n            return None\n            \n        # Compute regularization loss with safe handling\n        try:\n            reg_loss = model.gating_network.get_regularization_loss(gating_weights)\n            if torch.isnan(reg_loss):\n                print(\"Warning: NaN detected in regularization loss\")\n                reg_loss = torch.tensor(0.0, device=device)\n        except Exception as e:\n            print(f\"Warning: Error in regularization loss computation: {e}\")\n            reg_loss = torch.tensor(0.0, device=device)\n            \n        total_loss = main_loss + reg_weight * reg_loss\n        return total_loss\n\ndef train_mixture_of_experts(model, train_loader, val_loader, device, \n                           epochs=100, lr=1e-4, plot_interval=10, \n                           gradient_clip_val=1.0, reg_weight=0.1):\n    model.to(device)\n    \n    # Initialize optimizer with gradient clipping\n    optimizer = optim.Adam(model.parameters(), lr=lr, weight_decay=1e-5)\n    scheduler = optim.lr_scheduler.ReduceLROnPlateau(\n        optimizer, mode='min', patience=10, factor=0.5, min_lr=1e-6\n    )\n    early_stopping = EarlyStopping(patience=10, verbose=True)\n    \n    train_losses = []\n    val_losses = []\n    gating_weights_history = []\n    \n    for epoch in tqdm(range(epochs), desc=\"Training Mixture of Experts\"):\n        model.train()\n        train_loss = 0\n        epoch_gating_weights = []\n        num_batches = 0\n        \n        # Training loop\n        for batch, (X, y) in enumerate(train_loader):\n            try:\n                X, y = X.to(device), y.to(device)\n                \n                # Check for NaN in input\n                if torch.isnan(X).any() or torch.isnan(y).any():\n                    print(f\"Warning: NaN detected in input data at batch {batch}\")\n                    continue\n                    \n                optimizer.zero_grad()\n                \n                # Forward pass with gradient computation\n                with torch.set_grad_enabled(True):\n                    output, expert_outputs, gating_weights = model(X, y, return_info=True)\n                    \n                    # Check outputs for NaN\n                    if torch.isnan(output).any():\n                        print(f\"Warning: NaN detected in model output at batch {batch}\")\n                        continue\n                        \n                    loss = compute_loss(output, y, gating_weights, model, reg_weight)\n                    if loss is None:\n                        continue\n                        \n                    # Backward pass with gradient clipping\n                    loss.backward()\n                    torch.nn.utils.clip_grad_norm_(model.parameters(), gradient_clip_val)\n                    \n                    # Check for NaN gradients\n                    has_nan_grad = False\n                    for param in model.parameters():\n                        if param.grad is not None and torch.isnan(param.grad).any():\n                            has_nan_grad = True\n                            break\n                    \n                    if has_nan_grad:\n                        print(f\"Warning: NaN detected in gradients at batch {batch}\")\n                        optimizer.zero_grad()\n                        continue\n                        \n                    optimizer.step()\n                    \n                    train_loss += loss.item()\n                    epoch_gating_weights.append(gating_weights.mean(dim=0).cpu().detach().numpy())\n                    num_batches += 1\n                    \n            except RuntimeError as e:\n                print(f\"Error in batch {batch}: {e}\")\n                continue\n        \n        # Compute average training loss\n        if num_batches > 0:\n            avg_train_loss = train_loss / num_batches\n            train_losses.append(avg_train_loss)\n            gating_weights_history.append(np.mean(epoch_gating_weights, axis=0))\n        else:\n            print(\"Warning: No valid batches in epoch\")\n            continue\n            \n        # Validation loop\n        model.eval()\n        val_loss = 0\n        val_batches = 0\n        \n        with torch.no_grad():\n            for X, y in val_loader:\n                try:\n                    X, y = X.to(device), y.to(device)\n                    output = model(X)\n                    val_batch_loss = RMSE_rowwise_loss(output, y).item()\n                    \n                    if not np.isnan(val_batch_loss):\n                        val_loss += val_batch_loss\n                        val_batches += 1\n                        \n                except RuntimeError as e:\n                    print(f\"Error in validation: {e}\")\n                    continue\n        \n        if val_batches > 0:\n            avg_val_loss = val_loss / val_batches\n            val_losses.append(avg_val_loss)\n            \n            # Model checkpointing\n            if avg_val_loss < getattr(early_stopping, 'val_loss_min', float('inf')):\n                best_model_state = model.state_dict()\n                \n            # Learning rate scheduling\n            scheduler.step(avg_val_loss)\n            \n            if early_stopping(avg_val_loss):\n                print(\"Early stopping triggered\")\n                model.load_state_dict(best_model_state)\n                break\n                \n            # Logging and plotting\n            if epoch % plot_interval == 0 or epoch == epochs - 1:\n                tqdm.write(f\"Epoch {epoch}: Train Loss: {avg_train_loss:.4f}, \"\n                          f\"Val Loss: {avg_val_loss:.4f}, \"\n                          f\"LR: {optimizer.param_groups[0]['lr']:.2e}\")\n                \n                # Plot training curves\n                plot_training_curves(train_losses, val_losses, \n                                  f\"Training Curves (Epoch {epoch})\", \n                                  save_path=f\"training_curves_epoch_{epoch}.png\")\n                \n                # Plot expert utilization\n                plot_expert_utilization_over_time(\n                    np.array(gating_weights_history), \n                    save_path=f\"expert_utilization_over_time_epoch_{epoch}.png\"\n                )\n        \n        model.gating_network.increment_epoch()\n        \n    # Final plots\n    plot_training_curves(train_losses, val_losses, \"Final Training Curves\", \n                        save_path=\"final_training_curves.png\")\n    \n    plot_expert_utilization_over_time(np.array(gating_weights_history), \n                                    save_path=\"final_expert_utilization_over_time.png\")\n    \n    return train_losses, val_losses, gating_weights_history\n\ndef train_mixture_of_experts_em(model, train_loader, val_loader, device, \n                              epochs=100, lr=1e-4, plot_interval=10,\n                              gradient_clip_val=1.0):\n    model.to(device)\n    \n    # Initialize optimizer with gradient clipping\n    optimizer = optim.Adam(model.parameters(), lr=lr, weight_decay=1e-5)\n    scheduler = optim.lr_scheduler.ReduceLROnPlateau(\n        optimizer, mode='min', patience=10, factor=0.5, min_lr=1e-6\n    )\n    early_stopping = EarlyStopping(patience=10, verbose=True)\n    \n    train_losses = []\n    val_losses = []\n    gating_weights_history = []\n    best_model_state = None\n    \n    for epoch in tqdm(range(epochs), desc=\"Training MoE with EM\"):\n        model.train()\n        train_loss = 0\n        epoch_gating_weights = []\n        num_batches = 0\n        \n        # Training loop\n        for batch, (X, y) in enumerate(train_loader):\n            try:\n                X, y = X.to(device), y.to(device)\n                \n                # Check for NaN in input\n                if torch.isnan(X).any() or torch.isnan(y).any():\n                    print(f\"Warning: NaN detected in input data at batch {batch}\")\n                    continue\n                \n                optimizer.zero_grad()\n                \n                # E-step: Compute responsibilities\n                model.eval()\n                with torch.no_grad():\n                    output, expert_outputs, gating_weights = model(X, y, return_info=True)\n                    \n                    # Check outputs for NaN\n                    if torch.isnan(output).any():\n                        print(f\"Warning: NaN detected in model output at batch {batch}\")\n                        continue\n                    \n                    # Compute expert-wise losses\n                    expert_losses = torch.zeros((X.size(0), len(model.experts)), device=device)\n                    for i, expert_output in enumerate(expert_outputs.unbind(1)):\n                        expert_losses[:, i] = torch.sum((expert_output - y) ** 2, dim=1)\n                    \n                    # Compute responsibilities with temperature scaling\n                    responsibilities = F.softmax(-expert_losses / model.gating_network.temperature, dim=1)\n                \n                # M-step: Update model parameters\n                model.train()\n                \n                # Get fresh forward pass in training mode\n                output, expert_outputs, gating_weights = model(X, y, return_info=True)\n                \n                # Update experts using weighted losses\n                expert_loss = 0\n                for i, expert_output in enumerate(expert_outputs.unbind(1)):\n                    expert_loss += torch.mean(responsibilities[:, i].unsqueeze(1) * \n                                            torch.sum((expert_output - y) ** 2, dim=1))\n                \n                # Update gating network\n                gating_loss = F.kl_div(torch.log(gating_weights + 1e-10), \n                                     responsibilities.detach(), \n                                     reduction='batchmean')\n                \n                # Combined loss\n                loss = expert_loss + gating_loss\n                \n                # Backward pass with gradient clipping\n                loss.backward()\n                torch.nn.utils.clip_grad_norm_(model.parameters(), gradient_clip_val)\n                \n                # Check for NaN gradients\n                has_nan_grad = False\n                for param in model.parameters():\n                    if param.grad is not None and torch.isnan(param.grad).any():\n                        has_nan_grad = True\n                        break\n                \n                if has_nan_grad:\n                    print(f\"Warning: NaN detected in gradients at batch {batch}\")\n                    optimizer.zero_grad()\n                    continue\n                \n                optimizer.step()\n                \n                train_loss += loss.item()\n                epoch_gating_weights.append(gating_weights.mean(dim=0).cpu().detach().numpy())\n                num_batches += 1\n                \n            except RuntimeError as e:\n                print(f\"Error in batch {batch}: {e}\")\n                continue\n        \n        # Compute average training loss\n        if num_batches > 0:\n            avg_train_loss = train_loss / num_batches\n            train_losses.append(avg_train_loss)\n            gating_weights_history.append(np.mean(epoch_gating_weights, axis=0))\n        else:\n            print(\"Warning: No valid batches in epoch\")\n            continue\n        \n        # Validation loop\n        model.eval()\n        val_loss = 0\n        val_batches = 0\n        \n        with torch.no_grad():\n            for X, y in val_loader:\n                try:\n                    X, y = X.to(device), y.to(device)\n                    output = model(X)\n                    val_batch_loss = RMSE_rowwise_loss(output, y).item()\n                    \n                    if not np.isnan(val_batch_loss):\n                        val_loss += val_batch_loss\n                        val_batches += 1\n                        \n                except RuntimeError as e:\n                    print(f\"Error in validation: {e}\")\n                    continue\n        \n        if val_batches > 0:\n            avg_val_loss = val_loss / val_batches\n            val_losses.append(avg_val_loss)\n            \n            # Model checkpointing\n            if avg_val_loss < getattr(early_stopping, 'val_loss_min', float('inf')):\n                best_model_state = model.state_dict()\n            \n            # Learning rate scheduling\n            scheduler.step(avg_val_loss)\n            \n            # Early stopping check\n            if early_stopping(avg_val_loss):\n                print(\"Early stopping triggered\")\n                if best_model_state is not None:\n                    model.load_state_dict(best_model_state)\n                break\n            \n            # Logging and plotting\n            if epoch % plot_interval == 0 or epoch == epochs - 1:\n                tqdm.write(f\"Epoch {epoch}: Train Loss: {avg_train_loss:.4f}, \"\n                          f\"Val Loss: {avg_val_loss:.4f}, \"\n                          f\"LR: {optimizer.param_groups[0]['lr']:.2e}\")\n                \n                # Plot training curves\n                plot_training_curves(train_losses, val_losses, \n                                  f\"Training Curves (Epoch {epoch})\", \n                                  save_path=f\"training_curves_em_epoch_{epoch}.png\")\n                \n                # Plot expert utilization\n                plot_expert_utilization_over_time(\n                    np.array(gating_weights_history), \n                    save_path=f\"expert_utilization_em_over_time_epoch_{epoch}.png\"\n                )\n        \n        # Update gating network epoch counter and temperature\n        model.gating_network.increment_epoch()\n        current_temp = model.gating_network.temperature\n        model.gating_network.log_temperature.data = torch.log(torch.tensor(max(0.5, current_temp * 0.99)))\n    \n    # Final plots\n    plot_training_curves(train_losses, val_losses, \"Final Training Curves\", \n                        save_path=\"final_training_curves_em.png\")\n    \n    plot_expert_utilization_over_time(np.array(gating_weights_history), \n                                    save_path=\"final_expert_utilization_over_time_em.png\")\n    \n    return train_losses, val_losses, gating_weights_history","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-19T18:43:16.959016Z","iopub.execute_input":"2025-01-19T18:43:16.959290Z","iopub.status.idle":"2025-01-19T18:43:16.992544Z","shell.execute_reply.started":"2025-01-19T18:43:16.959264Z","shell.execute_reply":"2025-01-19T18:43:16.991908Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **Hybrid Training** ","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom tqdm import tqdm\nimport numpy as np\n\ndef pretrain_experts(model, train_loader, val_loader, device, epochs=50, lr=1e-4):\n    \"\"\"Pretrain each expert independently before MoE training\"\"\"\n    print(\"Pretraining experts independently...\")\n    \n    original_states = {name: expert.state_dict() for name, expert in model.experts.items()}\n    best_states = {}\n    expert_performances = {}\n    \n    # Store training history for each expert\n    expert_histories = {}\n    \n    for expert_name, expert in model.experts.items():\n        print(f\"\\nPretraining {expert_name}\")\n        expert.train()\n        optimizer = optim.Adam(expert.parameters(), lr=lr)\n        scheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, patience=5)\n        best_val_loss = float('inf')\n        \n        # Training history for this expert\n        train_losses = []\n        val_losses = []\n        \n        feature_indices = model.feature_indices[expert_name]\n        \n        for epoch in range(epochs):\n            expert.train()\n            train_loss = 0\n            batch_count = 0\n            \n            # Training loop\n            for batch_idx, (X, y) in enumerate(train_loader):\n                X, y = X.to(device), y.to(device)\n                X_expert = X[:, feature_indices]\n                \n                optimizer.zero_grad()\n                output = expert(X_expert)\n                loss = RMSE_rowwise_loss(output, y)\n                loss.backward()\n                optimizer.step()\n                train_loss += loss.item()\n                batch_count += 1\n            \n            avg_train_loss = train_loss / batch_count\n            train_losses.append(avg_train_loss)\n            \n            # Validation loop\n            expert.eval()\n            val_loss = 0\n            batch_count = 0\n            with torch.no_grad():\n                for X, y in val_loader:\n                    X, y = X.to(device), y.to(device)\n                    X_expert = X[:, feature_indices]\n                    output = expert(X_expert)\n                    val_loss += RMSE_rowwise_loss(output, y).item()\n                    batch_count += 1\n            \n            avg_val_loss = val_loss / batch_count\n            val_losses.append(avg_val_loss)\n            \n            if avg_val_loss < best_val_loss:\n                best_val_loss = avg_val_loss\n                best_states[expert_name] = expert.state_dict()\n            \n            if epoch % 10 == 0:\n                print(f\"Epoch {epoch}: Train Loss = {avg_train_loss:.4f}, Val Loss = {avg_val_loss:.4f}\")\n            \n            scheduler.step(avg_val_loss)\n        \n        # Plot training curves for this expert\n        plot_training_curves(\n            train_losses, \n            val_losses, \n            f\"Training Curves for {expert_name}\",\n            save_path=f\"expert_{expert_name}_training.png\"\n        )\n        \n        expert_performances[expert_name] = best_val_loss\n        expert_histories[expert_name] = {'train': train_losses, 'val': val_losses}\n        expert.load_state_dict(best_states[expert_name])\n    \n    return best_states, expert_performances, expert_histories\n\n\ndef train_gating_with_pretrained_experts(model, train_loader, val_loader, device, \n                                       expert_performances, epochs=100, lr=1e-4):\n    \"\"\"Train gating network with frozen pretrained experts\"\"\"\n    print(\"\\nTraining gating network with frozen experts...\")\n    \n    # Freeze expert parameters\n    for expert in model.experts.values():\n        for param in expert.parameters():\n            param.requires_grad = False\n    \n    optimizer = optim.Adam(model.gating_network.parameters(), lr=lr)\n    scheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, patience=5)\n    early_stopping = EarlyStopping(patience=10)\n    \n    best_val_loss = float('inf')\n    best_gating_state = None\n    \n    # Training history\n    train_losses = []\n    val_losses = []\n    gating_weights_history = []\n    \n    for epoch in range(epochs):\n        model.train()\n        train_loss = 0\n        batch_count = 0\n        epoch_gating_weights = []\n        \n        for batch_idx, (X, y) in enumerate(train_loader):\n            X, y = X.to(device), y.to(device)\n            \n            # First, get expert outputs and compute responsibilities\n            model.eval()  # Set model to eval mode for getting expert outputs\n            with torch.no_grad():\n                _, expert_outputs, _ = model(X, return_info=True)\n                expert_losses = torch.zeros((X.size(0), len(model.experts)), device=device)\n                \n                for i, expert_output in enumerate(expert_outputs.unbind(1)):\n                    expert_losses[:, i] = torch.sum((expert_output - y) ** 2, dim=1)\n                    expert_name = list(model.experts.keys())[i]\n                    expert_losses[:, i] *= expert_performances[expert_name]\n                \n                responsibilities = F.softmax(-expert_losses / model.gating_network.temperature, dim=1)\n            \n            # Now switch to training mode for gating network update\n            model.train()\n            optimizer.zero_grad()\n            \n            # Forward pass through gating network only\n            gating_input = X[:, model.total_features]\n            gating_weights = model.gating_network(gating_input)\n            \n            # Ensure gating_weights requires gradients\n            if not gating_weights.requires_grad:\n                gating_weights = gating_weights.detach().requires_grad_(True)\n            \n            # Compute KL divergence loss\n            log_gating_weights = torch.log(gating_weights + 1e-10)\n            gating_loss = F.kl_div(\n                log_gating_weights,\n                responsibilities.detach(),  # Detach responsibilities\n                reduction='batchmean'\n            )\n            \n            # Backward pass and optimization\n            gating_loss.backward()\n            optimizer.step()\n            \n            # Store metrics\n            train_loss += gating_loss.item()\n            epoch_gating_weights.append(gating_weights.detach().cpu().numpy().mean(axis=0))\n            batch_count += 1\n        \n        avg_train_loss = train_loss / batch_count\n        train_losses.append(avg_train_loss)\n        gating_weights_history.append(np.mean(epoch_gating_weights, axis=0))\n        \n        # Validation\n        model.eval()\n        val_loss = 0\n        batch_count = 0\n        with torch.no_grad():\n            for X, y in val_loader:\n                X, y = X.to(device), y.to(device)\n                output = model(X)\n                val_loss += RMSE_rowwise_loss(output, y).item()\n                batch_count += 1\n        \n        avg_val_loss = val_loss / batch_count\n        val_losses.append(avg_val_loss)\n        \n        if avg_val_loss < best_val_loss:\n            best_val_loss = avg_val_loss\n            best_gating_state = model.gating_network.state_dict()\n        \n        if epoch % 10 == 0:\n            print(f\"Epoch {epoch}: Train Loss = {avg_train_loss:.4f}, Val Loss = {avg_val_loss:.4f}, \"\n                  f\"Temp = {model.gating_network.log_temperature.item():.2f}\")\n        \n        # Update temperature\n        # current_temp = model.gating_network.get_temperature()\n        # model.gating_network.set_temperature(max(0.5, current_temp * 0.95))\n        \n        scheduler.step(avg_val_loss)\n        \n        if early_stopping(avg_val_loss):\n            print(\"Early stopping triggered\")\n            break\n    \n    # Plot training curves and expert utilization\n    plot_training_curves(\n        train_losses, \n        val_losses, \n        \"Gating Network Training Curves\",\n        save_path=\"gating_network_training.png\"\n    )\n    \n    plot_expert_utilization_over_time(\n        np.array(gating_weights_history),\n        save_path=\"expert_utilization_gating.png\"\n    )\n    \n    # Load best gating network state\n    model.gating_network.load_state_dict(best_gating_state)\n    \n    # Unfreeze experts for potential fine-tuning\n    for expert in model.experts.values():\n        for param in expert.parameters():\n            param.requires_grad = True\n    \n    return best_val_loss, {'train': train_losses, 'val': val_losses, 'gating_weights': gating_weights_history}\n\ndef fine_tune_moe(model, train_loader, val_loader, device, epochs=20, lr=1e-5):\n    \"\"\"Optional fine-tuning of entire model with small learning rate\"\"\"\n    print(\"\\nFine-tuning entire model...\")\n    \n    optimizer = optim.Adam(model.parameters(), lr=lr)\n    scheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, patience=5)\n    early_stopping = EarlyStopping(patience=10)\n    \n    # Training history\n    train_losses = []\n    val_losses = []\n    gating_weights_history = []\n    \n    for epoch in range(epochs):\n        model.train()\n        train_loss = 0\n        batch_count = 0\n        epoch_gating_weights = []\n        \n        for batch_idx, (X, y) in enumerate(train_loader):\n            X, y = X.to(device), y.to(device)\n            optimizer.zero_grad()\n            \n            output, _, gating_weights = model(X, return_info=True)\n            loss = RMSE_rowwise_loss(output, y)\n            loss.backward()\n            optimizer.step()\n            \n            train_loss += loss.item()\n            epoch_gating_weights.append(gating_weights.mean(dim=0).cpu().detach().numpy())\n            batch_count += 1\n        \n        avg_train_loss = train_loss / batch_count\n        train_losses.append(avg_train_loss)\n        gating_weights_history.append(np.mean(epoch_gating_weights, axis=0))\n        \n        # Validation\n        model.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 = model(X)\n                val_loss += RMSE_rowwise_loss(output, y).item()\n        \n        val_loss /= len(val_loader)\n        train_loss /= len(train_loader)\n        \n        # Store losses for plotting\n        train_losses.append(train_loss)\n        val_losses.append(val_loss)\n        \n        if epoch % 5 == 0:\n            print(f\"Epoch {epoch}: Val Loss = {val_loss:.4f}\")\n            plot_expert_utilization_over_time(np.array(gating_weights_history), f\"expert_weights_history_finetune_{epoch}\")\n        \n        scheduler.step(val_loss)\n        \n        if early_stopping(val_loss):\n            print(\"Early stopping triggered\")\n            break\n    \n    # Plot training curve\n    plot_training_curves(train_losses, val_losses, title='Fine-Tuning Training Curve')\n    \n    return val_loss\n\n# Update the main training function to handle the returned dictionary\ndef train_pretrained_moe(model, train_loader, val_loader, device, \n                        pretrain_epochs=50, gating_epochs=100, fine_tune_epochs=20):\n    \"\"\"Complete training pipeline for pretrained MoE\"\"\"\n    # 1. Pretrain experts\n    best_states, expert_performances, expert_histories = pretrain_experts(\n        model, train_loader, val_loader, device, epochs=pretrain_epochs\n    )\n    \n    # 2. Train gating network with frozen experts\n    gating_val_loss, gating_history = train_gating_with_pretrained_experts(\n        model, train_loader, val_loader, device, \n        expert_performances, epochs=gating_epochs\n    )\n    \n    # 3. Optional: Fine-tune entire model\n    final_val_loss = fine_tune_moe(\n        model, train_loader, val_loader, device, epochs=fine_tune_epochs\n    )\n    \n    training_history = {\n        'expert_histories': expert_histories,\n        'gating_history': gating_history,\n        'final_val_loss': final_val_loss\n    }\n    \n    return final_val_loss, training_history","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-19T18:43:16.993736Z","iopub.execute_input":"2025-01-19T18:43:16.994006Z","iopub.status.idle":"2025-01-19T18:43:17.022750Z","shell.execute_reply.started":"2025-01-19T18:43:16.993981Z","shell.execute_reply":"2025-01-19T18:43:17.021963Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **Monitoring Functions**","metadata":{}},{"cell_type":"code","source":"import seaborn as sns\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport pandas as pd\n\n# Set Seaborn style globally\nsns.set(style=\"whitegrid\", palette=\"muted\")\nsns.set_context(\"talk\")\n\ndef plot_training_curves(train_losses, val_losses, title, save_path=None):\n    plt.figure(figsize=(12, 6))\n    \n    # Ensure both lists have the same length\n    min_length = min(len(train_losses), len(val_losses))\n    train_losses = train_losses[:min_length]\n    val_losses = val_losses[:min_length]\n    epochs = list(range(1, min_length + 1))\n    \n    # Create a DataFrame for easier plotting\n    df = pd.DataFrame({\n        'Epoch': epochs + epochs,\n        'Loss': train_losses + val_losses,\n        'Type': ['Train'] * min_length + ['Validation'] * min_length\n    })\n    \n    # Use sns.lineplot for both raw data and moving average\n    sns.lineplot(data=df, x='Epoch', y='Loss', hue='Type', style='Type', markers=True, dashes=False)\n    \n    # Add moving average\n    window_size = min(5, min_length)  # Ensure window size is not larger than data length\n    for loss_type in ['Train', 'Validation']:\n        ma = df[df['Type'] == loss_type]['Loss'].rolling(window=window_size).mean()\n        plt.plot(epochs[window_size-1:], ma[window_size-1:], linestyle='--', \n                 label=f'{loss_type} MA', alpha=0.7)\n    \n    plt.title(title)\n    plt.xlabel('Epoch')\n    plt.ylabel('Loss')\n    plt.legend(title='', loc='center left', bbox_to_anchor=(1, 0.5))\n    plt.tight_layout()\n    \n    if save_path:\n        plt.savefig(save_path)\n    plt.show()\n\ndef plot_expert_utilization_over_time(gating_weights_history, save_path=None):\n    plt.figure(figsize=(12, 6))\n    \n    # Create a DataFrame for easier plotting\n    df = pd.DataFrame(gating_weights_history, columns=[f'Expert {i+1}' for i in range(gating_weights_history.shape[1])])\n    df['Epoch'] = range(1, len(df) + 1)\n    df_melted = df.melt('Epoch', var_name='Expert', value_name='Usage')\n    \n    # Use sns.lineplot with the melted DataFrame\n    sns.lineplot(data=df_melted, x='Epoch', y='Usage', hue='Expert', style='Expert', markers=True)\n    plt.title(\"Expert Utilization Over Time\")\n    plt.xlabel(\"Epoch\")\n    plt.ylabel(\"Usage\")\n    plt.legend(title='', loc='center left', bbox_to_anchor=(1, 0.5))\n    plt.grid(True, linestyle='--', alpha=0.7)\n    \n    plt.tight_layout()\n    if save_path:\n        plt.savefig(save_path)\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-19T18:43:17.024045Z","iopub.execute_input":"2025-01-19T18:43:17.024830Z","iopub.status.idle":"2025-01-19T18:43:17.221667Z","shell.execute_reply.started":"2025-01-19T18:43:17.024791Z","shell.execute_reply":"2025-01-19T18:43:17.220758Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Set Seaborn style globally\nsns.set(style=\"whitegrid\", palette=\"muted\")\nsns.set_context(\"talk\")\n\ndef plot_training_curves(train_losses, val_losses, title, save_path=None):\n    plt.figure(figsize=(12, 6))\n\n    # Ensure both lists have the same length\n    min_length = min(len(train_losses), len(val_losses))\n    train_losses = train_losses[:min_length]\n    val_losses = val_losses[:min_length]\n    epochs = list(range(1, min_length + 1))\n\n    # Create a DataFrame for easier plotting\n    df = pd.DataFrame({\n        'Epoch': epochs + epochs,\n        'Loss': train_losses + val_losses,\n        'Type': ['Train'] * min_length + ['Validation'] * min_length\n    })\n\n    # Use sns.lineplot for both raw data and moving average\n    sns.lineplot(data=df, x='Epoch', y='Loss', hue='Type', style='Type', markers=True, dashes=False)\n\n    # Add moving average\n    window_size = min(5, min_length)  # Ensure window size is not larger than data length\n    for loss_type in ['Train', 'Validation']:\n        ma = df[df['Type'] == loss_type]['Loss'].rolling(window=window_size).mean()\n        plt.plot(epochs[window_size-1:], ma[window_size-1:], linestyle='--',\n                 label=f'{loss_type} MA', alpha=0.7)\n\n    plt.title(title)\n    plt.xlabel('Epoch')\n    plt.ylabel('Loss')\n    plt.legend(title='', loc='center left', bbox_to_anchor=(1, 0.5))\n    plt.tight_layout()\n\n    if save_path:\n        plt.savefig(save_path)\n    plt.show()\n\ndef plot_expert_utilization_over_time(gating_weights_history, save_path=None):\n    plt.figure(figsize=(12, 6))\n\n    # Create a DataFrame for easier plotting\n    df = pd.DataFrame(gating_weights_history, columns=[f'Expert {i+1}' for i in range(gating_weights_history.shape[1])])\n    df['Epoch'] = range(1, len(df) + 1)\n    df_melted = df.melt('Epoch', var_name='Expert', value_name='Usage')\n\n    # Use sns.lineplot with the melted DataFrame\n    sns.lineplot(data=df_melted, x='Epoch', y='Usage', hue='Expert', style='Expert', markers=True)\n    plt.title(\"Expert Utilization Over Time\")\n    plt.xlabel(\"Epoch\")\n    plt.ylabel(\"Usage\")\n    plt.legend(title='', loc='center left', bbox_to_anchor=(1, 0.5))\n    plt.grid(True, linestyle='--', alpha=0.7)\n\n    plt.tight_layout()\n    if save_path:\n        plt.savefig(save_path)\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-19T18:43:17.222940Z","iopub.execute_input":"2025-01-19T18:43:17.223330Z","iopub.status.idle":"2025-01-19T18:43:17.234004Z","shell.execute_reply.started":"2025-01-19T18:43:17.223303Z","shell.execute_reply":"2025-01-19T18:43:17.233047Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from 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', 'shapley']):\n    \"\"\"\n    Compute and visualize feature importance using various interpretability methods.\n\n    Args:\n        model (torch.nn.Module): The model to analyze.\n        X (torch.Tensor): The input samples.\n        feature_names (list): The names of the features.\n        k (int): The number of top features to plot.\n        baseline (torch.Tensor): The baseline input for the interpretability methods. If None, it will be set to zero.\n        target (int): The target class index for the interpretability methods.\n        methods (list): The list of interpretability methods to use.\n\n    Returns:\n        dict: A dictionary containing the top features, their importances, and the full importances for the model and gating network.\n    \"\"\"\n    model.eval()\n\n    if baseline is None:\n        baseline = torch.zeros_like(X)\n\n    importances = {}\n    for method in methods:\n        # Set up the explainer based on method\n        if method == 'integrated_gradients':\n            explainer = IntegratedGradients(model)\n            gating_explainer = IntegratedGradients(model.gating_network)\n        elif method == 'saliency':\n            explainer = Saliency(model)\n            gating_explainer = Saliency(model.gating_network)\n        elif method == 'deeplift':\n            explainer = DeepLift(model)\n            gating_explainer = DeepLift(model.gating_network)\n        elif method == 'shapley':\n            explainer = ShapleyValueSampling(model)\n            gating_explainer = ShapleyValueSampling(model.gating_network)\n        else:\n            raise ValueError(f\"Unsupported method: {method}\")\n\n        # Compute attributions\n        attributions = explainer.attribute(X,\n                                         baselines=baseline if method != 'saliency' else None,\n                                         target=target)\n        gating_attributions = gating_explainer.attribute(X, baselines=baseline, target=target)\n\n        # Process attributions\n        attributions_np = attributions.cpu().detach().numpy()\n        gating_attributions_np = gating_attributions.cpu().detach().numpy()\n        avg_importance = attributions_np.mean(axis=0)\n        gating_avg_importance = gating_attributions_np.mean(axis=0)\n        top_k_indices = np.argsort(np.abs(avg_importance))[-k:][::-1]\n        gating_top_k_indices = np.argsort(np.abs(gating_avg_importance))[-k:][::-1]\n\n        # Create DataFrames for plotting\n        plot_data = pd.DataFrame({\n            'Feature': np.array(feature_names)[top_k_indices],\n            'Importance': avg_importance[top_k_indices]\n        })\n        gating_plot_data = pd.DataFrame({\n            'Feature': np.array(feature_names)[gating_top_k_indices],\n            'Importance': gating_avg_importance[gating_top_k_indices]\n        })\n\n        importances[method] = {\n            'model_top_features': plot_data['Feature'].tolist(),\n            'model_importances': plot_data['Importance'].tolist(),\n            'model_all_importances': avg_importance,\n            'gating_top_features': gating_plot_data['Feature'].tolist(),\n            'gating_importances': gating_plot_data['Importance'].tolist(),\n            'gating_all_importances': gating_avg_importance\n        }\n\n    # Set up the plot style\n    sns.set_style(\"whitegrid\")\n    plt.figure(figsize=(18, 12))\n\n    # Create the main figure\n    gs = plt.GridSpec(2, len(methods), figsize=(18, 12))\n\n    for i, method in enumerate(methods):\n        # Plot the feature importance for the model\n        ax = plt.subplot(gs[0, i])\n        colors = sns.color_palette(\"RdYlBu_r\", n_colors=k)\n        ax.barh(np.arange(k), importances[method]['model_importances'], height=0.8, color=colors)\n        ax.set_title(f\"Model Feature Importance ({method.replace('_', ' ').title()})\", fontsize=14)\n        ax.set_xlabel(\"Average Importance\", fontsize=12)\n        ax.set_yticks(np.arange(k))\n        ax.set_yticklabels(importances[method]['model_top_features'], fontsize=10)\n\n        # Plot the feature importance for the gating network\n        ax = plt.subplot(gs[1, i])\n        colors = sns.color_palette(\"RdYlBu_r\", n_colors=k)\n        ax.barh(np.arange(k), importances[method]['gating_importances'], height=0.8, color=colors)\n        ax.set_title(f\"Gating Network Feature Importance ({method.replace('_', ' ').title()})\", fontsize=14)\n        ax.set_xlabel(\"Average Importance\", fontsize=12)\n        ax.set_yticks(np.arange(k))\n        ax.set_yticklabels(importances[method]['gating_top_features'], fontsize=10)\n\n    plt.suptitle('Feature Importance Comparison', fontsize=16, y=0.95)\n    plt.tight_layout()\n    plt.show()\n\n    return importances","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-19T18:43:17.235075Z","iopub.execute_input":"2025-01-19T18:43:17.235327Z","iopub.status.idle":"2025-01-19T18:43:17.306579Z","shell.execute_reply.started":"2025-01-19T18:43:17.235303Z","shell.execute_reply":"2025-01-19T18:43:17.305751Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from matplotlib.gridspec import GridSpec\nfrom sklearn.manifold import TSNE\nfrom sklearn.cluster import KMeans\nfrom sklearn.metrics import silhouette_score\nfrom matplotlib.gridspec import GridSpec\nfrom typing import Tuple, List, Dict\nimport umap\n\ndef visualize_expert_weights(model, data_loader, expert_names, num_samples=10):\n    \"\"\"\n    Visualize expert weights from a mixture-of-experts model with improved plots.\n    \n    Args:\n        model: The MoE model\n        data_loader: DataLoader containing input samples\n        expert_names: List of expert names/identifiers\n        num_samples: Number of samples to visualize\n        \n    Returns:\n        fig: matplotlib figure containing the visualizations\n        expert_stats: DataFrame containing detailed statistics\n    \"\"\"\n    model.eval()\n    all_weights = []\n    all_labels = []\n    \n    # Collect data\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    # Calculate statistics for better visualization\n    weight_mean = np.mean(all_weights)\n    weight_std = np.std(all_weights)\n    vmin = max(0, weight_mean - 2*weight_std)\n    vmax = min(1, weight_mean + 2*weight_std)\n    \n    # Create figure with GridSpec for better layout\n    fig = plt.figure(figsize=(15, 12))\n    gs = GridSpec(2, 2, figure=fig, height_ratios=[1.2, 1])\n    \n    # 1. Enhanced Heatmap with better contrast\n    ax1 = fig.add_subplot(gs[0, :])\n    sns.heatmap(all_weights,\n                cmap=\"OrRd\",\n                vmin=vmin,\n                vmax=vmax,\n                annot=False,\n                fmt=\".2f\",\n                cbar_kws={'label': 'Expert Weight'},\n                xticklabels=expert_names,\n                ax=ax1)\n    ax1.set_title(\"Expert Weights Distribution Across Samples\", pad=20)\n    ax1.set_xlabel(\"Expert\")\n    ax1.set_ylabel(\"Sample Index\")\n    \n    # Add text annotation for weight range\n    ax1.text(-0.15, -0.15, \n             f'Weight Range: [{vmin:.3f}, {vmax:.3f}]',\n             transform=ax1.transAxes,\n             fontsize=8)\n    \n    # 2. Average Weights with Error Bars and Confidence Intervals\n    ax2 = fig.add_subplot(gs[1, 0])\n    avg_weights = all_weights.mean(axis=0)\n    std_weights = all_weights.std(axis=0)\n    conf_interval = 1.96 * std_weights / np.sqrt(len(all_weights))  # 95% CI\n    \n    # Create enhanced bar plot\n    bars = ax2.bar(range(len(expert_names)), \n                   avg_weights,\n                   yerr=conf_interval,\n                   capsize=5,\n                   alpha=0.6)\n    \n    # Add individual points for distribution visualization\n    for idx in range(len(expert_names)):\n        x = np.random.normal(idx, 0.04, size=len(all_weights))\n        ax2.scatter(x, all_weights[:, idx], alpha=0.1, s=20, color='gray')\n    \n    ax2.set_xticks(range(len(expert_names)))\n    ax2.set_xticklabels(expert_names, rotation=45)\n    ax2.set_title(\"Average Expert Weights with 95% CI\")\n    ax2.set_xlabel(\"Expert\")\n    ax2.set_ylabel(\"Weight\")\n    \n    # Add mean weight annotations\n    for idx, rect in enumerate(bars):\n        height = rect.get_height()\n        ax2.text(rect.get_x() + rect.get_width()/2., height + conf_interval[idx],\n                f'{avg_weights[idx]:.3f}',\n                ha='center', va='bottom', rotation=0)\n    \n    # 3. Expert Co-activation Analysis (replacing time series)\n    ax3 = fig.add_subplot(gs[1, 1])\n    \n    # Calculate co-activation matrix\n    n_experts = len(expert_names)\n    co_activation = np.zeros((n_experts, n_experts))\n    \n    # Consider an expert \"active\" if its weight is above mean weight for that sample\n    threshold = np.mean(all_weights, axis=1, keepdims=True)\n    active_mask = all_weights > threshold\n    \n    for i in range(n_experts):\n        for j in range(n_experts):\n            if i != j:\n                # Calculate how often experts i and j are active together\n                co_activation[i, j] = np.mean(\n                    np.logical_and(active_mask[:, i], active_mask[:, j])\n                )\n            else:\n                # Diagonal shows individual activation rate\n                co_activation[i, i] = np.mean(active_mask[:, i])\n    \n    sns.heatmap(co_activation, \n                annot=True, \n                fmt='.2f', \n                xticklabels=expert_names,\n                yticklabels=expert_names,\n                cmap='RdYlBu_r',\n                ax=ax3)\n    ax3.set_title(\"Expert Co-activation Analysis\")\n    ax3.set_xlabel(\"Expert\")\n    ax3.set_ylabel(\"Expert\")\n    \n    # Calculate comprehensive statistics\n    expert_stats = pd.DataFrame({\n        'Mean': avg_weights,\n        'Std': std_weights,\n        'CI_95_Lower': avg_weights - conf_interval,\n        'CI_95_Upper': avg_weights + conf_interval,\n        'Max': np.max(all_weights, axis=0),\n        'Min': np.min(all_weights, axis=0),\n        'Median': np.median(all_weights, axis=0),\n        'Q1': np.percentile(all_weights, 25, axis=0),\n        'Q3': np.percentile(all_weights, 75, axis=0),\n        'Activation_Rate': np.mean(active_mask, axis=0),\n        'Dominant_Rate': (all_weights.argmax(axis=1)[:, None] == np.arange(n_experts)).mean(axis=0)\n    }, index=expert_names)\n    \n    plt.tight_layout()\n    plt.show()\n    \n    # Print summary statistics\n    print(\"\\nExpert Statistics:\")\n    print(expert_stats)\n    \n    return fig, expert_stats\n\ndef analyze_expert_clusters(\n    model: torch.nn.Module,\n    data_loader: torch.utils.data.DataLoader,\n    expert_names: List[str],\n    num_samples: int = 600,\n    n_clusters: int = 5,\n    random_state: int = 42\n) -> Tuple[plt.Figure, plt.Figure, Dict]:\n    \"\"\"\n    Perform comprehensive analysis of expert clustering patterns.\n\n    Args:\n        model: The MoE model\n        data_loader: DataLoader containing input samples\n        expert_names: List of expert names/identifiers\n        num_samples: Number of samples to analyze\n        n_clusters: Number of clusters for KMeans\n        random_state: Random seed for reproducibility\n\n    Returns:\n        cluster_fig: matplotlib figure for cluster visualization\n        quality_fig: matplotlib figure for cluster quality metrics\n        analysis_results: Dictionary containing analysis metrics\n    \"\"\"\n    model.eval()\n    all_inputs = []\n    all_weights = []\n\n    # Data collection\n    with torch.no_grad():\n        for i, (inputs, _) in enumerate(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    all_inputs = np.concatenate(all_inputs, axis=0)[:num_samples]\n    all_weights = np.concatenate(all_weights, axis=0)[:num_samples]\n\n    # Clustering\n    kmeans = KMeans(n_clusters=n_clusters, n_init=10, random_state=random_state)\n    cluster_labels = kmeans.fit_predict(all_inputs)\n    centroids = kmeans.cluster_centers_\n\n    # Calculate silhouette score\n    silhouette_avg = silhouette_score(all_inputs, cluster_labels)\n\n    # Dimensionality reduction using both t-SNE and UMAP\n    tsne = TSNE(n_components=2, random_state=random_state)\n    umap_reducer = umap.UMAP(random_state=random_state)\n\n    tsne_data = tsne.fit_transform(all_inputs)\n    umap_data = umap_reducer.fit_transform(all_inputs)\n\n    # Dominant experts\n    dominant_experts = np.argmax(all_weights, axis=1)\n\n    # Create figures\n    umap_fig = plt.figure(figsize=(8, 8))\n    ax1 = umap_fig.add_subplot(1, 1, 1)\n    scatter1 = ax1.scatter(umap_data[:, 0], umap_data[:, 1],\n                          c=dominant_experts, cmap='Spectral',\n                          alpha=0.6, s=50)\n    ax1.set_title('UMAP: Samples Colored by Dominant Expert')\n    ax1.set_xlabel('UMAP Component 1')\n    ax1.set_ylabel('UMAP Component 2')\n    plt.colorbar(scatter1, ax=ax1, label='Dominant Expert')\n\n    tsne_fig = plt.figure(figsize=(8, 8))\n    ax2 = tsne_fig.add_subplot(1, 1, 1)\n    scatter2 = ax2.scatter(tsne_data[:, 0], tsne_data[:, 1],\n                          c=dominant_experts, cmap='Spectral',\n                          alpha=0.6, s=50)\n    ax2.set_title('t-SNE: Samples Colored by Dominant Expert')\n    ax2.set_xlabel('t-SNE Component 1')\n    ax2.set_ylabel('t-SNE Component 2')\n    plt.colorbar(scatter2, ax=ax2, label='Dominant Expert')\n\n    expert_cluster_fig = plt.figure(figsize=(12, 8))\n    ax3 = expert_cluster_fig.add_subplot(1, 1, 1)\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    sns.heatmap(expert_cluster_percentages,\n                annot=True,\n                fmt='.2f',\n                cmap='RdYlBu_r',\n                xticklabels=expert_names,\n                ax=ax3)\n    ax3.set_title('Expert Utilization per Cluster')\n    ax3.set_xlabel('Expert')\n    ax3.set_ylabel('Cluster')\n\n    quality_fig = plt.figure(figsize=(12, 8))\n    ax4 = quality_fig.add_subplot(1, 1, 1)\n    inertias = []\n    sil_scores = []\n    cluster_range = range(2, min(11, num_samples))\n\n    for k in cluster_range:\n        kmeans_temp = KMeans(n_clusters=k, n_init=10, random_state=random_state)\n        labels_temp = kmeans_temp.fit_predict(all_inputs)\n        inertias.append(kmeans_temp.inertia_)\n        sil_scores.append(silhouette_score(all_inputs, labels_temp))\n\n    ax4.plot(list(cluster_range), inertias, 'bo-', label='Inertia')\n    ax4_twin = ax4.twinx()\n    ax4_twin.plot(list(cluster_range), sil_scores, 'ro-', label='Silhouette')\n    ax4.set_title('Cluster Quality Metrics')\n    ax4.set_xlabel('Number of Clusters')\n    ax4.set_ylabel('Inertia', color='b')\n    ax4_twin.set_ylabel('Silhouette Score', color='r')\n    lines1, labels1 = ax4.get_legend_handles_labels()\n    lines2, labels2 = ax4_twin.get_legend_handles_labels()\n    ax4.legend(lines1 + lines2, labels1 + labels2, loc='upper right')\n\n    # Compile analysis results\n    analysis_results = {\n        'silhouette_score': silhouette_avg,\n        'cluster_sizes': pd.Series(cluster_labels).value_counts().sort_index().to_dict(),\n        'expert_cluster_distributions': expert_cluster_percentages.tolist(),\n        'inertia_scores': inertias,\n        'silhouette_scores': sil_scores\n    }\n\n    return umap_fig, tsne_fig, expert_cluster_fig, quality_fig, analysis_results\n\ndef plot_expert_clusters(umap_fig, tsne_fig, expert_cluster_fig, quality_fig):\n    # Plot UMAP figure\n    plt.figure(figsize=(8, 8))\n    plt.imshow(umap_fig)\n    plt.title('UMAP: Samples Colored by Dominant Expert')\n    plt.xlabel('UMAP Component 1')\n    plt.ylabel('UMAP Component 2')\n    plt.colorbar()\n    plt.show()\n\n    # Plot t-SNE figure\n    plt.figure(figsize=(8, 8))\n    plt.imshow(tsne_fig)\n    plt.title('t-SNE: Samples Colored by Dominant Expert')\n    plt.xlabel('t-SNE Component 1')\n    plt.ylabel('t-SNE Component 2')\n    plt.colorbar()\n    plt.show()\n\n    # Plot expert-cluster heatmap\n    plt.figure(figsize=(12, 8))\n    plt.imshow(expert_cluster_fig)\n    plt.title('Expert Utilization per Cluster')\n    plt.xlabel('Expert')\n    plt.ylabel('Cluster')\n    plt.colorbar()\n    plt.show()\n\n    # Plot cluster quality metrics\n    plt.figure(figsize=(12, 8))\n    plt.imshow(quality_fig)\n    plt.title('Cluster Quality Metrics')\n    plt.xlabel('Number of Clusters')\n    plt.ylabel('Inertia/Silhouette Score')\n    plt.legend()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-19T18:43:17.307918Z","iopub.execute_input":"2025-01-19T18:43:17.308245Z","iopub.status.idle":"2025-01-19T18:43:41.402671Z","shell.execute_reply.started":"2025-01-19T18:43:17.308208Z","shell.execute_reply":"2025-01-19T18:43:41.401944Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"expert_specs = {\n    \"lstm_expert_1\": [\"TargetEmbedding\", \"MorganFingerPrintEmbedding\"],\n    \"lstm_expert_2\": [\"TargetEmbedding\", \"Mol2VecEmbedding\"],\n    \"lstm_expert_3\": [\"TargetEmbedding\", \"SMILESEmbedding\"],\n    # \"lstm_expert_4\": [\"TargetEmbedding\", \"SMILESEmbedding\", \"Mol2VecEmbedding\"],\n    # \"lstm_expert_5\": [\"TargetEmbedding\", \"SMILESEmbedding\", \"MorganFingerPrintEmbedding\"],\n    # \"lstm_expert_6\": [\"TargetEmbedding\", \"MorganFingerPrintEmbedding\", \"Mol2VecEmbedding\"],  \n}\n\nexpert_config = preprocessor_train.generate_expert_config(expert_specs)\nexpert_names = list(expert_config.keys())\noutput_size = 18211\nepochs = 300\nfolds = 5\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nplt.style.use('ggplot')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-19T18:43:41.403736Z","iopub.execute_input":"2025-01-19T18:43:41.404364Z","iopub.status.idle":"2025-01-19T18:43:41.410492Z","shell.execute_reply.started":"2025-01-19T18:43:41.404330Z","shell.execute_reply":"2025-01-19T18:43:41.409521Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **End-to-End**","metadata":{}},{"cell_type":"code","source":"# Normal MOE\nfinal_model = MixtureOfExperts(expert_config, output_size)\nfinal_model.to(device)\n\ntrain_losses, val_losses, gating_weights_history = train_mixture_of_experts(\n    final_model, train_loader, val_loader, device, epochs=epochs, lr=1e-4, plot_interval=10\n)\n\ntorch.save(final_model.state_dict(), 'final_moe_model.pth')\nprint(\"Final model trained and saved.\")\n\nvisualize_expert_weights(final_model, train_loader, expert_names, 10)\n\nfinal_model.save('e2e_moe_model')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-19T18:43:41.411523Z","iopub.execute_input":"2025-01-19T18:43:41.411777Z","iopub.status.idle":"2025-01-19T18:43:55.960750Z","shell.execute_reply.started":"2025-01-19T18:43:41.411752Z","shell.execute_reply":"2025-01-19T18:43:55.959818Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **EM MOE**","metadata":{}},{"cell_type":"code","source":"em_model = MixtureOfExperts(expert_config, output_size)\nem_model.to(device)\n\ntrain_losses, val_losses, gating_weights_history = train_mixture_of_experts_em(\n    em_model, train_loader, val_loader, device, epochs=epochs, lr=1e-4, plot_interval=10\n)\n\nvisualize_expert_weights(em_model, train_loader, expert_names, 10)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-19T18:43:55.961974Z","iopub.execute_input":"2025-01-19T18:43:55.962260Z","iopub.status.idle":"2025-01-19T18:44:12.420746Z","shell.execute_reply.started":"2025-01-19T18:43:55.962234Z","shell.execute_reply":"2025-01-19T18:44:12.419948Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **Hybrid MOE**","metadata":{}},{"cell_type":"code","source":"# Hybrid MOE\n\nhybrid_model = MixtureOfExperts(expert_config, output_size)\nhybrid_model.to(device)\n\nfinal_loss = train_pretrained_moe(\n    model=hybrid_model,\n    train_loader=train_loader,\n    val_loader=val_loader,\n    device=device,\n    pretrain_epochs=50,\n    gating_epochs=20,\n    fine_tune_epochs=50,\n)\n\n\ntorch.save(hybrid_model.state_dict(), 'hybrid_moe_model.pth')\nprint(\"Hybrid model trained and saved.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-19T18:44:12.422033Z","iopub.execute_input":"2025-01-19T18:44:12.422400Z","iopub.status.idle":"2025-01-19T18:44:37.358415Z","shell.execute_reply.started":"2025-01-19T18:44:12.422359Z","shell.execute_reply":"2025-01-19T18:44:37.357439Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **Testing**","metadata":{}},{"cell_type":"code","source":"import torch\nfrom torch.utils.data import DataLoader, TensorDataset, SubsetRandomSampler\nfrom sklearn.model_selection import KFold\nfrom collections import defaultdict\nimport time\n\ndef train_model(model, train_loader, val_loader, criterion, optimizer, device, num_epochs=100, patience=10):\n    \"\"\"Helper function to train individual experts\"\"\"\n    early_stopping = EarlyStopping(patience=patience)\n    scheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, patience=5)\n    train_losses = []\n    val_losses = []\n\n    for epoch in range(num_epochs):\n        # Training phase\n        model.train()\n        train_loss = 0\n        for batch_X, batch_y in train_loader:\n            batch_X, batch_y = batch_X.to(device), batch_y.to(device)\n\n            optimizer.zero_grad()\n            outputs = model(batch_X)\n            loss = criterion(outputs, batch_y)\n            loss.backward()\n            optimizer.step()\n\n            train_loss += loss.item()\n\n        train_loss /= len(train_loader)\n        train_losses.append(train_loss)\n\n        # Validation phase\n        model.eval()\n        val_loss = 0\n        with torch.no_grad():\n            for batch_X, batch_y in val_loader:\n                batch_X, batch_y = batch_X.to(device), batch_y.to(device)\n                outputs = model(batch_X)\n                loss = criterion(outputs, batch_y)\n                val_loss += loss.item()\n\n        val_loss /= len(val_loader)\n        val_losses.append(val_loss)\n\n        if epoch % 10 == 0:\n            print(f\"Epoch {epoch}: Train Loss = {train_loss:.4f}, Val Loss = {val_loss:.4f}\")\n\n        # Early stopping check\n        # if early_stopping(val_loss):\n        #     break\n\n        scheduler.step(val_loss)\n\n    return train_losses, val_losses\n\ndef evaluate_model(model, dataloader, device):\n    \"\"\"Evaluate model performance on a dataset\"\"\"\n    model.eval()\n    predictions = []\n    actuals = []\n\n    with torch.no_grad():\n        for X_batch, y_batch in dataloader:\n            X_batch, y_batch = X_batch.to(device), y_batch.to(device)\n            pred = model(X_batch)\n            predictions.append(pred.cpu().numpy())\n            actuals.append(y_batch.cpu().numpy())\n\n    predictions = np.concatenate(predictions)\n    actuals = np.concatenate(actuals)\n\n    rmse = np.sqrt(np.mean((predictions - actuals) ** 2))\n    mae = np.mean(np.abs(predictions - actuals))\n    mrrmse = RMSE_rowwise_loss_numpy(predictions, actuals)\n\n    # Calculate R² for each target variable\n    r2_scores = []\n    for i in range(actuals.shape[1]):\n        ss_res = np.sum((actuals[:, i] - predictions[:, i]) ** 2)\n        ss_tot = np.sum((actuals[:, i] - np.mean(actuals[:, i])) ** 2)\n        r2_scores.append(1 - (ss_res / (ss_tot + 1e-10)))\n    r2 = np.mean(r2_scores)\n\n    return {\n        'rmse': rmse,\n        'mae': mae,\n        'mrrmse': mrrmse,\n        'r2': r2,\n        'predictions': predictions,\n        'actuals': actuals\n    }\n\ndef cross_validate_models(X, y, expert_configs, n_splits=5, batch_size=32, learning_rate=0.001):\n    \"\"\"Perform k-fold cross validation comparing MoE against individual experts\"\"\"\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    kf = KFold(n_splits=n_splits, shuffle=True, random_state=42)\n\n    results = {\n        'moe': defaultdict(list),\n        'individual_experts': defaultdict(lambda: defaultdict(list))\n    }\n\n    # Add fold tracking\n    results['fold_predictions'] = []\n\n    for fold, (train_idx, val_idx) in enumerate(tqdm(kf.split(X), total=n_splits, desc=\"Cross-validation folds\")):\n        print(f\"\\nProcessing fold {fold + 1}/{n_splits}\")\n\n        train_sampler = SubsetRandomSampler(train_idx)\n        val_sampler = SubsetRandomSampler(val_idx)\n\n        train_loader = DataLoader(TensorDataset(X, y), batch_size=batch_size, sampler=train_sampler)\n        val_loader = DataLoader(TensorDataset(X, y), batch_size=batch_size, sampler=val_sampler)\n\n        # Train and evaluate MoE\n        print(\"Training MoE model...\")\n        moe_model = MixtureOfExperts(expert_configs, y.shape[1]).to(device)\n        optimizer = torch.optim.Adam(moe_model.parameters(), lr=learning_rate)\n\n        start_time = time.time()\n        train_losses, val_losses, gating_weights_history = train_mixture_of_experts(\n            moe_model, train_loader, val_loader, device, epochs=epochs, lr=1e-4, plot_interval=10\n        )\n        \n        training_time = time.time() - start_time\n\n        # Evaluate MoE\n        fold_results = evaluate_model(moe_model, val_loader, device)\n\n        # Store MoE results\n        for metric in ['rmse', 'mae', 'mrrmse', 'r2']:\n            results['moe'][metric].append(fold_results[metric])\n        results['moe']['training_time'].append(training_time)\n        results['moe']['val_losses'].append(val_losses)\n\n        # Store fold predictions\n        results['fold_predictions'].append({\n            'fold': fold,\n            'moe_pred': fold_results['predictions'],\n            'moe_actual': fold_results['actuals'],\n            'val_indices': val_idx\n        })\n\n        # Train and evaluate individual experts\n        for expert_name, indices in expert_configs.items():\n            print(f\"Training individual expert: {expert_name}\")\n            X_expert = X[:, indices]\n\n            expert = LSTMExpert(len(indices), y.shape[1]).to(device)\n            optimizer = torch.optim.Adam(expert.parameters(), lr=learning_rate)\n\n            train_loader_expert = DataLoader(TensorDataset(X_expert, y),\n                                          batch_size=batch_size, sampler=train_sampler)\n            val_loader_expert = DataLoader(TensorDataset(X_expert, y),\n                                        batch_size=batch_size, sampler=val_sampler)\n\n            start_time = time.time()\n            train_losses, val_losses = train_model(expert, train_loader_expert,\n                                                 val_loader_expert, RMSE_rowwise_loss,\n                                                 optimizer, device)\n            training_time = time.time() - start_time\n\n            # Evaluate expert\n            fold_results = evaluate_model(expert, val_loader_expert, device)\n\n            # Store results\n            for metric in ['rmse', 'mae', 'mrrmse', 'r2']:\n                results['individual_experts'][expert_name][metric].append(fold_results[metric])\n            results['individual_experts'][expert_name]['training_time'].append(training_time)\n            results['individual_experts'][expert_name]['val_losses'].append(val_losses)\n\n            # Store predictions\n            results['fold_predictions'][-1][f'{expert_name}_pred'] = fold_results['predictions']\n\n    summary = compute_summary_statistics(results)\n    return results, summary\n\ndef compute_summary_statistics(results):\n    \"\"\"Compute comprehensive summary statistics\"\"\"\n    summary = {\n        'moe': {},\n        'individual_experts': defaultdict(dict)\n    }\n\n    # Compute MoE statistics\n    for metric in ['rmse', 'mae', 'mrrmse', 'r2', 'training_time']:\n        values = np.array(results['moe'][metric])\n        summary['moe'][metric] = {\n            'mean': np.mean(values),\n            'std': np.std(values),\n            'min': np.min(values),\n            'max': np.max(values),\n            'median': np.median(values)\n        }\n\n    # Compute individual expert statistics\n    for expert_name in results['individual_experts'].keys():\n        for metric in ['rmse', 'mae', 'mrrmse', 'r2', 'training_time']:\n            values = np.array(results['individual_experts'][expert_name][metric])\n            summary['individual_experts'][expert_name][metric] = {\n                'mean': np.mean(values),\n                'std': np.std(values),\n                'min': np.min(values),\n                'max': np.max(values),\n                'median': np.median(values)\n            }\n\n    return summary\n\ndef plot_performance_comparison(results, summary, save_path=None):\n    \"\"\"Create comprehensive performance visualizations\"\"\"\n    metrics = ['rmse', 'mae', 'mrrmse', 'r2']\n\n    # Set style\n    plt.style.use('ggplot')\n\n    # Create figure with subplots\n    fig = plt.figure(figsize=(20, 15))\n    gs = fig.add_gridspec(3, 2, height_ratios=[3, 3, 1])\n\n    # 1. Box plots for each metric\n    for idx, metric in enumerate(metrics):\n        ax = fig.add_subplot(gs[idx // 2, idx % 2])\n\n        # Prepare data\n        plot_data = {\n            'MoE': results['moe'][metric]\n        }\n        plot_data.update({\n            expert: results['individual_experts'][expert][metric]\n            for expert in results['individual_experts'].keys()\n        })\n\n        # Create violin plot with box plot inside\n        sns.violinplot(data=pd.DataFrame(plot_data), ax=ax, inner='box')\n\n        # Customize plot\n        ax.set_title(f'{metric.upper()} Distribution', fontsize=12, pad=20)\n        ax.set_ylabel(metric.upper())\n        ax.tick_params(axis='x', rotation=45)\n\n        # Add mean values and confidence intervals\n        for i, (model, values) in enumerate(plot_data.items()):\n            mean = np.mean(values)\n            ci = np.percentile(values, [2.5, 97.5])\n            ax.text(i, ax.get_ylim()[1], f'μ={mean:.4f}\\nCI=[{ci[0]:.4f}, {ci[1]:.4f}]',\n                   ha='center', va='bottom', fontsize=8)\n\n    # 2. Training time comparison\n    ax_time = fig.add_subplot(gs[2, :])\n    models = ['MoE'] + list(results['individual_experts'].keys())\n    times = [np.mean(results['moe']['training_time'])]\n    times.extend([np.mean(results['individual_experts'][expert]['training_time'])\n                 for expert in results['individual_experts'].keys()])\n\n    # Create bar plot with error bars\n    time_std = [np.std(results['moe']['training_time'])]\n    time_std.extend([np.std(results['individual_experts'][expert]['training_time'])\n                    for expert in results['individual_experts'].keys()])\n\n    bars = ax_time.bar(models, times, yerr=time_std, capsize=5)\n\n    # Customize time plot\n    ax_time.set_title('Training Time Comparison', fontsize=12, pad=20)\n    ax_time.set_ylabel('Time (seconds)')\n    ax_time.tick_params(axis='x', rotation=45)\n\n    # Add value labels on bars\n    for bar in bars:\n        height = bar.get_height()\n        ax_time.text(bar.get_x() + bar.get_width()/2., height,\n                    f'{height:.2f}s',\n                    ha='center', va='bottom')\n\n    plt.tight_layout()\n    if save_path:\n        plt.savefig(f\"{save_path}/performance_comparison.png\", dpi=300, bbox_inches='tight')\n    plt.show()\n\ndef plot_learning_curves(results, save_path=None):\n    \"\"\"Plot learning curves for all models\"\"\"\n    plt.figure(figsize=(15, 8))\n\n    # Plot MoE learning curves\n    mean_losses = np.mean(results['moe']['val_losses'], axis=0)\n    std_losses = np.std(results['moe']['val_losses'], axis=0)\n    epochs = range(1, len(mean_losses) + 1)\n\n    plt.plot(epochs, mean_losses, label='MoE', linewidth=2)\n    plt.fill_between(epochs, mean_losses - std_losses, mean_losses + std_losses, alpha=0.2)\n\n    # Plot individual expert learning curves\n    for expert_name in results['individual_experts'].keys():\n        mean_losses = np.mean(results['individual_experts'][expert_name]['val_losses'], axis=0)\n        std_losses = np.std(results['individual_experts'][expert_name]['val_losses'], axis=0)\n\n        plt.plot(epochs, mean_losses, label=expert_name, linewidth=2)\n        plt.fill_between(epochs, mean_losses - std_losses, mean_losses + std_losses, alpha=0.2)\n\n    plt.title('Validation Learning Curves', fontsize=14)\n    plt.xlabel('Epoch')\n    plt.ylabel('Loss')\n    plt.legend(bbox_to_anchor=(1.05, 1), loc='upper left')\n    plt.grid(True)\n\n    if save_path:\n        plt.savefig(f\"{save_path}/learning_curves.png\", dpi=300, bbox_inches='tight')\n    plt.show()\n\ndef create_performance_summary_table(summary):\n    \"\"\"Create a comprehensive performance summary table\"\"\"\n    # Initialize data for table\n    data = []\n\n    # Add MoE data\n    moe_row = ['Mixture of Experts']\n    for metric in ['rmse', 'mae', 'mrrmse', 'r2']:\n        stats = summary['moe'][metric]\n        moe_row.append(f\"{stats['mean']:.4f} ± {stats['std']:.4f}\\n(min={stats['min']:.4f}, max={stats['max']:.4f})\")\n    moe_row.append(f\"{summary['moe']['training_time']['mean']:.2f} ± {summary['moe']['training_time']['std']:.2f}\")\n    data.append(moe_row)\n\n    # Add individual experts data\n    for expert in summary['individual_experts'].keys():\n        expert_row = [expert]\n        for metric in ['rmse', 'mae', 'mrrmse', 'r2']:\n            stats = summary['individual_experts'][expert][metric]\n            expert_row.append(f\"{stats['mean']:.4f} ± {stats['std']:.4f}\\n(min={stats['min']:.4f}, max={stats['max']:.4f})\")\n        expert_row.append(f\"{summary['individual_experts'][expert]['training_time']['mean']:.2f} ± {summary['individual_experts'][expert]['training_time']['std']:.2f}\")\n        data.append(expert_row)\n\n    # Create DataFrame with all metrics\n    columns = ['Model', 'RMSE', 'MAE', 'MRRMSE', 'R²', 'Training Time (s)']\n    df = pd.DataFrame(data, columns=columns)\n\n    # Style the DataFrame for better visualization\n    styled_df = df.style.set_properties(**{\n        'text-align': 'center',\n        'white-space': 'pre-wrap'\n    }).set_properties(subset=['Model'], **{\n        'text-align': 'left',\n        'font-weight': 'bold'\n    }).set_table_styles([\n        {'selector': 'th', 'props': [('text-align', 'center'), ('font-weight', 'bold')]},\n        {'selector': 'caption', 'props': [('caption-side', 'top')]}\n    ]).set_caption('Model Performance Summary')\n\n    return styled_df\n\ndef visualize_all_results(results, summary, save_path=None):\n    \"\"\"Generate all visualizations and summary tables\"\"\"\n    # 1. Plot performance comparisons\n    plot_performance_comparison(results, summary, save_path)\n\n    # 2. Plot learning curves\n    plot_learning_curves(results, save_path)\n\n    # 3. Create and display summary table\n    summary_table = create_performance_summary_table(summary)\n    display(summary_table)  # For Jupyter notebooks\n\n    return summary_table\n\ndef print_summary_results(summary):\n    \"\"\"Print formatted summary results including all metrics\"\"\"\n    print(\"\\nMoE Performance:\")\n    for metric in ['rmse', 'mae', 'mrrmse', 'r2']:\n        print(f\"{metric.upper()}: {summary['moe'][metric]['mean']:.4f} ± {summary['moe'][metric]['std']:.4f}\")\n\n    print(\"\\nIndividual Expert Performances:\")\n    for expert_name, metrics in summary['individual_experts'].items():\n        print(f\"\\n{expert_name}:\")\n        for metric in ['rmse', 'mae', 'mrrmse', 'r2']:\n            print(f\"{metric.upper()}: {metrics[metric]['mean']:.4f} ± {metrics[metric]['std']:.4f}\")\n\n# Example usage\ndef run_complete_analysis(X, y, expert_configs, n_splits=5, save_path=None):\n    \"\"\"\n    Run complete analysis pipeline\n\n    Parameters:\n    -----------\n    X : numpy.ndarray or torch.Tensor\n        Input features of shape (n_samples, n_features)\n    y : numpy.ndarray or torch.Tensor\n        Target values of shape (n_samples, n_targets)\n    expert_configs : dict\n        Dictionary mapping expert names to their feature indices\n        Example: {'expert1': [0,1,2], 'expert2': [3,4,5]}\n    n_splits : int, optional (default=5)\n        Number of cross-validation folds\n    save_path : str, optional\n        Path to save visualizations\n\n    Returns:\n    --------\n    results : dict\n        Raw results from cross-validation\n    summary : dict\n        Summary statistics\n    summary_table : pandas.DataFrame\n        Formatted summary table\n    \"\"\"\n    # Convert inputs to torch tensors if they aren't already\n    if not isinstance(X, torch.Tensor):\n        X = torch.FloatTensor(X)\n    if not isinstance(y, torch.Tensor):\n        y = torch.FloatTensor(y)\n\n    # Run cross-validation\n    print(\"Starting cross-validation...\")\n    results, summary = cross_validate_models(X, y, expert_configs, n_splits=n_splits)\n\n    # Generate visualizations and summary\n    print(\"\\nGenerating visualizations and summary...\")\n    summary_table = visualize_all_results(results, summary, save_path)\n\n    # Print detailed summary\n    print(\"\\nDetailed Performance Summary:\")\n    print_summary_results(summary)\n\n    return results, summary, summary_table","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-19T18:44:37.359826Z","iopub.execute_input":"2025-01-19T18:44:37.360151Z","iopub.status.idle":"2025-01-19T18:44:37.403841Z","shell.execute_reply.started":"2025-01-19T18:44:37.360121Z","shell.execute_reply":"2025-01-19T18:44:37.403046Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"results, summary, summary_table = run_complete_analysis(\n    X=X,\n    y=y,\n    expert_configs=expert_config,\n    n_splits=5,\n    # save_path='./visualization_outputs'\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-19T18:44:37.404813Z","iopub.execute_input":"2025-01-19T18:44:37.405123Z","iopub.status.idle":"2025-01-19T18:48:15.259044Z","shell.execute_reply.started":"2025-01-19T18:44:37.405098Z","shell.execute_reply":"2025-01-19T18:48:15.257431Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **Interpretations**","metadata":{}},{"cell_type":"code","source":"import torch\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\nclass MOEAnalyzer:\n    def __init__(self, moe_model, save_dir=\"moe_analysis\"):\n        self.model = moe_model\n        self.expert_names = list(moe_model.experts.keys())\n        self.save_dir = save_dir\n        os.makedirs(self.save_dir, exist_ok=True)\n        \n    def analyze_expert_usage(self, dataloader, device='cuda'):\n        self.model.eval()\n        all_gating_weights = []\n        \n        with torch.no_grad():\n            for batch in dataloader:\n                data = batch[0].to(device)\n                _, _, gating_weights = self.model(data, return_info=True)\n                all_gating_weights.append(gating_weights.cpu())\n        \n        all_gating_weights = torch.cat(all_gating_weights, dim=0)\n        \n        expert_usage = {\n            'mean': all_gating_weights.mean(dim=0),\n            'std': all_gating_weights.std(dim=0)\n        }\n        expert_corr = torch.corrcoef(all_gating_weights.T)\n        \n        plt.figure(figsize=(10, 8))\n        sns.heatmap(expert_corr, annot=True, cmap=\"coolwarm\", \n                    xticklabels=self.expert_names, yticklabels=self.expert_names)\n        plt.title(\"Expert Correlation Matrix\")\n        plt.savefig(f\"{self.save_dir}/expert_correlation_matrix.png\")\n        plt.show()\n        plt.close()\n        \n        return expert_usage, expert_corr\n\n    def visualize_expert_weights(self, data_loader, num_samples=100, device='cuda'):\n        self.model.eval()\n        all_weights = []\n        all_labels = []\n        samples_processed = 0\n        \n        with torch.no_grad():\n            for inputs, labels in data_loader:\n                if samples_processed >= num_samples:\n                    break\n                    \n                inputs = inputs.to(device)\n                _ = self.model(inputs)\n                weights = self.model.last_gating_weights\n                \n                weight_sums = weights.sum(dim=1)\n                assert torch.allclose(weight_sums, torch.ones_like(weight_sums), rtol=1e-3), \"Weights don't sum to 1\"\n                \n                remaining = num_samples - samples_processed\n                batch_size = min(remaining, len(inputs))\n                \n                all_weights.append(weights[:batch_size].cpu().numpy())\n                all_labels.extend(labels[:batch_size].cpu().numpy())\n                samples_processed += batch_size\n        \n        all_weights = np.concatenate(all_weights, axis=0)\n        avg_weights = all_weights.mean(axis=0)\n        std_weights = all_weights.std(axis=0)\n        \n        fig, axes = plt.subplots(2, 2, figsize=(20, 16))\n        sns.heatmap(all_weights[:min(50, len(all_weights))], cmap=\"YlOrRd\", \n                    xticklabels=self.expert_names, ax=axes[0,0])\n        axes[0,0].set_title(\"Expert Weights for First 50 Instances\")\n        \n        axes[0,1].bar(self.expert_names, avg_weights, yerr=std_weights, capsize=5)\n        axes[0,1].set_title(\"Average Expert Weights with Standard Deviation\")\n        axes[0,1].set_xticklabels(self.expert_names, rotation=45, ha='right')\n        \n        sns.boxplot(data=all_weights, ax=axes[1,0])\n        axes[1,0].set_title(\"Distribution of Expert Weights\")\n        axes[1,0].set_xticklabels(self.expert_names, rotation=45, ha='right')\n        \n        x = range(len(all_weights))\n        axes[1,1].stackplot(x, all_weights.T)\n        axes[1,1].set_title(\"Expert Weight Distribution Across Samples\")\n        \n        plt.tight_layout()\n        plt.savefig(f\"{self.save_dir}/expert_weights_visualization.png\")\n        plt.close()\n        plt.show()\n        \n        return fig, all_weights, all_labels\n\n    def analyze_expert_specialization(self, data_loader, num_samples=100, device='cuda'):\n        self.model.eval()\n        all_weights = []\n        all_features = []\n        samples_processed = 0\n        \n        with torch.no_grad():\n            for inputs, _ in data_loader:\n                if samples_processed >= num_samples:\n                    break\n                    \n                inputs = inputs.to(device)\n                features = inputs\n                outputs = self.model(inputs)\n                weights = self.model.last_gating_weights\n                \n                remaining = num_samples - samples_processed\n                batch_size = min(remaining, len(inputs))\n                \n                all_weights.append(weights[:batch_size].cpu().numpy())\n                all_features.append(features[:batch_size].cpu().numpy())\n                samples_processed += batch_size\n        \n        all_weights = np.concatenate(all_weights, axis=0)\n        all_features = np.concatenate(all_features, axis=0)\n        \n        weight_entropy = -np.sum(all_weights * np.log(all_weights + 1e-10), axis=1)\n        expert_usage = (all_weights > 0.2).sum(axis=1)\n        \n        expert_similarities = np.zeros((len(self.expert_names), len(self.expert_names)))\n        for i in range(len(self.expert_names)):\n            for j in range(len(self.expert_names)):\n                mask_i = all_weights[:, i] > 0.2\n                mask_j = all_weights[:, j] > 0.2\n                if mask_i.sum() > 0 and mask_j.sum() > 0:\n                    expert_similarities[i, j] = np.corrcoef(\n                        all_features[mask_i].mean(axis=0),\n                        all_features[mask_j].mean(axis=0)\n                    )[0, 1]\n        \n        fig, axes = plt.subplots(2, 2, figsize=(20, 16))\n        sns.heatmap(expert_similarities, annot=True, fmt='.2f', cmap=\"coolwarm\",\n                    xticklabels=self.expert_names, yticklabels=self.expert_names,\n                    ax=axes[0,0])\n        axes[0,0].set_title(\"Expert Similarity Matrix\")\n        axes[0,0].tick_params(axis='both', which='major', labelsize=8)\n        \n        sns.histplot(weight_entropy, bins=30, ax=axes[0,1], kde=True)\n        axes[0,1].set_title(\"Distribution of Gate Entropy\")\n        axes[0,1].set_xlabel(\"Entropy\")\n        \n        sns.histplot(expert_usage, bins=range(len(self.expert_names)+2), ax=axes[1,0])\n        axes[1,0].set_title(\"Number of Active Experts per Sample\")\n        axes[1,0].set_xlabel(\"Number of Active Experts\")\n        \n        linkage_matrix = linkage(expert_similarities, method='ward')\n        dendrogram(linkage_matrix, labels=self.expert_names, ax=axes[1,1])\n        axes[1,1].set_title(\"Expert Hierarchy Dendrogram\")\n        \n        plt.tight_layout()\n        plt.savefig(f\"{self.save_dir}/expert_specialization_analysis.png\")\n        plt.close()\n        plt.show()\n        \n        return fig, {\n            'weight_entropy': weight_entropy,\n            'expert_usage': expert_usage,\n            'expert_similarities': expert_similarities\n        }\n\n    def evaluate_expert_performance(self, data_loader, device='cuda'):\n        self.model.eval()\n        all_predictions = []\n        all_labels = []\n        expert_performance = {name: [] for name in self.expert_names}\n        moe_performance = []\n        \n        with torch.no_grad():\n            for inputs, labels in data_loader:\n                inputs, labels = inputs.to(device), labels.to(device)\n                \n                # Model and expert predictions\n                outputs, expert_outputs, _ = self.model(inputs, return_info=True)\n                all_predictions.append(outputs.cpu())\n                all_labels.append(labels.cpu())\n                \n                moe_performance.extend(mean_absolute_error(outputs.cpu(), labels.cpu(), multioutput='raw_values'))\n                \n                for i, expert_name in enumerate(self.expert_names):\n                    expert_perf = mean_absolute_error(expert_outputs[:, i, :].cpu(), labels.cpu(), multioutput='raw_values')\n                    expert_performance[expert_name].extend(expert_perf)\n        \n        all_predictions = torch.cat(all_predictions, dim=0).numpy()\n        all_labels = torch.cat(all_labels, dim=0).numpy()\n        \n        performance_df = pd.DataFrame(expert_performance)\n        plt.figure(figsize=(10, 8))\n        sns.boxplot(data=performance_df)\n        plt.title(\"Expert Performance Distribution\")\n        plt.ylabel(\"MAE per Output Dimension\")\n        plt.xticks(rotation=45, ha='right')\n        plt.tight_layout()\n        plt.savefig(f\"{self.save_dir}/expert_performance_vs_moe.png\")\n        plt.close()\n        plt.show()\n        \n        return performance_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-19T18:49:09.639427Z","iopub.execute_input":"2025-01-19T18:49:09.639815Z","iopub.status.idle":"2025-01-19T18:49:09.667362Z","shell.execute_reply.started":"2025-01-19T18:49:09.639783Z","shell.execute_reply":"2025-01-19T18:49:09.666413Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"feature_names","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-19T18:48:15.262365Z","iopub.status.idle":"2025-01-19T18:48:15.262702Z","shell.execute_reply.started":"2025-01-19T18:48:15.262547Z","shell.execute_reply":"2025-01-19T18:48:15.262563Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import re\nfrom collections import defaultdict\n\ndef categorize_features(feature_names):\n    \"\"\"\n    Categorize features with specific handling for molecular and cell type embeddings.\n    \n    Categories:\n    - Cell type embeddings (target_ct_emb)\n    - Small molecule embeddings (target_sm_emb)\n    - Morgan fingerprints (CompressedFP_64)\n    - Mol2vec embeddings (vector)\n    - Molecular properties\n    - Other features\n    \n    Parameters:\n    feature_names: List of feature names\n    \n    Returns:\n    dict: Categories of features and their corresponding names/indices\n    \"\"\"\n    categories = {\n        'cell_type_embeddings': [],  # target_ct_emb\n        'molecule_embeddings': [],   # target_sm_emb\n        'morgan_fingerprints': [],   # CompressedFP\n        'mol2vec_embeddings': [],    # vector\n        'molecular_properties': [],\n        'other': []\n    }\n    \n    # Regular expressions for different feature types\n    embedding_pattern = re.compile(r'(.+?)_emb_\\d+')\n    vector_pattern = re.compile(r'vector_\\d+')\n    \n    # Common molecular property keywords\n    molecular_properties = {\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    for idx, feature in enumerate(feature_names):\n        feature_info = {\n            'name': feature,\n            'index': idx\n        }\n        \n        # Extract dimension if present\n        dim_match = re.search(r'_(\\d+)$', feature)\n        if dim_match:\n            feature_info['dimension'] = int(dim_match.group(1))\n            \n        # Categorize based on prefixes\n        if feature.startswith('target_ct_emb_'):\n            categories['cell_type_embeddings'].append(feature_info)\n        elif feature.startswith('target_sm_emb_'):\n            categories['molecule_embeddings'].append(feature_info)\n        elif feature.startswith('CompressedFP'):\n            categories['morgan_fingerprints'].append(feature_info)\n        elif feature.startswith('vector_'):\n            categories['mol2vec_embeddings'].append(feature_info)\n        # Check for molecular properties\n        elif any(prop in feature for prop in molecular_properties):\n            categories['molecular_properties'].append(feature_info)\n        else:\n            categories['other'].append(feature_info)\n    \n    return categories\n\ndef analyze_feature_distribution(categories):\n    \"\"\"\n    Analyze the distribution of features across categories.\n    \n    Parameters:\n    categories: Output from categorize_features\n    \n    Returns:\n    dict: Statistical summary of feature distribution\n    \"\"\"\n    summary = {\n        'total_features': 0,\n        'distribution': {}\n    }\n    \n    # Analyze each category\n    for category_name, features in categories.items():\n        if not features:\n            continue\n            \n        # Get dimensions if available\n        dimensions = [f.get('dimension') for f in features if 'dimension' in f]\n        \n        category_summary = {\n            'count': len(features),\n            'features': [f['name'] for f in features]\n        }\n        \n        if dimensions:\n            category_summary.update({\n                'min_dim': min(dimensions),\n                'max_dim': max(dimensions)\n            })\n            \n        summary['distribution'][category_name] = category_summary\n        summary['total_features'] += len(features)\n    \n    return summary\n\ndef print_feature_summary(feature_names):\n    \"\"\"\n    Print a human-readable summary of feature categories.\n    \n    Parameters:\n    feature_names: List of feature names\n    \"\"\"\n    categories = categorize_features(feature_names)\n    summary = analyze_feature_distribution(categories)\n    \n    print(f\"Feature Distribution Summary:\")\n    print(f\"Total features: {summary['total_features']}\\n\")\n    \n    for category, info in summary['distribution'].items():\n        print(f\"{category.replace('_', ' ').title()}:\")\n        print(f\"  Count: {info['count']}\")\n        \n        if 'min_dim' in info:\n            print(f\"  Dimension range: {info['min_dim']}-{info['max_dim']}\")\n            \n        if category in ['molecular_properties', 'other']:\n            print(\"  Features:\")\n            for feature in info['features']:\n                print(f\"    - {feature}\")\n        print()\n\ndef get_feature_groups():\n    \"\"\"\n    Get the predefined feature groups for the model.\n    \n    Returns:\n    dict: Mapping of feature group names to their prefixes\n    \"\"\"\n    return {\n        'cell_type_embeddings': 'target_ct_emb_',\n        'molecule_embeddings': 'target_sm_emb_',\n        'morgan_fingerprints': 'CompressedFP',\n        'mol2vec_embeddings': 'vector_'\n    }\n\ndef get_features_by_prefix(feature_names, prefix):\n    \"\"\"\n    Get all features that start with a specific prefix.\n    \n    Parameters:\n    feature_names: List of feature names\n    prefix: Prefix to filter by\n    \n    Returns:\n    list: Features matching the prefix\n    \"\"\"\n    return [name for name in feature_names if name.startswith(prefix)]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-19T18:49:14.435723Z","iopub.execute_input":"2025-01-19T18:49:14.436085Z","iopub.status.idle":"2025-01-19T18:49:14.449962Z","shell.execute_reply.started":"2025-01-19T18:49:14.436051Z","shell.execute_reply":"2025-01-19T18:49:14.448905Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"categories = categorize_features(feature_names)\n\n# Print detailed summary\nprint_feature_summary(feature_names)\n\n# Get specific feature groups\nfeature_groups = get_feature_groups()\nfor group_name, prefix in feature_groups.items():\n    matching_features = get_features_by_prefix(feature_names, prefix)\n    print(f\"\\n{group_name}:\")\n    # print(matching_features)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-19T18:49:19.271661Z","iopub.execute_input":"2025-01-19T18:49:19.272402Z","iopub.status.idle":"2025-01-19T18:49:19.283297Z","shell.execute_reply.started":"2025-01-19T18:49:19.272366Z","shell.execute_reply":"2025-01-19T18:49:19.282381Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport numpy as np\nimport seaborn as sns\nimport matplotlib.pyplot as plt\nfrom captum.attr import IntegratedGradients, LayerGradCam, LayerAttribution, ShapleyValueSampling\nfrom matplotlib.gridspec import GridSpec\n\ndef analyze_features(model, X, feature_names, k=10, baseline=None, target=0):\n    \"\"\"\n    Analyze feature importance with both overall and group-wise analysis.\n    \n    Parameters:\n    model: MoE model\n    X: Input tensor\n    feature_names: List of feature names\n    k: Number of top features to show overall\n    baseline: Baseline for attribution\n    target: Target expert/output\n    \n    Returns:\n    Dict with feature importance analysis results\n    \"\"\"\n    if baseline is None:\n        baseline = torch.zeros_like(X)\n    \n    # Define feature groups and their prefixes\n    embedding_prefixes = {\n        'cell_type': 'target_ct_emb_',\n        'small_molecule': 'target_sm_emb_',\n        'morgan_fingerprint': 'CompressedFP',\n        'mol2vec': 'vector_',\n    }\n    \n    # Molecular property keywords\n    molecular_properties = {\n        'Molecular Weight', 'LogP', 'TPSA', 'Number of Atoms', 'Number of Bonds',\n        'Number of Rotatable Bonds', 'Number of Hydrogen Bond Acceptors', \n        'Number of Hydrogen Bond Donors', 'Number of Rings', 'Number of Aromatic Rings',\n        'Number of Stereocenters', 'Fraction of sp3 Carbons', 'Balaban J Index',\n        'Bertz CT', 'QED Score'\n    }\n    \n    # Get attributions\n    ig = ShapleyValueSampling(model.gating_network)\n    attributions = ig.attribute(X, baselines=baseline, target=target)\n    attr_np = attributions.detach().cpu().numpy().mean(axis=0)\n    \n    # Initialize results\n    results = {\n        'overall_features': {\n            'names': feature_names,\n            'values': attr_np,\n            'groups': []  # Store group membership for each feature\n        },\n        'group_analysis': {}\n    }\n    \n    # Analyze each feature and assign to groups\n    for idx, feature in enumerate(feature_names):\n        # Check embedding groups\n        group_assigned = False\n        for group, prefix in embedding_prefixes.items():\n            if feature.startswith(prefix):\n                if group not in results['group_analysis']:\n                    results['group_analysis'][group] = {\n                        'indices': [],\n                        'values': []\n                    }\n                results['group_analysis'][group]['indices'].append(idx)\n                results['group_analysis'][group]['values'].append(attr_np[idx])\n                results['overall_features']['groups'].append(group)\n                group_assigned = True\n                break\n        \n        if not group_assigned:\n            # Check if molecular feature\n            is_molecular = any(prop in feature for prop in molecular_properties)\n            group = 'molecular_features' if is_molecular else 'other_features'\n            \n            if group not in results['group_analysis']:\n                results['group_analysis'][group] = {\n                    'indices': [],\n                    'values': []\n                }\n            results['group_analysis'][group]['indices'].append(idx)\n            results['group_analysis'][group]['values'].append(attr_np[idx])\n            results['overall_features']['groups'].append(group)\n    \n    # Calculate top k features overall\n    abs_values = np.abs(results['overall_features']['values'])\n    top_k_idx = np.argsort(abs_values)[-k:][::-1]\n    \n    results['top_k'] = {\n        'names': [results['overall_features']['names'][i] for i in top_k_idx],\n        'values': results['overall_features']['values'][top_k_idx],\n        'groups': [results['overall_features']['groups'][i] for i in top_k_idx]\n    }\n    \n    # Calculate group-wise importance\n    for group, data in results['group_analysis'].items():\n        values = np.array(data['values'])\n        n_features = len(values)\n        \n        # Raw importance (sum of absolute values)\n        data['raw_importance'] = np.sum(np.abs(values))\n        \n        # Normalized importance (by feature count)\n        data['normalized_importance'] = data['raw_importance'] / n_features\n        data['n_features'] = n_features\n    \n    return results\n\ndef plot_top_k_features(results, figsize=(12, 6)):\n    \"\"\"\n    Visualize top k features across all groups with enhanced styling.\n    \"\"\"\n    plt.figure(figsize=figsize)\n    \n    # Create color palette for groups\n    unique_groups = list(set(results['top_k']['groups']))\n    color_palette = sns.color_palette(\"husl\", len(unique_groups))\n    group_colors = dict(zip(unique_groups, color_palette))\n    \n    # Create bar colors based on groups\n    colors = [group_colors[group] for group in results['top_k']['groups']]\n    \n    # Create plot\n    ax = plt.gca()\n    bars = plt.barh(range(len(results['top_k']['values'])), \n                   np.abs(results['top_k']['values']),\n                   color=colors)\n    \n    # Customize plot\n    plt.yticks(range(len(results['top_k']['names'])), results['top_k']['names'])\n    plt.xlabel('Absolute Feature Importance')\n    plt.title('Top K Features Overall', pad=20)\n    \n    # Add value labels\n    for i, bar in enumerate(bars):\n        width = bar.get_width()\n        plt.text(width, bar.get_y() + bar.get_height()/2,\n                f'{results[\"top_k\"][\"values\"][i]:.3f}',\n                ha='left', va='center', fontsize=10)\n    \n    # Add legend\n    legend_elements = [plt.Rectangle((0,0),1,1, facecolor=color, label=group)\n                      for group, color in group_colors.items()]\n    plt.legend(handles=legend_elements, title='Feature Groups',\n              bbox_to_anchor=(1.05, 1), loc='upper left')\n    \n    plt.grid(True, axis='x', alpha=0.3)\n    plt.tight_layout()\n    plt.show()\n\n\ndef plot_group_importance(results, figsize=(20, 10)):\n    \"\"\"\n    Visualize group-wise importance with both raw and normalized values, including pie charts.\n    \"\"\"\n    groups = list(results['group_analysis'].keys())\n    raw_importance = [data['raw_importance'] for data in results['group_analysis'].values()]\n    normalized_importance = [data['normalized_importance'] for data in results['group_analysis'].values()]\n    feature_counts = [data['n_features'] for data in results['group_analysis'].values()]\n    \n    fig = plt.figure(figsize=figsize)\n    gs = GridSpec(2, 3, width_ratios=[1.2, 1.2, 1], height_ratios=[1, 1])\n    \n    # Raw importance plot and pie chart\n    ax1 = fig.add_subplot(gs[0, 0])\n    sns.barplot(x=raw_importance, y=groups, palette=\"husl\", ax=ax1)\n    ax1.set_title('Raw Group Importance\\n(Sum of Absolute Values)', pad=20)\n    ax1.set_xlabel('Raw Importance')\n    \n    # Add value labels\n    for i, v in enumerate(raw_importance):\n        ax1.text(v, i, f'{v:.3f}', va='center')\n    \n    # Raw importance pie chart\n    ax1_pie = fig.add_subplot(gs[0, 1])\n    total_raw = sum(raw_importance)\n    raw_percentages = [v/total_raw * 100 for v in raw_importance]\n    wedges, texts, autotexts = ax1_pie.pie(raw_percentages, labels=groups, autopct='%1.1f%%',\n                                          colors=sns.color_palette(\"husl\", len(groups)))\n    ax1_pie.set_title('Relative Raw Importance Distribution', pad=20)\n    plt.setp(autotexts, size=8, weight=\"bold\")\n    plt.setp(texts, size=8)\n    \n    # Normalized importance plot and pie chart\n    ax2 = fig.add_subplot(gs[1, 0])\n    sns.barplot(x=normalized_importance, y=groups, palette=\"husl\", ax=ax2)\n    ax2.set_title('Normalized Group Importance\\n(Raw Importance / Feature Count)', pad=20)\n    ax2.set_xlabel('Normalized Importance')\n    \n    # Add value labels\n    for i, v in enumerate(normalized_importance):\n        ax2.text(v, i, f'{v:.3f}', va='center')\n    \n    # Normalized importance pie chart\n    ax2_pie = fig.add_subplot(gs[1, 1])\n    total_norm = sum(normalized_importance)\n    norm_percentages = [v/total_norm * 100 for v in normalized_importance]\n    wedges, texts, autotexts = ax2_pie.pie(norm_percentages, labels=groups, autopct='%1.1f%%',\n                                          colors=sns.color_palette(\"husl\", len(groups)))\n    ax2_pie.set_title('Relative Normalized Importance Distribution', pad=20)\n    plt.setp(autotexts, size=8, weight=\"bold\")\n    plt.setp(texts, size=8)\n    \n    # Feature count bar plot\n    ax3 = fig.add_subplot(gs[:, 2])\n    bars = sns.barplot(y=groups, x=feature_counts, palette=\"husl\", ax=ax3, orient='h')\n    ax3.set_title('Number of Features per Group', pad=20)\n    ax3.set_xlabel('Feature Count')\n    \n    # Add value labels\n    for i, v in enumerate(feature_counts):\n        ax3.text(v, i, str(v), ha='left', va='center')\n    \n    plt.tight_layout()\n    plt.show()\n\ndef create_comprehensive_report(results, figsize=(20, 15)):\n    \"\"\"\n    Create a comprehensive visualization report including all analyses and pie charts.\n    \"\"\"\n    # Create figure with subplots\n    fig = plt.figure(figsize=figsize)\n    gs = GridSpec(3, 2, height_ratios=[1.5, 1, 1])\n    \n    # Top K features plot\n    ax1 = fig.add_subplot(gs[0, 0])\n    \n    # Color palette for groups\n    unique_groups = list(set(results['top_k']['groups']))\n    color_palette = sns.color_palette(\"husl\", len(unique_groups))\n    group_colors = dict(zip(unique_groups, color_palette))\n    colors = [group_colors[group] for group in results['top_k']['groups']]\n    \n    # Plot top k features\n    bars = ax1.barh(range(len(results['top_k']['values'])), \n                    np.abs(results['top_k']['values']),\n                    color=colors)\n    \n    ax1.set_yticks(range(len(results['top_k']['names'])))\n    ax1.set_yticklabels(results['top_k']['names'])\n    ax1.set_xlabel('Absolute Feature Importance')\n    ax1.set_title('Top 20 Features Overall', pad=20)\n    \n    # Add value labels\n    for i, bar in enumerate(bars):\n        width = bar.get_width()\n        ax1.text(width, bar.get_y() + bar.get_height()/2,\n                f'{results[\"top_k\"][\"values\"][i]:.3f}',\n                ha='left', va='center')\n    \n    # Add legend\n    legend_elements = [plt.Rectangle((0,0),1,1, facecolor=color, label=group)\n                      for group, color in group_colors.items()]\n    ax1.legend(handles=legend_elements, title='Feature Groups',\n              bbox_to_anchor=(1.05, 1), loc='upper left')\n    \n    # Add pie chart for top k feature distribution\n    ax1_pie = fig.add_subplot(gs[0, 1])\n    group_counts = {}\n    group_importance = {}\n    for group, value in zip(results['top_k']['groups'], np.abs(results['top_k']['values'])):\n        group_counts[group] = group_counts.get(group, 0) + 1\n        group_importance[group] = group_importance.get(group, 0) + value\n    \n    # Plot pie chart of importance distribution\n    pie_labels = list(group_importance.keys())\n    pie_values = list(group_importance.values())\n    total_importance = sum(pie_values)\n    pie_percentages = [v/total_importance * 100 for v in pie_values]\n    \n    wedges, texts, autotexts = ax1_pie.pie(pie_percentages, \n                                          labels=pie_labels,\n                                          colors=[group_colors[group] for group in pie_labels],\n                                          autopct='%1.1f%%')\n    ax1_pie.set_title(f\"Distribution of Top 20 Feature Importance by Group\", pad=20)\n    plt.setp(autotexts, size=8, weight=\"bold\")\n    plt.setp(texts, size=8)\n    \n    # Group importance plots\n    groups = list(results['group_analysis'].keys())\n    raw_importance = [data['raw_importance'] for data in results['group_analysis'].values()]\n    normalized_importance = [data['normalized_importance'] for data in results['group_analysis'].values()]\n    feature_counts = [data['n_features'] for data in results['group_analysis'].values()]\n    \n    # Raw importance\n    ax2 = fig.add_subplot(gs[1, 0])\n    sns.barplot(x=raw_importance, y=groups, palette=\"husl\", ax=ax2)\n    ax2.set_title('Raw Group Importance', pad=20)\n    ax2.set_xlabel('Raw Importance')\n    \n    for i, v in enumerate(raw_importance):\n        ax2.text(v, i, f'{v:.3f}', va='center')\n    \n    # Raw importance pie chart\n    ax2_pie = fig.add_subplot(gs[1, 1])\n    total_raw = sum(raw_importance)\n    raw_percentages = [v/total_raw * 100 for v in raw_importance]\n    wedges, texts, autotexts = ax2_pie.pie(raw_percentages, labels=groups, autopct='%1.1f%%',\n                                          colors=sns.color_palette(\"husl\", len(groups)))\n    ax2_pie.set_title('Relative Raw Importance Distribution', pad=20)\n    plt.setp(autotexts, size=8, weight=\"bold\")\n    plt.setp(texts, size=8)\n    \n    # Feature counts\n    ax4 = fig.add_subplot(gs[2, :])\n    sns.barplot(x=groups, y=feature_counts, palette=\"husl\", ax=ax4)\n    ax4.set_title('Number of Features per Group', pad=20)\n    ax4.set_xlabel('Feature Groups')\n    ax4.set_ylabel('Feature Count')\n    \n    for i, v in enumerate(feature_counts):\n        ax4.text(i, v, str(v), ha='center', va='bottom')\n    \n    plt.xticks(rotation=45)\n    plt.suptitle('Comprehensive Feature Importance Analysis', y=0.95, fontsize=16)\n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-19T19:06:16.889768Z","iopub.execute_input":"2025-01-19T19:06:16.890140Z","iopub.status.idle":"2025-01-19T19:06:16.927184Z","shell.execute_reply.started":"2025-01-19T19:06:16.890108Z","shell.execute_reply":"2025-01-19T19:06:16.926439Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Run the analysis\nresults = analyze_features(em_model, X, feature_names, k=20)\n\n# View individual plots\nplot_top_k_features(results)\nplot_group_importance(results)\n\n# Or view comprehensive report\ncreate_comprehensive_report(results)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-19T19:10:18.555298Z","iopub.execute_input":"2025-01-19T19:10:18.556044Z","iopub.status.idle":"2025-01-19T19:11:02.453317Z","shell.execute_reply.started":"2025-01-19T19:10:18.556000Z","shell.execute_reply":"2025-01-19T19:11:02.452431Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def analyze_features_stable(model, X, feature_names, k=10, baseline=None, target=0, n_runs=10, random_seed=42):\n    \"\"\"\n    Analyze feature importance with stability through averaging multiple runs.\n    \n    Parameters:\n    model: MoE model\n    X: Input tensor\n    feature_names: List of feature names\n    k: Number of top features to show overall\n    baseline: Baseline for attribution\n    target: Target expert/output\n    n_runs: Number of runs to average over (default: 10)\n    random_seed: Random seed for reproducibility\n    \n    Returns:\n    Dict with feature importance analysis results\n    \"\"\"\n    if baseline is None:\n        baseline = torch.zeros_like(X)\n    \n    # Set random seed for reproducibility\n    torch.manual_seed(random_seed)\n    np.random.seed(random_seed)\n    \n    # Define feature groups and their prefixes\n    embedding_prefixes = {\n        'cell_type': 'target_ct_emb_',\n        'small_molecule': 'target_sm_emb_',\n        'morgan_fingerprint': 'CompressedFP',\n        'mol2vec': 'vector_',\n    }\n    \n    # Molecular property keywords\n    molecular_properties = {\n        'Molecular Weight', 'LogP', 'TPSA', 'Number of Atoms', 'Number of Bonds',\n        'Number of Rotatable Bonds', 'Number of Hydrogen Bond Acceptors', \n        'Number of Hydrogen Bond Donors', 'Number of Rings', 'Number of Aromatic Rings',\n        'Number of Stereocenters', 'Fraction of sp3 Carbons', 'Balaban J Index',\n        'Bertz CT', 'QED Score'\n    }\n    \n    # Initialize array to store multiple runs\n    n_features = len(feature_names)\n    all_runs_attr = np.zeros((n_runs, n_features))\n    \n    # Run multiple times and collect results\n    ig = ShapleyValueSampling(model.gating_network)\n    \n    print(f\"Running {n_runs} iterations for stable feature importance...\")\n    for i in range(n_runs):\n        with torch.no_grad():\n            attributions = ig.attribute(X, baselines=baseline, target=target)\n            all_runs_attr[i] = attributions.detach().cpu().numpy().mean(axis=0)\n        print(f\"Completed run {i+1}/{n_runs}\")\n    \n    # Calculate mean and standard deviation across runs\n    attr_mean = np.mean(all_runs_attr, axis=0)\n    attr_std = np.std(all_runs_attr, axis=0)\n    \n    # Initialize results\n    results = {\n        'overall_features': {\n            'names': feature_names,\n            'values': attr_mean,\n            'std_devs': attr_std,\n            'groups': []  # Store group membership for each feature\n        },\n        'group_analysis': {},\n        'stability_metrics': {\n            'mean_std': np.mean(attr_std),\n            'max_std': np.max(attr_std),\n            'min_std': np.min(attr_std)\n        }\n    }\n    \n    # Analyze each feature and assign to groups\n    for idx, feature in enumerate(feature_names):\n        # Check embedding groups\n        group_assigned = False\n        for group, prefix in embedding_prefixes.items():\n            if feature.startswith(prefix):\n                if group not in results['group_analysis']:\n                    results['group_analysis'][group] = {\n                        'indices': [],\n                        'values': [],\n                        'std_devs': []\n                    }\n                results['group_analysis'][group]['indices'].append(idx)\n                results['group_analysis'][group]['values'].append(attr_mean[idx])\n                results['group_analysis'][group]['std_devs'].append(attr_std[idx])\n                results['overall_features']['groups'].append(group)\n                group_assigned = True\n                break\n        \n        if not group_assigned:\n            # Check if molecular feature\n            is_molecular = any(prop in feature for prop in molecular_properties)\n            group = 'molecular_features' if is_molecular else 'other_features'\n            \n            if group not in results['group_analysis']:\n                results['group_analysis'][group] = {\n                    'indices': [],\n                    'values': [],\n                    'std_devs': []\n                }\n            results['group_analysis'][group]['indices'].append(idx)\n            results['group_analysis'][group]['values'].append(attr_mean[idx])\n            results['group_analysis'][group]['std_devs'].append(attr_std[idx])\n            results['overall_features']['groups'].append(group)\n    \n    # Calculate top k features overall\n    abs_values = np.abs(results['overall_features']['values'])\n    top_k_idx = np.argsort(abs_values)[-k:][::-1]\n    \n    results['top_k'] = {\n        'names': [results['overall_features']['names'][i] for i in top_k_idx],\n        'values': results['overall_features']['values'][top_k_idx],\n        'std_devs': results['overall_features']['std_devs'][top_k_idx],\n        'groups': [results['overall_features']['groups'][i] for i in top_k_idx],\n        'confidence_intervals': [\n            (results['overall_features']['values'][i] - 2*results['overall_features']['std_devs'][i],\n             results['overall_features']['values'][i] + 2*results['overall_features']['std_devs'][i])\n            for i in top_k_idx\n        ]\n    }\n    \n    # Calculate group-wise importance\n    for group, data in results['group_analysis'].items():\n        values = np.array(data['values'])\n        stds = np.array(data['std_devs'])\n        n_features = len(values)\n        \n        # Raw importance (sum of absolute values)\n        data['raw_importance'] = np.sum(np.abs(values))\n        data['raw_importance_std'] = np.sqrt(np.sum(stds**2))  # Propagate uncertainty\n        \n        # Normalized importance (by feature count)\n        data['normalized_importance'] = data['raw_importance'] / n_features\n        data['normalized_importance_std'] = data['raw_importance_std'] / n_features\n        data['n_features'] = n_features\n        \n        # Stability metric for this group\n        data['stability_metric'] = np.mean(stds / np.abs(values))  # Coefficient of variation\n    \n    return results\n\ndef plot_stable_feature_importance(results, figsize=(15, 8)):\n    \"\"\"\n    Plot feature importance with error bars from multiple runs.\n    \n    Parameters:\n    results: Output from analyze_features_stable\n    figsize: Figure size for the plot\n    \"\"\"\n    plt.figure(figsize=figsize)\n    \n    # Create color map for different groups\n    unique_groups = list(set(results['top_k']['groups']))\n    colors = plt.cm.get_cmap('tab10')(np.linspace(0, 1, len(unique_groups)))\n    group_colors = dict(zip(unique_groups, colors))\n    \n    # Plot bars with error bars\n    y_pos = np.arange(len(results['top_k']['names']))\n    bars = plt.barh(y_pos, \n                   results['top_k']['values'],\n                   xerr=results['top_k']['std_devs'],\n                   color=[group_colors[group] for group in results['top_k']['groups']])\n    \n    # Customize plot\n    plt.yticks(y_pos, results['top_k']['names'])\n    plt.xlabel('Feature Importance (mean ± std)')\n    plt.title('Top Feature Importance (Averaged over multiple runs)')\n    \n    # Add legend\n    legend_elements = [plt.Rectangle((0,0),1,1, facecolor=color, label=group)\n                      for group, color in group_colors.items()]\n    plt.legend(handles=legend_elements, title='Feature Groups')\n    \n    # Add stability metrics\n    stability_text = f\"Stability Metrics:\\n\"\n    stability_text += f\"Mean Std: {results['stability_metrics']['mean_std']:.3f}\\n\"\n    stability_text += f\"Max Std: {results['stability_metrics']['max_std']:.3f}\\n\"\n    stability_text += f\"Min Std: {results['stability_metrics']['min_std']:.3f}\"\n    \n    plt.text(1.1, 0.5, stability_text,\n             transform=plt.gca().transAxes,\n             bbox=dict(facecolor='white', alpha=0.8))\n    \n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-19T19:15:16.294509Z","iopub.execute_input":"2025-01-19T19:15:16.294995Z","iopub.status.idle":"2025-01-19T19:15:16.325621Z","shell.execute_reply.started":"2025-01-19T19:15:16.294951Z","shell.execute_reply":"2025-01-19T19:15:16.324532Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Run stable analysis\nstable_results = analyze_features_stable(\n    em_model,\n    X,\n    feature_names,\n    k=10,\n    n_runs=10,\n    random_seed=42\n)\n\n# Plot results with error bars\nplot_stable_feature_importance(stable_results)\n\n# Access stability metrics\nprint(\"\\nStability Metrics:\")\nfor group, data in stable_results['group_analysis'].items():\n    print(f\"\\n{group}:\")\n    print(f\"Raw importance: {data['raw_importance']:.3f} ± {data['raw_importance_std']:.3f}\")\n    print(f\"Normalized importance: {data['normalized_importance']:.3f} ± {data['normalized_importance_std']:.3f}\")\n    print(f\"Stability metric (CV): {data['stability_metric']:.3f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-19T19:21:27.465155Z","iopub.execute_input":"2025-01-19T19:21:27.465772Z","iopub.status.idle":"2025-01-19T19:28:09.102315Z","shell.execute_reply.started":"2025-01-19T19:21:27.465735Z","shell.execute_reply":"2025-01-19T19:28:09.101388Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import matplotlib.pyplot as plt\nimport seaborn as sns\nimport captum\nfrom captum.attr import IntegratedGradients, LayerGradCam, LayerAttribution\n\ndef compute_feature_importance(model, X, feature_names, k=10, baseline=None, target=0, methods=('integrated_gradients', 'saliency', 'deeplift', 'shapley')):\n    \"\"\"\n    Compute feature importance using various interpretability methods for the entire Mixture of Experts model and the gating network.\n\n    Parameters:\n    model (MixtureOfExperts): Trained Mixture of Experts 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\n    Returns:\n    top_k_features (dict): Dictionary mapping method name to list of top k feature names for the entire model.\n    top_k_gating_features (dict): Dictionary mapping method name to list of top k feature names for the gating network.\n    top_k_importances (dict): Dictionary mapping method name to list of top k feature importances for the entire model.\n    top_k_gating_importances (dict): Dictionary mapping method name to list of top k feature importances for the gating network.\n    all_importances (dict): Dictionary mapping method name to list of all feature importances for the entire model.\n    all_gating_importances (dict): Dictionary mapping method name to list of all feature importances for the gating network.\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_gating_features = {}\n    top_k_importances = {}\n    top_k_gating_importances = {}\n    all_importances = {}\n    all_gating_importances = {}\n\n    # explainer = IntegratedGradients(model)\n    # attributions = explainer.attribute(X, baselines=baseline, target=target)\n    # gating_explainer = IntegratedGradients(model.gating_network)\n    # gating_attributions = gating_explainer.attribute(X, baselines=baseline, target=target)\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            gating_explainer = IntegratedGradients(model.gating_network)\n            gating_attributions = gating_explainer.attribute(X, baselines=baseline, target=target)\n        elif method == 'saliency':\n            explainer = Saliency(model)\n            attributions = explainer.attribute(X, target=target)\n            gating_explainer = Saliency(model.gating_network)\n            gating_attributions = gating_explainer.attribute(X, target=target)\n        elif method == 'deeplift':\n            explainer = DeepLift(model)\n            attributions = explainer.attribute(X, baselines=baseline, target=target)\n            gating_explainer = DeepLift(model.gating_network)\n            gating_attributions = gating_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            gating_explainer = ShapleyValueSampling(model.gating_network)\n            gating_attributions = gating_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    gating_attributions_np = gating_attributions.cpu().detach().numpy()\n    \n    # Average feature importance across all samples\n    avg_importance = attributions_np.mean(axis=0)\n    avg_gating_importance = gating_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    top_k_gating_indices = np.argsort(np.abs(avg_gating_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    top_k_gating_features[method] = [feature_names[i] for i in top_k_gating_indices]\n    top_k_gating_importances[method] = avg_gating_importance[top_k_gating_indices]\n    all_gating_importances[method] = avg_gating_importance\n    \n    # Plot the attributions for top k features using Seaborn\n    fig, ax = plt.subplots(1, 2, figsize=(16, 6))\n    \n    # Plot top k features for the entire model\n    sns.barplot(x=top_k_importances[method], y=np.array(feature_names)[top_k_indices], orient='h', ax=ax[0])\n    ax[0].set_xlabel(\"Average Importance across samples\")\n    ax[0].set_ylabel(\"Top Features\")\n    ax[0].set_title(f\"Top {k} Feature Importance ({method}) - Model\")\n    \n    # Plot top k features for the gating network\n    sns.barplot(x=top_k_gating_importances[method], y=np.array(feature_names)[top_k_gating_indices], orient='h', ax=ax[1])\n    ax[1].set_xlabel(\"Average Importance across samples\")\n    ax[1].set_ylabel(\"Top Features\")\n    ax[1].set_title(f\"Top {k} Feature Importance ({method}) - Gating Network\")\n    \n    plt.tight_layout()\n    plt.show()\n    \n    return (top_k_features, top_k_gating_features, top_k_importances, top_k_gating_importances, all_importances, all_gating_importances)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-19T13:21:47.958746Z","iopub.status.idle":"2025-01-19T13:21:47.959279Z","shell.execute_reply.started":"2025-01-19T13:21:47.958999Z","shell.execute_reply":"2025-01-19T13:21:47.959024Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"compute_feature_importance(final_model, X, feature_names)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-19T13:21:47.960291Z","iopub.status.idle":"2025-01-19T13:21:47.960751Z","shell.execute_reply.started":"2025-01-19T13:21:47.960513Z","shell.execute_reply":"2025-01-19T13:21:47.960536Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **Predictions**","metadata":{}},{"cell_type":"code","source":"id_map = pd.read_csv(\"/kaggle/input/open-problems-single-cell-perturbations/id_map.csv\")\nsm_name_to_smiles = df.set_index('sm_name')['SMILES'].to_dict()\nlen(sm_name_to_smiles)\nid_map['SMILES'] = id_map['sm_name'].map(sm_name_to_smiles)\nprint(id_map.isna().sum())\n# Preprocesing test dataset\ndf_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":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-02T15:26:13.440823Z","iopub.execute_input":"2025-01-02T15:26:13.441587Z","iopub.status.idle":"2025-01-02T15:26:14.950569Z","shell.execute_reply.started":"2025-01-02T15:26:13.441551Z","shell.execute_reply":"2025-01-02T15:26:14.949687Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_array =  processed_df_test.to_numpy()\nhybrid_model.to(device)\ntarget_pred = hybrid_model(torch.Tensor(test_array).to(device))\nres = target_pred.cpu()\nsample_submission = pd.read_csv(\"/kaggle/input/open-problems-single-cell-perturbations/sample_submission.csv\")\nsample_columns = sample_submission.columns\nsample_columns= sample_columns[1:]\nsubmission_df = pd.DataFrame(res.detach().numpy(), columns=sample_columns)\nsubmission_df.insert(0, 'id', range(255))\ndisplay(submission_df)\nsubmission_df.to_csv(\"submission_hybrid.csv\", index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-02T15:26:14.951991Z","iopub.execute_input":"2025-01-02T15:26:14.952271Z","iopub.status.idle":"2025-01-02T15:26:22.219607Z","shell.execute_reply.started":"2025-01-02T15:26:14.952244Z","shell.execute_reply":"2025-01-02T15:26:22.218873Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_array =  processed_df_test.to_numpy()\nfinal_model.to(device)\ntarget_pred = final_model(torch.Tensor(test_array).to(device))\nres = target_pred.cpu()\nsample_submission = pd.read_csv(\"/kaggle/input/open-problems-single-cell-perturbations/sample_submission.csv\")\nsample_columns = sample_submission.columns\nsample_columns= sample_columns[1:]\nsubmission_df = pd.DataFrame(res.detach().numpy(), columns=sample_columns)\nsubmission_df.insert(0, 'id', range(255))\ndisplay(submission_df)\nsubmission_df.to_csv(\"submission_final.csv\", index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-02T15:26:22.220773Z","iopub.execute_input":"2025-01-02T15:26:22.221114Z","iopub.status.idle":"2025-01-02T15:26:29.227004Z","shell.execute_reply.started":"2025-01-02T15:26:22.221080Z","shell.execute_reply":"2025-01-02T15:26:29.226063Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_array =  processed_df_test.to_numpy()\nem_model.to(device)\ntarget_pred = em_model(torch.Tensor(test_array).to(device))\nres = target_pred.cpu()\nsample_submission = pd.read_csv(\"/kaggle/input/open-problems-single-cell-perturbations/sample_submission.csv\")\nsample_columns = sample_submission.columns\nsample_columns= sample_columns[1:]\nsubmission_df = pd.DataFrame(res.detach().numpy(), columns=sample_columns)\nsubmission_df.insert(0, 'id', range(255))\ndisplay(submission_df)\nsubmission_df.to_csv(\"submission_em.csv\", index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-02T15:26:29.228912Z","iopub.execute_input":"2025-01-02T15:26:29.229185Z","iopub.status.idle":"2025-01-02T15:26:36.228926Z","shell.execute_reply.started":"2025-01-02T15:26:29.229158Z","shell.execute_reply":"2025-01-02T15:26:36.228210Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# !kaggle competitions submit -c open-problems-single-cell-perturbations -f \"/kaggle/working/submission_hybrid.csv\" -m \"kaggle_test_OP_MOE_V2\"\n# !kaggle competitions submit -c open-problems-single-cell-perturbations -f \"/kaggle/working/submission_final.csv\" -m \"kaggle_test_OP_MOE_V2\"\n!kaggle competitions submit -c open-problems-single-cell-perturbations -f \"/kaggle/working/submission_em.csv\" -m \"kaggle_test_OP_MOE_V2\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-02T15:26:36.230006Z","iopub.execute_input":"2025-01-02T15:26:36.230375Z","iopub.status.idle":"2025-01-02T15:26:42.305197Z","shell.execute_reply.started":"2025-01-02T15:26:36.230318Z","shell.execute_reply":"2025-01-02T15:26:42.304249Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **Visualisation**","metadata":{}},{"cell_type":"code","source":"!pip install -q gradio plotly rdkit torch pillow","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-02T15:26:42.306690Z","iopub.execute_input":"2025-01-02T15:26:42.307062Z","iopub.status.idle":"2025-01-02T15:26:57.215429Z","shell.execute_reply.started":"2025-01-02T15:26:42.307012Z","shell.execute_reply":"2025-01-02T15:26:57.214239Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def prepare_single_prediction(preprocessor, cell_type, sm_name):\n        \"\"\"\n        Prepare input data for a single prediction using specified cell type and small molecule.\n        \n        Args:\n            cell_type (str): Cell type name\n            sm_name (str): Small molecule name\n            \n        Returns:\n            pd.DataFrame: Processed input data ready for model prediction\n        \"\"\"\n        # Validate inputs\n        if cell_type not in preprocessor.unique_cell_types:\n            raise ValueError(f\"Unknown cell type: {cell_type}. Available cell types: {sorted(self.unique_cell_types)}\")\n        if sm_name not in preprocessor.unique_sm_names:\n            raise ValueError(f\"Unknown small molecule: {sm_name}. Available small molecules: {sorted(self.unique_sm_names)}\")\n        \n        # Create a single-row DataFrame with the required columns\n        input_df = pd.DataFrame({\n            'cell_type': [cell_type],\n            'sm_name': [sm_name]\n        })\n        \n        # Get the processed features\n        processed_features = preprocessor.preprocess(input_df, fit=False)\n        \n        return processed_features\n\ndef get_available_options(preprocessor):\n        \"\"\"\n        Get available cell types and small molecules for the interface.\n        \n        Returns:\n            dict: Dictionary containing lists of available cell types and small molecules\n        \"\"\"\n        return {\n            'cell_types': sorted(list(preprocessor.unique_cell_types)),\n            'small_molecules': sorted(list(preprocessor.unique_sm_names))\n        }\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-02T15:26:57.217004Z","iopub.execute_input":"2025-01-02T15:26:57.217366Z","iopub.status.idle":"2025-01-02T15:26:57.224431Z","shell.execute_reply.started":"2025-01-02T15:26:57.217323Z","shell.execute_reply":"2025-01-02T15:26:57.223575Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import gradio as gr\nimport plotly.graph_objects as go\nimport plotly.express as px\nfrom rdkit import Chem\nfrom rdkit.Chem import Draw\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom PIL import Image\nfrom sklearn.cluster import KMeans\nfrom sklearn.manifold import TSNE\n\ndef create_molecule_image(smiles):\n    \"\"\"Create molecule image and return PIL Image directly\"\"\"\n    mol = Chem.MolFromSmiles(smiles)\n    img = Draw.MolToImage(mol, size=(400, 400))\n    return img\n\ndef create_expert_importance_plot(gating_weights):\n    \"\"\"Create plotly figure for expert importance visualization\"\"\"\n    if torch.is_tensor(gating_weights):\n        gating_weights = gating_weights.cpu().numpy()\n    \n    if len(gating_weights.shape) > 1:\n        gating_weights = gating_weights.squeeze()\n    \n    fig = go.Figure()\n    experts = [f\"Expert {i+1}\" for i in range(len(gating_weights))]\n    \n    fig.add_trace(go.Bar(\n        x=experts,\n        y=gating_weights,\n        marker_color='rgb(55, 83, 109)'\n    ))\n    \n    fig.update_layout(\n        title='Expert Importance Weights',\n        xaxis_title='Experts',\n        yaxis_title='Weight',\n        template='plotly_white',\n        height=400,\n        yaxis=dict(range=[0, 1])\n    )\n    return fig\n\ndef create_cluster_visualization(input_features, cached_data, n_clusters=5):\n    \"\"\"Create cluster visualization plot\"\"\"\n    # Get cached cluster data\n    reduced_inputs = cached_data['reduced_inputs']\n    cluster_centers = cached_data['cluster_centers']\n    kmeans_model = cached_data['kmeans_model']\n    all_inputs = cached_data['all_inputs']  # Get all_inputs from cached data\n    \n    # Ensure input_features has the same dtype\n    input_features = input_features.astype(np.float64)\n    \n    # Use the cached model to predict\n    try:\n        new_cluster = kmeans_model.predict(input_features)[0]\n    except AttributeError:  # Handle older/newer scikit-learn versions\n        new_cluster = kmeans_model.fit_predict(input_features)[0]\n    \n    # Project the new point into the same t-SNE space\n    # We'll approximate its position by finding the nearest neighbors\n    distances = np.linalg.norm(kmeans_model.cluster_centers_ - input_features, axis=1)\n    nearest_cluster = np.argmin(distances)\n    \n    # Find points in the nearest cluster\n    cluster_points = reduced_inputs[cached_data['cluster_labels'] == nearest_cluster]\n    \n    # Use the mean position of the 5 nearest points in that cluster as the approximate position\n    original_space_distances = np.linalg.norm(\n        all_inputs[cached_data['cluster_labels'] == nearest_cluster] - input_features, \n        axis=1\n    )\n    k = min(5, len(original_space_distances))\n    nearest_indices = np.argpartition(original_space_distances, k)[:k]\n    current_point_position = cluster_points[nearest_indices].mean(axis=0)\n    \n    # Create visualization using plotly\n    fig = go.Figure()\n    \n    # Plot existing points with consistent color scheme\n    colors = px.colors.qualitative.Set3[:n_clusters]\n    for cluster in range(n_clusters):\n        mask = cached_data['cluster_labels'] == cluster\n        fig.add_trace(go.Scatter(\n            x=reduced_inputs[mask, 0],\n            y=reduced_inputs[mask, 1],\n            mode='markers',\n            name=f'Cluster {cluster}',\n            marker=dict(\n                size=8, \n                opacity=0.6,\n                color=colors[cluster]\n            ),\n            showlegend=True\n        ))\n    \n    # Plot cluster centers\n    fig.add_trace(go.Scatter(\n        x=cluster_centers[:, 0],\n        y=cluster_centers[:, 1],\n        mode='markers',\n        name='Centroids',\n        marker=dict(\n            symbol='x',\n            size=12,\n            line=dict(width=2),\n            color='black'\n        )\n    ))\n    \n    # Plot current point\n    fig.add_trace(go.Scatter(\n        x=[current_point_position[0]],\n        y=[current_point_position[1]],\n        mode='markers',\n        name='Current Input',\n        marker=dict(\n            symbol='star',\n            size=15,\n            color=colors[new_cluster],\n            line=dict(color='black', width=2)\n        )\n    ))\n    \n    # Add current input cluster highlight\n    fig.add_annotation(\n        text=f\"Current Input: Cluster {new_cluster}\",\n        xref=\"paper\", yref=\"paper\",\n        x=0.02, y=0.98,\n        showarrow=False,\n        font=dict(size=14, color=\"black\"),\n        bgcolor=\"white\",\n        bordercolor=colors[new_cluster],\n        borderwidth=2\n    )\n    \n    fig.update_layout(\n        title='Cluster Analysis of Input Data',\n        xaxis_title='t-SNE Component 1',\n        yaxis_title='t-SNE Component 2',\n        template='plotly_white',\n        height=400,\n        legend=dict(\n            yanchor=\"top\",\n            y=0.99,\n            xanchor=\"right\",\n            x=0.99\n        ),\n        margin=dict(l=20, r=20, t=40, b=20)\n    )\n    \n    return fig\n\ndef create_gene_expression_plot(predictions):\n    \"\"\"Create plotly figure for gene expression visualization\"\"\"\n    if torch.is_tensor(predictions):\n        predictions = predictions.cpu().numpy()\n    \n    if len(predictions.shape) > 1:\n        predictions = predictions.squeeze()\n    \n    display_genes = predictions[:100]\n    \n    fig = go.Figure()\n    fig.add_trace(go.Scatter(\n        y=display_genes,\n        mode='lines+markers',\n        name='Expression',\n        line=dict(color='rgb(67, 147, 195)')\n    ))\n    \n    fig.update_layout(\n        title='Gene Expression Predictions (First 100 Genes)',\n        xaxis_title='Gene Index',\n        yaxis_title='Log Fold Change',\n        template='plotly_white',\n        height=400,\n        showlegend=False\n    )\n    return fig\n\ndef predict_and_visualize(cell_type, sm_name, preprocessor, model, device, cached_cluster_data):\n    \"\"\"Main prediction and visualization function\"\"\"\n    try:\n        # Prepare input data\n        input_df = prepare_single_prediction(preprocessor, cell_type, sm_name)\n        test_array = input_df.to_numpy().astype(np.float64)  # Ensure consistent dtype\n        \n        # Get model prediction\n        model.to(device)\n        with torch.no_grad():\n            target_pred = model(torch.Tensor(test_array).to(device))\n            predictions = target_pred.cpu().numpy().squeeze()\n            gating_weights = model.last_gating_weights.detach()\n        \n        # Get molecule visualization\n        smiles = preprocessor.get_smiles_mapping()[sm_name]\n        mol_img = create_molecule_image(smiles)\n        \n        # Create plots\n        expert_plot = create_expert_importance_plot(gating_weights)\n        expression_plot = create_gene_expression_plot(predictions)\n        cluster_plot = create_cluster_visualization(test_array, cached_cluster_data)\n        \n        return mol_img, expert_plot, expression_plot, cluster_plot, \"Prediction completed successfully!\"\n        \n    except Exception as e:\n        import traceback\n        return None, None, None, None, f\"Error: {str(e)}\\n{traceback.format_exc()}\"\n\ndef prepare_cluster_data(model, data_loader, n_clusters=5, n_samples=600):\n    \"\"\"Prepare cached cluster data\"\"\"\n    all_inputs = []\n    with torch.no_grad():\n        for inputs, _ in data_loader:\n            if len(all_inputs) * inputs.shape[0] >= n_samples:\n                break\n            all_inputs.append(inputs.cpu().numpy())\n    \n    # Ensure consistent dtype (float64/double)\n    all_inputs = np.concatenate(all_inputs, axis=0)[:n_samples].astype(np.float64)\n    \n    # Perform clustering\n    kmeans = KMeans(n_clusters=n_clusters, random_state=42)\n    cluster_labels = kmeans.fit_predict(all_inputs)\n    \n    # Perform t-SNE on all inputs\n    tsne = TSNE(n_components=2, random_state=42)\n    reduced_data = tsne.fit_transform(all_inputs)\n    \n    # Project cluster centers into the same space\n    cluster_centers = np.array([\n        reduced_data[cluster_labels == i].mean(axis=0)\n        for i in range(n_clusters)\n    ])\n    \n    # Store cached data\n    cached_data = {\n        'reduced_inputs': reduced_data,\n        'cluster_labels': cluster_labels,\n        'kmeans_model': kmeans,\n        'cluster_centers': cluster_centers,\n        'all_inputs': all_inputs  # Add all_inputs to cached data\n    }\n    \n    return cached_data\n\ndef create_gradio_interface(preprocessor, model, device, data_loader):\n    \"\"\"Create Gradio interface\"\"\"\n    # Prepare cluster data\n    cached_cluster_data = prepare_cluster_data(model, data_loader)\n    \n    # Get available options\n    options = get_available_options(preprocessor)\n    cell_types = options['cell_types']\n    small_molecules = options['small_molecules']\n    \n    # Create the interface\n    with gr.Blocks(title=\"Gene Expression Prediction\") as interface:\n        gr.Markdown(\"# scMOE: Gene Expression Prediction Interface\")\n        \n        with gr.Row():\n            with gr.Column(scale=1):\n                cell_type_input = gr.Dropdown(\n                    choices=cell_types,\n                    label=\"Select Cell Type\",\n                    info=\"Choose one of the available cell types\"\n                )\n                sm_name_input = gr.Dropdown(\n                    choices=small_molecules,\n                    label=\"Select Small Molecule\",\n                    info=\"Choose a small molecule compound\",\n                )\n                predict_btn = gr.Button(\"Predict Gene Expression\", variant=\"primary\")\n            \n            with gr.Column(scale=1):\n                mol_image = gr.Image(\n                    label=\"Molecule Structure\",\n                    type=\"pil\"\n                )\n        \n        with gr.Row():\n            with gr.Column(scale=1):\n                expert_plot = gr.Plot(label=\"Expert Importance\")\n            with gr.Column(scale=1):\n                expression_plot = gr.Plot(label=\"Gene Expression Predictions\")\n        \n        with gr.Row():\n            cluster_plot = gr.Plot(label=\"Cluster Analysis\")\n            \n        output_message = gr.Textbox(label=\"Status\")\n        \n        # Set up the prediction function\n        predict_btn.click(\n            fn=lambda ct, sm: predict_and_visualize(ct, sm, preprocessor, model, device, cached_cluster_data),\n            inputs=[cell_type_input, sm_name_input],\n            outputs=[mol_image, expert_plot, expression_plot, cluster_plot, output_message]\n        )\n        \n        # Add example inputs\n        gr.Examples(\n            examples=[\n                [cell_types[0], small_molecules[0]],\n                # [cell_types[1], small_molecules[1]],\n                [cell_types[1], small_molecules[21]]\n            ],\n            inputs=[cell_type_input, sm_name_input],\n            label=\"Example Inputs\"\n        )\n    \n    return interface","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-02T15:26:57.225815Z","iopub.execute_input":"2025-01-02T15:26:57.226077Z","iopub.status.idle":"2025-01-02T15:27:00.055746Z","shell.execute_reply.started":"2025-01-02T15:26:57.226052Z","shell.execute_reply":"2025-01-02T15:27:00.055001Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    \ninterface = create_gradio_interface(preprocessor_train, final_model, device, train_loader)\ninterface.launch(share=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-02T15:27:14.485771Z","iopub.execute_input":"2025-01-02T15:27:14.486480Z","iopub.status.idle":"2025-01-02T15:27:18.927748Z","shell.execute_reply.started":"2025-01-02T15:27:14.486442Z","shell.execute_reply":"2025-01-02T15:27:18.926943Z"}},"outputs":[],"execution_count":null}]}