{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","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},{"sourceId":255833,"sourceType":"modelInstanceVersion","modelInstanceId":218714,"modelId":240454}],"dockerImageVersionId":30840,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# **scMOE: Embedding Comparsion**","metadata":{}},{"cell_type":"code","source":"# installations\n!pip install -q umap-learn rdkit captum git+https://github.com/samoturk/mol2vec;\n!pip install -q selfies==2.1.1  simpletransformers==0.63.9 pandarallel==1.6.4 wandb==0.13.10","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-20T17:58:03.56923Z","iopub.execute_input":"2025-02-20T17:58:03.5697Z","iopub.status.idle":"2025-02-20T17:58:44.90852Z","shell.execute_reply.started":"2025-02-20T17:58:03.569662Z","shell.execute_reply":"2025-02-20T17:58:44.907437Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Imports\n\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!-----')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-02-20T17:58:44.909866Z","iopub.execute_input":"2025-02-20T17:58:44.910181Z","iopub.status.idle":"2025-02-20T17:58:48.746183Z","shell.execute_reply.started":"2025-02-20T17:58:44.910144Z","shell.execute_reply":"2025-02-20T17:58:48.745235Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# uncomment to download data\n# !kaggle competitions download -c open-problems-single-cell-perturbations\n# !unzip -q  open-problems-single-cell-perturbations.zip -d opxmoe_data","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-20T17:58:48.747771Z","iopub.execute_input":"2025-02-20T17:58:48.748104Z","iopub.status.idle":"2025-02-20T17:58:48.75146Z","shell.execute_reply.started":"2025-02-20T17:58:48.748083Z","shell.execute_reply":"2025-02-20T17:58:48.750566Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from gensim.models import word2vec\nmodel1 = word2vec.Word2Vec.load(f\"/kaggle/input/mol2vec/pytorch/default/1/model_300dim.pkl\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-20T17:58:48.752759Z","iopub.execute_input":"2025-02-20T17:58:48.75297Z","iopub.status.idle":"2025-02-20T17:59:05.669914Z","shell.execute_reply.started":"2025-02-20T17:58:48.752941Z","shell.execute_reply":"2025-02-20T17:59:05.669246Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = pd.read_parquet('../input/open-problems-single-cell-perturbations/de_train.parquet')\ndf.tail()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-20T17:59:05.670755Z","iopub.execute_input":"2025-02-20T17:59:05.671011Z","iopub.status.idle":"2025-02-20T17:59:07.456763Z","shell.execute_reply.started":"2025-02-20T17:59:05.670977Z","shell.execute_reply":"2025-02-20T17:59:07.455813Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df.columns","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-20T17:59:07.457828Z","iopub.execute_input":"2025-02-20T17:59:07.458171Z","iopub.status.idle":"2025-02-20T17:59:07.464917Z","shell.execute_reply.started":"2025-02-20T17:59:07.458137Z","shell.execute_reply":"2025-02-20T17:59:07.463841Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Function to convert SMILES to SELFIES\nfrom selfies import encoder\ndef smiles_to_selfies(smiles):\n    try:\n        return encoder(smiles)\n    except Exception as e:\n        print(f\"Error converting SMILES '{smiles}': {e}\") \n        return None  ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-20T17:59:07.466044Z","iopub.execute_input":"2025-02-20T17:59:07.46646Z","iopub.status.idle":"2025-02-20T17:59:07.489855Z","shell.execute_reply.started":"2025-02-20T17:59:07.466423Z","shell.execute_reply":"2025-02-20T17:59:07.488865Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **Embedding Classes**","metadata":{}},{"cell_type":"code","source":"from abc import ABC, abstractmethod\nimport json\nimport pickle\nfrom typing import Dict, Optional, Union, List, Tuple, Literal\nimport umap\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom sklearn.preprocessing import StandardScaler, LabelEncoder\nfrom sklearn.metrics import silhouette_score, silhouette_samples\n\nNON_GENE_COLUMNS = ['cell_type', 'sm_name', 'sm_lincs_id', 'SMILES', 'control', 'SELFIES']\n\nclass BaseEmbedding(ABC):\n    \"\"\"\n    Abstract base class for embeddings with enhanced functionality for gene expression prediction.\n    \n    This class provides a framework for different types of embeddings (e.g., cell type, small molecule)\n    with consistent interfaces for preprocessing, computing, storing, and retrieving embeddings.\n    \"\"\"\n    \n    def __init__(self, name: str = \"base\"):\n        \"\"\"\n        Initialize the embedding class.\n        \n        Args:\n            name (str): Identifier for the embedding type\n        \"\"\"\n        self.embedding_dict: Dict[str, np.ndarray] = {}\n        self.metadata: Dict = {}\n        self.name = name\n        self._is_fitted = False\n    \n    @abstractmethod\n    def preprocess(self, df: pd.DataFrame) -> pd.DataFrame:\n        \"\"\"\n        Preprocess the input data before computing embeddings.\n        \n        Args:\n            df (pd.DataFrame): Input DataFrame containing entity information\n            \n        Returns:\n            pd.DataFrame: Preprocessed DataFrame\n        \"\"\"\n        pass\n    \n    @abstractmethod\n    def get_embedding(self, df: pd.DataFrame, fit: bool = True) -> np.ndarray:\n        \"\"\"\n        Compute embeddings for the input data.\n        \n        Args:\n            df (pd.DataFrame): Input DataFrame\n            fit (bool): Whether to fit the embedding model or use pre-fitted model\n            \n        Returns:\n            np.ndarray: Computed embeddings\n        \"\"\"\n        pass\n    \n    @property\n    @abstractmethod\n    def embedding_size(self) -> int:\n        \"\"\"\n        Returns the size of the embedding vector.\n        \n        Returns:\n            int: Size of embedding vector\n        \"\"\"\n        pass\n    \n    def compute_and_store_embeddings(\n        self, \n        df: pd.DataFrame, \n        entity_column: str,\n        batch_size: Optional[int] = None\n    ) -> None:\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            batch_size (Optional[int]): Batch size for processing large datasets\n        \"\"\"\n        unique_entities = df[entity_column].unique()\n        \n        if batch_size:\n            # Process in batches\n            for i in range(0, len(unique_entities), batch_size):\n                batch_entities = unique_entities[i:i + batch_size]\n                batch_df = df[df[entity_column].isin(batch_entities)].copy()\n                self._process_batch(batch_df, entity_column, batch_entities)\n        else:\n            self._process_batch(df, entity_column, unique_entities)\n        \n        self._is_fitted = True\n    \n    def _process_batch(\n        self, \n        df: pd.DataFrame, \n        entity_column: str, \n        entities: np.ndarray\n    ) -> None:\n        \"\"\"\n        Process a batch of entities and compute their embeddings.\n        \n        Args:\n            df (pd.DataFrame): Input DataFrame\n            entity_column (str): Column name containing entities\n            entities (np.ndarray): Array of entity names to process\n        \"\"\"\n        for entity in entities:\n            entity_df = df[df[entity_column] == entity].copy()\n            entity_df = self.preprocess(entity_df)\n            embedding = self.get_embedding(entity_df, fit=False)\n            \n            # Handle different embedding return types\n            if isinstance(embedding, (torch.Tensor, np.ndarray)):\n                self.embedding_dict[entity] = (\n                    embedding.mean(axis=0).cpu().numpy() \n                    if isinstance(embedding, torch.Tensor) \n                    else embedding.mean(axis=0)\n                )\n            else:\n                raise ValueError(f\"Unsupported embedding type: {type(embedding)}\")\n    \n    def get_entity_embedding(\n        self, \n        entity_name: str,\n        allow_missing: bool = True\n    ) -> np.ndarray:\n        \"\"\"\n        Retrieve embedding for a specific entity.\n        \n        Args:\n            entity_name (str): Name of the entity\n            allow_missing (bool): If True, return zero vector for missing entities\n            \n        Returns:\n            np.ndarray: Embedding vector for the entity\n            \n        Raises:\n            KeyError: If entity is missing and allow_missing is False\n        \"\"\"\n        if entity_name in self.embedding_dict:\n            return self.embedding_dict[entity_name]\n        elif allow_missing:\n            return np.zeros(self.embedding_size)\n        else:\n            raise KeyError(f\"No embedding found for entity: {entity_name}\")\n    \n    def get_multiple_embeddings(\n        self, \n        entity_names: List[str]\n    ) -> np.ndarray:\n        \"\"\"\n        Retrieve embeddings for multiple entities at once.\n        \n        Args:\n            entity_names (List[str]): List of entity names\n            \n        Returns:\n            np.ndarray: Stack of embedding vectors\n        \"\"\"\n        return np.stack([self.get_entity_embedding(name) for name in entity_names])\n    \n\n    def save_embedding(\n        self, \n        filepath: str, \n        metadata: Optional[Dict] = None\n    ) -> None:\n        \"\"\"\n        Save embeddings and metadata to disk.\n        \n        Args:\n            filepath (str): Path to save the embeddings\n            metadata (Optional[Dict]): Additional metadata to save\n        \"\"\"\n        if not self._is_fitted:\n            raise ValueError(\"Cannot save embeddings before computing them\")\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            'name': self.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, indent=2)\n                \n    def load_embeddings(self, filepath: str) -> bool:\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        Raises:\n            ValueError: If loaded embedding size doesn't match current size\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            self.name = save_data.get('name', self.name)\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            self._is_fitted = True\n            return True\n            \n        except Exception as e:\n            print(f\"Error loading embeddings: {str(e)}\")\n            return False","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-20T17:59:07.49072Z","iopub.execute_input":"2025-02-20T17:59:07.490995Z","iopub.status.idle":"2025-02-20T17:59:37.586845Z","shell.execute_reply.started":"2025-02-20T17:59:07.490971Z","shell.execute_reply":"2025-02-20T17:59:37.585913Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **One-Hot Encoding**","metadata":{}},{"cell_type":"code","source":"from sklearn.preprocessing import OneHotEncoder\n\nclass OneHotEmbedding(BaseEmbedding):\n    \"\"\"\n    One-hot encoding implementation of the BaseEmbedding class.\n    Handles both single and multiple column encodings.\n    \"\"\"\n    def __init__(self, columns: Union[str, List[str]], name: str = \"onehot\"):\n        \"\"\"\n        Initialize OneHotEmbedding.\n        \n        Args:\n            columns: Column name(s) to encode\n            name: Identifier for this embedding\n        \"\"\"\n        super().__init__(name=name)\n        self.columns = [columns] if isinstance(columns, str) else columns\n        self.encoder = OneHotEncoder(sparse=False, handle_unknown='ignore')\n        self._embedding_size = None\n        self.feature_names = None\n        self._is_fitted = False\n    \n    def preprocess(self, df: pd.DataFrame) -> pd.DataFrame:\n        \"\"\"\n        Select relevant columns for encoding.\n        \n        Args:\n            df: Input DataFrame\n            \n        Returns:\n            DataFrame with selected columns\n        \"\"\"\n        return df[self.columns]\n    \n    def get_embedding(self, df: pd.DataFrame, fit: bool = True) -> np.ndarray:\n        \"\"\"\n        Get one-hot encoded features.\n        \n        Args:\n            df: Input DataFrame\n            fit: Whether to fit the encoder or use pre-fitted encoder\n            \n        Returns:\n            Array of one-hot encoded features\n        \"\"\"\n        input_data = df[self.columns].values\n        \n        if fit:\n            encoded_features = self.encoder.fit_transform(input_data)\n            self._is_fitted = True\n            self.feature_names = self.encoder.get_feature_names_out(self.columns)\n        else:\n            if not self._is_fitted:\n                raise ValueError(\"Encoder must be fitted before transform\")\n            encoded_features = self.encoder.transform(input_data)\n        \n        self._embedding_size = encoded_features.shape[1]\n        \n        # Store embeddings for each unique combination\n        for idx, row in df[self.columns].iterrows():\n            key = tuple(row.values) if len(self.columns) > 1 else row.values[0]\n            self.embedding_dict[key] = encoded_features[idx]\n        \n        return encoded_features\n    \n    @property\n    def embedding_size(self) -> int:\n        \"\"\"\n        Get the size of the one-hot encoded vector.\n        \n        Returns:\n            Size of the embedding vector\n        \"\"\"\n        if self._embedding_size is None:\n            raise ValueError(\"Embedding size not set. Call get_embedding first.\")\n        return self._embedding_size\n    \n    def get_feature_names(self) -> List[str]:\n        \"\"\"\n        Get names of the one-hot encoded features.\n        \n        Returns:\n            List of feature names\n        \"\"\"\n        if self.feature_names is None:\n            raise ValueError(\"Feature names not available. Call get_embedding with fit=True first.\")\n        return self.feature_names.tolist()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-20T17:59:37.589878Z","iopub.execute_input":"2025-02-20T17:59:37.590393Z","iopub.status.idle":"2025-02-20T17:59:37.598531Z","shell.execute_reply.started":"2025-02-20T17:59:37.590363Z","shell.execute_reply":"2025-02-20T17:59:37.597468Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **SMILES Embeddings**","metadata":{}},{"cell_type":"code","source":"from rdkit import Chem\nfrom rdkit.Chem import Descriptors\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\nfrom sklearn.preprocessing import StandardScaler\nfrom dataclasses import dataclass\n\n@dataclass\nclass MoleculeFeatures:\n    \"\"\"Dataclass to store molecular features with their descriptions.\"\"\"\n    name: str\n    function: callable\n    description: str\n\nclass SMILESEmbedding(BaseEmbedding):\n    \"\"\"\n    SMILES-based molecular embedding using RDKit descriptors.\n    Converts SMILES strings into numerical feature vectors using chemical descriptors.\n    \"\"\"\n    \n    # Define molecular features as class attribute for easy access and modification\n    MOLECULAR_FEATURES = [\n        MoleculeFeatures(\"Molecular Weight\", Descriptors.MolWt, \"Molecular weight of the compound\"),\n        MoleculeFeatures(\"LogP\", Descriptors.MolLogP, \"Octanol-water partition coefficient\"),\n        MoleculeFeatures(\"TPSA\", CalcTPSA, \"Topological polar surface area\"),\n        MoleculeFeatures(\"Number of Atoms\", lambda m: m.GetNumAtoms(), \"Total number of atoms\"),\n        MoleculeFeatures(\"Number of Bonds\", lambda m: m.GetNumBonds(), \"Total number of bonds\"),\n        MoleculeFeatures(\"Number of Rotatable Bonds\", CalcNumRotatableBonds, \"Number of rotatable bonds\"),\n        MoleculeFeatures(\"Number of Hydrogen Bond Acceptors\", CalcNumHBA, \"Number of H-bond acceptors\"),\n        MoleculeFeatures(\"Number of Hydrogen Bond Donors\", CalcNumHBD, \"Number of H-bond donors\"),\n        MoleculeFeatures(\"Number of Rings\", Descriptors.RingCount, \"Total number of rings\"),\n        MoleculeFeatures(\"Number of Aromatic Rings\", rdMolDescriptors.CalcNumAromaticRings, \"Number of aromatic rings\"),\n        MoleculeFeatures(\"Number of Stereocenters\", \n                        lambda m: len(Chem.FindMolChiralCenters(m, includeUnassigned=True)),\n                        \"Number of stereogenic centers\"),\n        MoleculeFeatures(\"Fraction of sp3 Carbons\", CalcFractionCSP3, \"Fraction of sp3 hybridized carbons\"),\n        MoleculeFeatures(\"Balaban J Index\", Descriptors.BalabanJ, \"Topological connectivity index\"),\n        MoleculeFeatures(\"Bertz CT\", Descriptors.BertzCT, \"Complexity index\"),\n        MoleculeFeatures(\"QED Score\", QED.qed, \"Drug-likeness score\")\n    ]\n\n    def __init__(self, name: str = \"smiles\", handle_errors: bool = True):\n        \"\"\"\n        Initialize SMILESEmbedding.\n        \n        Args:\n            name: Identifier for this embedding\n            handle_errors: If True, return null values for invalid SMILES\n        \"\"\"\n        super().__init__(name=name)\n        self.scaler = StandardScaler()\n        self.handle_errors = handle_errors\n        self._feature_names = [f.name for f in self.MOLECULAR_FEATURES]\n        self._null_embedding = np.zeros(self.embedding_size)\n    \n    def preprocess(self, df: pd.DataFrame) -> pd.DataFrame:\n        \"\"\"Verify SMILES column exists and remove invalid entries.\"\"\"\n        if 'SMILES' not in df.columns:\n            raise ValueError(\"DataFrame must contain 'SMILES' column\")\n        \n        # Remove invalid SMILES\n        valid_mask = df['SMILES'].apply(lambda x: bool(x) and Chem.MolFromSmiles(x) is not None)\n        if not valid_mask.all() and not self.handle_errors:\n            invalid_count = (~valid_mask).sum()\n            raise ValueError(f\"Found {invalid_count} invalid SMILES strings\")\n        \n        return df[valid_mask]\n    \n    def create_molecule_embedding_dict(self, df: pd.DataFrame) -> None:\n        \"\"\"\n        Create a dictionary of molecule embeddings.\n        \n        Args:\n            df: DataFrame containing SMILES column\n        \"\"\"\n        df = self.preprocess(df)\n        self.compute_and_store_embeddings(df, 'SMILES')\n    \n    def get_embedding(self, df: pd.DataFrame, fit: bool = True) -> np.ndarray:\n        \"\"\"\n        Get molecular descriptor-based embeddings for SMILES strings.\n        \n        Args:\n            df: DataFrame containing SMILES column\n            fit: Whether to fit the scaler or use pre-fitted scaler\n            \n        Returns:\n            Array of molecular descriptors\n        \"\"\"\n        df = self.preprocess(df)\n        \n        # Extract features for each SMILES\n        features_list = []\n        for smiles in df['SMILES']:\n            try:\n                features = self.extract_smiles_info(smiles)\n                features_list.append([features[f.name] for f in self.MOLECULAR_FEATURES])\n            except Exception as e:\n                if self.handle_errors:\n                    features_list.append(self._null_embedding)\n                else:\n                    raise ValueError(f\"Error processing SMILES {smiles}: {str(e)}\")\n        \n        # Convert to array and scale\n        features_array = np.array(features_list)\n        if fit:\n            features_scaled = self.scaler.fit_transform(features_array)\n        else:\n            if not hasattr(self.scaler, 'mean_'):\n                raise ValueError(\"Scaler must be fitted before transform\")\n            features_scaled = self.scaler.transform(features_array)\n        \n        # Store embeddings\n        for idx, smiles in enumerate(df['SMILES']):\n            self.embedding_dict[smiles] = features_scaled[idx]\n        \n        return features_scaled\n    \n    @staticmethod\n    def extract_smiles_info(smiles: str) -> Dict[str, float]:\n        \"\"\"\n        Extract molecular descriptors from SMILES string.\n        \n        Args:\n            smiles: SMILES string\n            \n        Returns:\n            Dictionary of molecular descriptors\n        \"\"\"\n        if not smiles:\n            raise ValueError(\"Empty SMILES string\")\n            \n        mol = Chem.MolFromSmiles(smiles)\n        if mol is None:\n            raise ValueError(f\"Invalid SMILES string: {smiles}\")\n        \n        return {\n            feature.name: feature.function(mol)\n            for feature in SMILESEmbedding.MOLECULAR_FEATURES\n        }\n    \n    @property\n    def embedding_size(self) -> int:\n        \"\"\"Size of the molecular descriptor vector.\"\"\"\n        return len(self.MOLECULAR_FEATURES)\n    \n    def get_feature_importance(self, target_values: np.ndarray) -> pd.DataFrame:\n        \"\"\"\n        Calculate correlation between features and target values.\n        \n        Args:\n            target_values: Array of target values\n            \n        Returns:\n            DataFrame with feature importances\n        \"\"\"\n        if len(self.embedding_dict) == 0:\n            raise ValueError(\"No embeddings computed yet\")\n            \n        embeddings = np.stack(list(self.embedding_dict.values()))\n        correlations = np.corrcoef(embeddings.T, target_values.reshape(1, -1))[-1, :-1]\n        \n        return pd.DataFrame({\n            'Feature': self._feature_names,\n            'Correlation': correlations,\n            'Absolute Correlation': np.abs(correlations)\n        }).sort_values('Absolute Correlation', ascending=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-20T17:59:37.600804Z","iopub.execute_input":"2025-02-20T17:59:37.601171Z","iopub.status.idle":"2025-02-20T17:59:38.103986Z","shell.execute_reply.started":"2025-02-20T17:59:37.601137Z","shell.execute_reply":"2025-02-20T17:59:38.103311Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **Common Autoencoder**","metadata":{}},{"cell_type":"code","source":"import torch.nn as nn\nclass Autoencoder(nn.Module):\n    def __init__(self, input_size, hidden_size):\n        super().__init__()\n        self.encoder = nn.Sequential(\n            nn.Linear(input_size, hidden_size * 2),\n            nn.ReLU(),\n            nn.Linear(hidden_size * 2, hidden_size)\n        )\n        self.decoder = nn.Sequential(\n            nn.Linear(hidden_size, hidden_size * 2),\n            nn.ReLU(),\n            nn.Linear(hidden_size * 2, 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 Target_autoencoder(nn.Module):\n    \"\"\"Autoencoder module with matching architecture from the previous implementation.\"\"\"\n    \n    def __init__(self, input_size: int, hidden_size: int, emb_size: int):\n        \"\"\"\n        Initialize the autoencoder with encoder and decoder networks.\n        \n        Args:\n            input_size (int): Input dimension size\n            hidden_size (int): Size of hidden layers\n            emb_size (int): Size of the latent embedding\n        \"\"\"\n        super().__init__()\n        \n        self.encoder = nn.Sequential(\n            nn.Linear(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.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, input_size)\n        )\n    \n    def forward(self, x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:\n        \"\"\"\n        Forward pass through the autoencoder.\n        \n        Args:\n            x (torch.Tensor): Input tensor\n            \n        Returns:\n            Tuple[torch.Tensor, torch.Tensor]: Encoded and decoded tensors\n        \"\"\"\n        encoded = self.encoder(x)\n        decoded = self.decoder(encoded)\n        return encoded, decoded","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-20T17:59:38.104974Z","iopub.execute_input":"2025-02-20T17:59:38.105306Z","iopub.status.idle":"2025-02-20T17:59:38.113036Z","shell.execute_reply.started":"2025-02-20T17:59:38.105275Z","shell.execute_reply":"2025-02-20T17:59:38.112044Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **Target Embeddings**","metadata":{}},{"cell_type":"code","source":"import torch.nn as nn\nfrom dataclasses import dataclass\nfrom torch.utils.data import DataLoader, TensorDataset\n\nclass TargetEmbedding(BaseEmbedding):\n    \"\"\"\n    Target embedding class using dual autoencoders to learn compressed representations \n    of cell type and small molecule medians from gene expression data.\n    \"\"\"\n    \n    def __init__(\n        self,\n        emb_size: int = 256,\n        hidden_size: int = 1024,\n        input_size: int = 18211,\n        name: str = \"target_embedding\"\n    ):\n        \"\"\"\n        Initialize the target embedding class.\n        \n        Args:\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            input_size (int): Input dimension of gene expression values\n            name (str): Identifier for this embedding type\n        \"\"\"\n        super().__init__(name=name)\n        \n        self.emb_size = emb_size\n        self._embedding_size = emb_size * 2  # Combined size of both embeddings\n        self.input_size = input_size\n        \n        # Initialize autoencoders\n        self.cell_type_autoencoder = Target_autoencoder(input_size, hidden_size, emb_size)\n        self.sm_autoencoder = Target_autoencoder(input_size, hidden_size, emb_size)\n        \n        # Initialize storage for median values\n        self.cell_type_medians: Dict = {}\n        self.sm_medians: Dict = {}\n        self.cell_type_tensors: Dict = {}\n        self.sm_tensors: Dict = {}\n        \n    def preprocess(self, df: pd.DataFrame) -> pd.DataFrame:\n        \"\"\"\n        Preprocess the input DataFrame and compute medians if needed.\n        \n        Args:\n            df (pd.DataFrame): Input DataFrame with cell_type, sm_name, and gene expression columns\n            \n        Returns:\n            pd.DataFrame: Preprocessed DataFrame\n            \n        Raises:\n            ValueError: If required columns are missing\n        \"\"\"\n        df = df.copy()\n        \n        # Validate required columns\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 missing values in gene expression columns\n        gene_cols = [col for col in df.columns if col not in NON_GENE_COLUMNS]\n        df[gene_cols] = df[gene_cols].fillna(0)\n        \n        # Compute medians if not already computed\n        if not self.cell_type_medians:\n            self._compute_medians(df, gene_cols)\n            \n        return df\n    \n    def _compute_medians(self, df: pd.DataFrame, gene_cols: list) -> None:\n        \"\"\"\n        Compute median gene expression values for cell types and small molecules.\n        \n        Args:\n            df (pd.DataFrame): Input DataFrame\n            gene_cols (list): List of gene expression column names\n        \"\"\"\n        # Compute cell type medians\n        cell_type_medians = df.groupby('cell_type')[gene_cols].median()\n        self.cell_type_medians = {\n            ct: values.values for ct, values in cell_type_medians.iterrows()\n        }\n        self.cell_type_tensors = {\n            k: torch.tensor(v, dtype=torch.float32) \n            for k, v in self.cell_type_medians.items()\n        }\n        \n        # Compute small molecule medians\n        sm_medians = df.groupby('sm_name')[gene_cols].median()\n        self.sm_medians = {\n            sm: values.values for sm, values in sm_medians.iterrows()\n        }\n        self.sm_tensors = {\n            k: torch.tensor(v, dtype=torch.float32) \n            for k, v in self.sm_medians.items()\n        }\n    \n    def train_autoencoders(\n        self,\n        num_epochs: int = 100,\n        batch_size: int = 32,\n        learning_rate: float = 1e-3\n    ) -> None:\n        \"\"\"\n        Train both autoencoders to learn compressed representations.\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        # Prepare data\n        cell_type_data = torch.stack(list(self.cell_type_tensors.values()))\n        sm_data = torch.stack(list(self.sm_tensors.values()))\n        \n        # Initialize optimizers\n        optimizer = torch.optim.Adam(\n            list(self.cell_type_autoencoder.parameters()) + \n            list(self.sm_autoencoder.parameters()),\n            lr=learning_rate\n        )\n        criterion = nn.MSELoss()\n        \n        for epoch in range(num_epochs):\n            # Train cell type autoencoder\n            _, ct_decoded = self.cell_type_autoencoder(cell_type_data)\n            ct_loss = criterion(ct_decoded, cell_type_data)\n            \n            # Train small molecule autoencoder\n            _, sm_decoded = self.sm_autoencoder(sm_data)\n            sm_loss = criterion(sm_decoded, sm_data)\n            \n            # Combined loss\n            total_loss = ct_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: {ct_loss.item():.4f}, '\n                      f'SM Loss: {sm_loss.item():.4f}')\n        \n        # Store embeddings in the base class dictionary\n        self._store_embeddings()\n        self._is_fitted = True\n    \n    def _store_embeddings(self) -> None:\n        \"\"\"Store computed embeddings in the base class dictionary.\"\"\"\n        with torch.no_grad():\n            # Store cell type embeddings\n            for ct, tensor in self.cell_type_tensors.items():\n                encoded = self.cell_type_autoencoder.encoder(tensor.unsqueeze(0))\n                self.embedding_dict[f\"cell_type_{ct}\"] = encoded.squeeze(0).numpy()\n            \n            # Store small molecule embeddings\n            for sm, tensor in self.sm_tensors.items():\n                encoded = self.sm_autoencoder.encoder(tensor.unsqueeze(0))\n                self.embedding_dict[f\"sm_name_{sm}\"] = encoded.squeeze(0).numpy()\n    \n    def get_embedding(self, df: pd.DataFrame, fit: bool = True) -> np.ndarray:\n        \"\"\"\n        Get embeddings for the input data.\n        \n        Args:\n            df (pd.DataFrame): Input DataFrame\n            fit (bool): If True, train the autoencoders before getting embeddings\n            \n        Returns:\n            np.ndarray: Combined embeddings for cell types and small molecules\n            \n        Raises:\n            ValueError: If autoencoders are not trained and fit=False\n        \"\"\"\n        df = self.preprocess(df)\n        \n        if fit and not self._is_fitted:\n            self.train_autoencoders()\n        elif not self._is_fitted:\n            raise ValueError(\"Autoencoders must be trained before getting embeddings with fit=False\")\n        \n        with torch.no_grad():\n            # Get embeddings for each entity\n            cell_type_embeddings = []\n            sm_embeddings = []\n            \n            for _, row in df.iterrows():\n                # Get cell type embedding\n                ct_tensor = self.cell_type_tensors.get(row['cell_type'], \n                    torch.zeros(self.input_size, dtype=torch.float32))\n                ct_emb = self.cell_type_autoencoder.encoder(ct_tensor.unsqueeze(0))\n                cell_type_embeddings.append(ct_emb.squeeze(0))\n                \n                # Get small molecule embedding\n                sm_tensor = self.sm_tensors.get(row['sm_name'],\n                    torch.zeros(self.input_size, dtype=torch.float32))\n                sm_emb = self.sm_autoencoder.encoder(sm_tensor.unsqueeze(0))\n                sm_embeddings.append(sm_emb.squeeze(0))\n            \n            # Combine embeddings\n            cell_type_tensor = torch.stack(cell_type_embeddings)\n            sm_tensor = torch.stack(sm_embeddings)\n            combined = torch.cat([cell_type_tensor, sm_tensor], dim=1)\n            \n        return combined.numpy()\n    \n    @property\n    def embedding_size(self) -> int:\n        \"\"\"\n        Returns the size of the combined embedding vector.\n        \n        Returns:\n            int: Size of the combined embedding vector\n        \"\"\"\n        return self._embedding_size","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-20T17:59:38.114001Z","iopub.execute_input":"2025-02-20T17:59:38.114294Z","iopub.status.idle":"2025-02-20T17:59:38.135187Z","shell.execute_reply.started":"2025-02-20T17:59:38.114273Z","shell.execute_reply":"2025-02-20T17:59:38.134348Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **Morgan Fingerprint Embedding**","metadata":{}},{"cell_type":"code","source":"from rdkit import Chem\nfrom rdkit.Chem import AllChem, DataStructs\nimport torch.nn as nn\nfrom torch.utils.data import DataLoader, TensorDataset\n\nclass MorganFingerPrintEmbedding(BaseEmbedding):\n    \"\"\"\n    Morgan Fingerprint embedding class that generates and compresses molecular fingerprints\n    using an autoencoder architecture.\n    \"\"\"\n    \n    def __init__(self, hidden_size: int = 128, name: str = \"morgan_fingerprint\"):\n        \"\"\"\n        Initialize the Morgan Fingerprint embedding.\n        \n        Args:\n            hidden_size (int): Size of the compressed fingerprint representation\n            name (str): Identifier for this embedding type\n        \"\"\"\n        super().__init__(name=name)\n        self.hidden_size = hidden_size\n        self.autoencoder = None\n        self._input_size = 2048  # Default Morgan fingerprint size\n        \n    def preprocess(self, df: pd.DataFrame) -> pd.DataFrame:\n        \"\"\"\n        Preprocess the input DataFrame. Verifies SMILES column exists.\n        \n        Args:\n            df (pd.DataFrame): Input DataFrame with SMILES column\n            \n        Returns:\n            pd.DataFrame: Preprocessed DataFrame\n            \n        Raises:\n            ValueError: If SMILES column is missing\n        \"\"\"\n        if 'SMILES' not in df.columns:\n            raise ValueError(\"DataFrame must contain 'SMILES' column\")\n        return df\n    \n    def get_embedding(self, df: pd.DataFrame, fit: bool = True) -> np.ndarray:\n        \"\"\"\n        Generate embeddings for the molecules in the DataFrame.\n        \n        Args:\n            df (pd.DataFrame): Input DataFrame with SMILES column\n            fit (bool): Whether to train a new autoencoder or use existing one\n            \n        Returns:\n            np.ndarray: Compressed fingerprint embeddings\n            \n        Raises:\n            ValueError: If autoencoder is not trained and fit=False\n        \"\"\"\n        morgan_fp_list = df['SMILES'].apply(self.extract_morgan_fingerprint)\n        morgan_fp_array = np.stack([fp for fp in morgan_fp_list if fp is not None])\n        \n        if fit:\n            self.autoencoder = self.train_autoencoder(\n                morgan_fp_array, \n                self._input_size, \n                self.hidden_size\n            )\n        elif self.autoencoder is None:\n            raise ValueError(\"Autoencoder must be trained before getting embeddings with fit=False\")\n        \n        # Convert to tensor and get embeddings\n        with torch.no_grad():\n            embeddings = self.autoencoder.encoder(\n                torch.tensor(morgan_fp_array, dtype=torch.float32)\n            ).cpu().numpy()\n            \n        return embeddings\n    \n    def extract_morgan_fingerprint(\n        self, \n        smiles: str, \n        radius: int = 2, \n        nBits: int = 2048\n    ) -> Optional[np.ndarray]:\n        \"\"\"\n        Extract Morgan fingerprint from SMILES string.\n        \n        Args:\n            smiles (str): SMILES representation of molecule\n            radius (int): Morgan fingerprint radius\n            nBits (int): Number of bits in fingerprint\n            \n        Returns:\n            Optional[np.ndarray]: Fingerprint array or None if conversion fails\n        \"\"\"\n        if not smiles or not isinstance(smiles, str):\n            return None\n            \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(\n        self, \n        data: np.ndarray, \n        input_size: int, \n        hidden_size: int,\n        num_epochs: int = 50,\n        batch_size: int = 32,\n        learning_rate: float = 0.001\n    ) -> Autoencoder:\n        \"\"\"\n        Train the autoencoder model for fingerprint compression.\n        \n        Args:\n            data (np.ndarray): Input fingerprint data\n            input_size (int): Size of input fingerprints\n            hidden_size (int): Size of compressed representation\n            num_epochs (int): Number of training epochs\n            batch_size (int): Training batch size\n            learning_rate (float): Learning rate for optimization\n            \n        Returns:\n            Autoencoder: Trained autoencoder model\n        \"\"\"\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                loss = criterion(reconstructed, inputs)\n                \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) -> int:\n        \"\"\"\n        Returns the size of the compressed fingerprint representation.\n        \n        Returns:\n            int: Size of embedding vector\n        \"\"\"\n        return self.hidden_size","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-20T17:59:38.135971Z","iopub.execute_input":"2025-02-20T17:59:38.136246Z","iopub.status.idle":"2025-02-20T17:59:38.154136Z","shell.execute_reply.started":"2025-02-20T17:59:38.136218Z","shell.execute_reply":"2025-02-20T17:59:38.153284Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **Mol2Vec Embedding**","metadata":{}},{"cell_type":"code","source":"from gensim.models import word2vec\nfrom rdkit import Chem\nfrom mol2vec.features import mol2alt_sentence, MolSentence\n\nclass Mol2VecEmbedding(BaseEmbedding):\n    \"\"\"\n    Mol2Vec embedding class that generates molecular embeddings using a pre-trained Word2Vec model.\n    The embeddings are based on molecular substructures represented as \"sentences\".\n    \"\"\"\n    \n    def __init__(\n        self, \n        model_path: str,\n        name: str = \"mol2vec\",\n        radius: int = 1,\n        unseen_vec: Optional[np.ndarray] = None\n    ):\n        \"\"\"\n        Initialize the Mol2Vec embedding class.\n        \n        Args:\n            model_path (str): Path to the pre-trained Word2Vec model\n            name (str): Identifier for this embedding type\n            radius (int): Radius for molecular substructure generation\n            unseen_vec (Optional[np.ndarray]): Vector to use for unseen substructures\n        \"\"\"\n        super().__init__(name=name)\n        \n        # Load the pre-trained model\n        self.model = word2vec.Word2Vec.load(model_path)\n        self.keys = set(self.model.wv.key_to_index.keys())\n        self.radius = radius\n        \n        # Initialize unseen vector if not provided\n        if unseen_vec is None:\n            self.unseen_vec = np.zeros(self.embedding_size)\n        else:\n            self.unseen_vec = unseen_vec\n            \n        self.metadata.update({\n            'model_path': model_path,\n            'radius': radius,\n            'vocabulary_size': len(self.keys)\n        })\n    \n    def preprocess(self, df: pd.DataFrame) -> pd.DataFrame:\n        \"\"\"\n        Preprocess the input DataFrame by converting SMILES to molecules and generating\n        substructure sentences.\n        \n        Args:\n            df (pd.DataFrame): Input DataFrame with SMILES column\n            \n        Returns:\n            pd.DataFrame: Preprocessed DataFrame\n            \n        Raises:\n            ValueError: If SMILES column is missing or contains invalid molecules\n        \"\"\"\n        if 'SMILES' not in df.columns:\n            raise ValueError(\"DataFrame must contain 'SMILES' column\")\n        \n        df = df.copy()\n        \n        # Convert SMILES to molecules\n        df['mol'] = df['SMILES'].apply(self._smiles_to_mol)\n        \n        # Check for failed conversions\n        failed_smiles = df[df['mol'].isna()]['SMILES'].tolist()\n        if failed_smiles:\n            raise ValueError(f\"Failed to parse {len(failed_smiles)} SMILES strings: {failed_smiles[:5]}\")\n        \n        # Generate substructure sentences\n        df['sentence'] = df['mol'].apply(\n            lambda x: MolSentence(mol2alt_sentence(x, self.radius))\n        )\n        \n        return df\n    \n    def _smiles_to_mol(self, smiles: str) -> Optional[Chem.Mol]:\n        \"\"\"\n        Convert SMILES string to RDKit molecule with error handling.\n        \n        Args:\n            smiles (str): SMILES string\n            \n        Returns:\n            Optional[Chem.Mol]: RDKit molecule object or None if conversion fails\n        \"\"\"\n        try:\n            mol = Chem.MolFromSmiles(smiles)\n            if mol is None:\n                print(f\"Warning: Could not parse SMILES: {smiles}\")\n            return mol\n        except Exception as e:\n            print(f\"Error processing SMILES {smiles}: {str(e)}\")\n            return None\n    \n    def get_embedding(\n        self, \n        df: pd.DataFrame, \n        fit: bool = False\n    ) -> np.ndarray:\n        \"\"\"\n        Generate embeddings for the molecules in the DataFrame.\n        \n        Args:\n            df (pd.DataFrame): Input DataFrame with SMILES column\n            fit (bool): Not used for Mol2Vec (pre-trained model)\n            \n        Returns:\n            np.ndarray: Molecular embeddings\n        \"\"\"\n        df = self.preprocess(df)\n        \n        # Generate embeddings for each molecule\n        vectors = []\n        for _, row in df.iterrows():\n            vec = self.sentence_to_vector(\n                sentence=row['sentence'],\n                handle_unseen=True\n            )\n            vectors.append(vec)\n            \n            # Store in base class dictionary\n            self.embedding_dict[row['SMILES']] = vec\n        \n        embeddings = np.stack(vectors)\n        self._is_fitted = True\n        \n        return embeddings\n    \n    def sentence_to_vector(\n        self, \n        sentence: MolSentence,\n        handle_unseen: bool = True\n    ) -> np.ndarray:\n        \"\"\"\n        Convert a molecular sentence to a vector by averaging substructure vectors.\n        \n        Args:\n            sentence (MolSentence): Molecular sentence (substructures)\n            handle_unseen (bool): Whether to handle unseen substructures with unseen_vec\n            \n        Returns:\n            np.ndarray: Molecular embedding vector\n        \"\"\"\n        vectors = []\n        for word in sentence:\n            if word in self.keys:\n                vectors.append(self.model.wv.get_vector(word))\n            elif handle_unseen:\n                vectors.append(self.unseen_vec)\n                \n        if not vectors:\n            return self.unseen_vec\n            \n        return np.mean(vectors, axis=0)\n    \n    def get_entity_embedding(\n        self, \n        smiles: str,\n        allow_missing: bool = True\n    ) -> np.ndarray:\n        \"\"\"\n        Get embedding for a specific molecule by SMILES.\n        \n        Args:\n            smiles (str): SMILES string of the molecule\n            allow_missing (bool): If True, compute embedding for missing molecules\n            \n        Returns:\n            np.ndarray: Molecular embedding vector\n        \"\"\"\n        if smiles in self.embedding_dict:\n            return self.embedding_dict[smiles]\n        elif allow_missing:\n            mol = self._smiles_to_mol(smiles)\n            if mol is not None:\n                sentence = MolSentence(mol2alt_sentence(mol, self.radius))\n                return self.sentence_to_vector(sentence)\n            return np.zeros(self.embedding_size)\n        else:\n            raise KeyError(f\"No embedding found for SMILES: {smiles}\")\n    \n    def save_embeddings(\n        self, \n        filepath: str,\n        metadata: Optional[Dict] = None\n    ) -> None:\n        \"\"\"\n        Save embeddings and metadata to disk.\n        \n        Args:\n            filepath (str): Path to save the embeddings\n            metadata (Optional[Dict]): Additional metadata to save\n        \"\"\"\n        if metadata:\n            self.metadata.update(metadata)\n        super().save_embeddings(filepath, self.metadata)\n    \n    @property\n    def embedding_size(self) -> int:\n        \"\"\"\n        Returns the size of the molecular embedding vector.\n        \n        Returns:\n            int: Size of embedding vector\n        \"\"\"\n        return len(self.model.wv.get_vector(next(iter(self.keys))))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-20T17:59:38.154933Z","iopub.execute_input":"2025-02-20T17:59:38.155188Z","iopub.status.idle":"2025-02-20T17:59:38.1988Z","shell.execute_reply.started":"2025-02-20T17:59:38.155156Z","shell.execute_reply":"2025-02-20T17:59:38.198095Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **Chemberta Embedding**","metadata":{}},{"cell_type":"code","source":"from transformers import AutoModelForMaskedLM, AutoTokenizer\n\nclass ChemBERTaEmbedding(BaseEmbedding):\n    \"\"\"\n    ChemBERTa embedding implementation of the BaseEmbedding class.\n    Provides molecular embeddings using the ChemBERTa transformer model.\n    \"\"\"\n    def __init__(self, \n                 columns: str = 'SMILES',\n                 name: str = \"chemberta\",\n                 model_name: str = \"DeepChem/ChemBERTa-77M-MTR\",\n                 embedding_type: str = 'mean_pooling',\n                 padding: bool = False):\n        \"\"\"\n        Initialize ChemBERTaEmbedding.\n\n        Args:\n            columns: Column name containing SMILES strings\n            name: Identifier for this embedding\n            model_name: Name of the pretrained ChemBERTa model\n            embedding_type: Type of embedding ('cls' or 'mean_pooling')\n            padding: Whether to use padding in tokenization\n        \"\"\"\n        super().__init__(name=name)\n        self.columns = columns\n        self.model_name = model_name\n        self.embedding_type = embedding_type\n        self.padding = padding\n        \n        # Initialize model and tokenizer\n        self.model = AutoModelForMaskedLM.from_pretrained(self.model_name)\n        self.tokenizer = AutoTokenizer.from_pretrained(self.model_name)\n        self.model.eval()\n        \n        self._embedding_size = None\n        self.feature_names = None\n        self._is_fitted = False\n\n    def preprocess(self, df: pd.DataFrame) -> pd.DataFrame:\n        \"\"\"\n        Select relevant columns for encoding.\n\n        Args:\n            df: Input DataFrame\n\n        Returns:\n            DataFrame with selected SMILES column\n        \"\"\"\n        if self.columns not in df.columns:\n            raise ValueError(f\"Column {self.columns} not found in DataFrame\")\n        return df[[self.columns]]\n\n    def get_embedding(self, df: pd.DataFrame, fit: bool = True) -> np.ndarray:\n        \"\"\"\n        Get ChemBERTa embeddings for SMILES strings.\n\n        Args:\n            df: Input DataFrame\n            fit: Whether to store embeddings in dictionary (ignored as model is pre-trained)\n\n        Returns:\n            Array of molecular embeddings\n        \"\"\"\n        smiles_list = df[self.columns].tolist()\n        embeddings_cls, embeddings_mean = self._featurize_batch(smiles_list)\n\n        if self.embedding_type == 'cls':\n            embeddings = embeddings_cls\n            self.feature_names = [f'ChemBERTa_cls_{i}' for i in range(embeddings.shape[1])]\n        elif self.embedding_type == 'mean_pooling':\n            embeddings = embeddings_mean\n            self.feature_names = [f'ChemBERTa_mean_{i}' for i in range(embeddings.shape[1])]\n        else:\n            raise ValueError(\"embedding_type must be 'cls' or 'mean_pooling'\")\n\n        self._embedding_size = embeddings.shape[1]\n        self._is_fitted = True\n\n        # Store embeddings for each SMILES string\n        for idx, smiles in enumerate(smiles_list):\n            self.embedding_dict[smiles] = embeddings[idx]\n\n        return embeddings\n\n    def _featurize_batch(self, smiles_list: List[str]) -> tuple:\n        \"\"\"\n        Generate embeddings for a batch of SMILES strings.\n\n        Args:\n            smiles_list: List of SMILES strings to encode\n\n        Returns:\n            Tuple of (CLS embeddings, mean pooled embeddings)\n        \"\"\"\n        embeddings_cls = []\n        embeddings_mean = []\n\n        with torch.no_grad():\n            for smiles in smiles_list:\n                encoded_input = self.tokenizer(\n                    smiles,\n                    return_tensors=\"pt\",\n                    padding=self.padding,\n                    truncation=True\n                )\n                model_output = self.model(**encoded_input)\n\n                # Get CLS token embedding\n                embedding_cls = model_output[0][:, 0, :]\n                embeddings_cls.append(embedding_cls)\n\n                # Get mean pooled embedding\n                embedding_mean = torch.mean(model_output[0], 1)\n                embeddings_mean.append(embedding_mean)\n\n        # Stack the tensors and convert to numpy\n        embeddings_cls = torch.cat(embeddings_cls).numpy()\n        embeddings_mean = torch.cat(embeddings_mean).numpy()\n        return embeddings_cls, embeddings_mean\n\n    @property\n    def embedding_size(self) -> int:\n        \"\"\"\n        Get the size of the embedding vector.\n\n        Returns:\n            Size of the embedding vector\n        \"\"\"\n        if self._embedding_size is None:\n            # This will be set correctly during first call to get_embedding\n            return self.model.config.hidden_size\n        return self._embedding_size\n\n    def get_feature_names(self) -> List[str]:\n        \"\"\"\n        Get names of the embedding features.\n\n        Returns:\n            List of feature names\n        \"\"\"\n        if self.feature_names is None:\n            raise ValueError(\"Feature names not available. Call get_embedding first.\")\n        return self.feature_names","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-20T17:59:38.199546Z","iopub.execute_input":"2025-02-20T17:59:38.199782Z","iopub.status.idle":"2025-02-20T17:59:39.578735Z","shell.execute_reply.started":"2025-02-20T17:59:38.199762Z","shell.execute_reply":"2025-02-20T17:59:39.577972Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **Molformer Embedding**","metadata":{}},{"cell_type":"code","source":"from transformers import AutoModel, AutoTokenizer\n\nclass MolformerEmbedding(BaseEmbedding):\n    \"\"\"\n    Molformer embedding implementation of the BaseEmbedding class.\n    Provides molecular embeddings using the MolFormer transformer model.\n    \"\"\"\n    def __init__(self, \n                 columns: str = 'SMILES',\n                 name: str = \"molformer\",\n                 model_path: str = 'ibm/MoLFormer-XL-both-10pct',\n                 padding: bool = True):\n        \"\"\"\n        Initialize MolformerEmbedding.\n\n        Args:\n            columns: Column name containing SMILES strings\n            name: Identifier for this embedding\n            model_path: Path to the pretrained MolFormer model\n            padding: Whether to use padding in tokenization (should be True for this model)\n        \"\"\"\n        super().__init__(name=name)\n        self.columns = columns\n        self.model_path = model_path\n        self.padding = padding\n\n        # Initialize model and tokenizer\n        self.model = AutoModel.from_pretrained(model_path, trust_remote_code=True)\n        self.tokenizer = AutoTokenizer.from_pretrained(model_path, trust_remote_code=True)\n        self.model.eval()\n\n        self._embedding_size = self.model.config.hidden_size  # Set from model config\n        self.feature_names = None\n        self._is_fitted = False\n\n    def preprocess(self, df: pd.DataFrame) -> pd.DataFrame:\n        \"\"\"\n        Select relevant columns for encoding.\n\n        Args:\n            df: Input DataFrame\n\n        Returns:\n            DataFrame with selected SMILES column\n        \"\"\"\n        if self.columns not in df.columns:\n            raise ValueError(f\"Column {self.columns} not found in DataFrame\")\n        return df[[self.columns]]\n\n    def get_embedding(self, df: pd.DataFrame, fit: bool = True) -> np.ndarray:\n        \"\"\"\n        Get MolFormer embeddings for SMILES strings.\n\n        Args:\n            df: Input DataFrame\n            fit: Whether to store embeddings in dictionary (ignored as model is pre-trained)\n\n        Returns:\n            Array of molecular embeddings\n        \"\"\"\n        smiles_list = df[self.columns].tolist()\n        embeddings = self._featurize_batch(smiles_list)\n\n        # Generate feature names if not already created\n        if self.feature_names is None:\n            self.feature_names = [f'Molformer_{i}' for i in range(embeddings.shape[1])]\n\n        self._is_fitted = True\n\n        # Store embeddings for each SMILES string\n        for idx, smiles in enumerate(smiles_list):\n            self.embedding_dict[smiles] = embeddings[idx]\n\n        return embeddings\n\n    def _featurize_batch(self, smiles_list: List[str]) -> np.ndarray:\n        \"\"\"\n        Generate embeddings for a batch of SMILES strings.\n\n        Args:\n            smiles_list: List of SMILES strings to encode\n\n        Returns:\n            Array of molecular embeddings\n        \"\"\"\n        inputs = self.tokenizer(\n            smiles_list,\n            padding=self.padding,\n            return_tensors=\"pt\"\n        )\n\n        with torch.no_grad():\n            outputs = self.model(**inputs)\n            \n        # Use pooler_output for the sentence-level representation\n        embeddings = outputs.pooler_output.cpu().numpy()\n        return embeddings\n\n    @property\n    def embedding_size(self) -> int:\n        \"\"\"\n        Get the size of the embedding vector.\n\n        Returns:\n            Size of the embedding vector\n        \"\"\"\n        return self._embedding_size\n\n    def get_feature_names(self) -> List[str]:\n        \"\"\"\n        Get names of the embedding features.\n\n        Returns:\n            List of feature names\n        \"\"\"\n        if self.feature_names is None:\n            raise ValueError(\"Feature names not available. Call get_embedding first.\")\n        return self.feature_names","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-20T17:59:39.579513Z","iopub.execute_input":"2025-02-20T17:59:39.580063Z","iopub.status.idle":"2025-02-20T17:59:39.591433Z","shell.execute_reply.started":"2025-02-20T17:59:39.58004Z","shell.execute_reply":"2025-02-20T17:59:39.590101Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **SmoleBart Embedding**","metadata":{}},{"cell_type":"code","source":"from transformers import AutoModel, AutoTokenizer\n\nclass SmoleBartEmbedding(BaseEmbedding):\n    \"\"\"\n    SmoleBART embedding implementation of the BaseEmbedding class.\n    Provides molecular embeddings using the SmoleBART transformer model.\n    \"\"\"\n    def __init__(self, \n                 columns: str = 'SMILES',\n                 name: str = \"smolebart\",\n                 model_path: str = \"UdS-LSV/smole-bart\",\n                 embedder: Literal[\"encoder\", \"decoder\"] = \"encoder\",\n                 padding: bool = True):\n        \"\"\"\n        Initialize SmoleBartEmbedding.\n\n        Args:\n            columns: Column name containing SMILES strings\n            name: Identifier for this embedding\n            model_path: Path to the pretrained SmoleBART model\n            embedder: Which part of the model to use for embeddings (\"encoder\" or \"decoder\")\n            padding: Whether to use padding in tokenization\n        \"\"\"\n        super().__init__(name=name)\n        self.columns = columns\n        self.model_path = model_path\n        self.embedder = embedder\n        self.padding = padding\n\n        # Initialize model and tokenizer\n        self.model = AutoModel.from_pretrained(model_path, trust_remote_code=True)\n        self.tokenizer = AutoTokenizer.from_pretrained(model_path, trust_remote_code=True)\n        self.model.eval()\n\n        self._embedding_size = self.model.config.hidden_size  # Set from model config\n        self.feature_names = None\n        self._is_fitted = False\n\n    def preprocess(self, df: pd.DataFrame) -> pd.DataFrame:\n        \"\"\"\n        Select relevant columns for encoding.\n\n        Args:\n            df: Input DataFrame\n\n        Returns:\n            DataFrame with selected SMILES column\n        \"\"\"\n        if self.columns not in df.columns:\n            raise ValueError(f\"Column {self.columns} not found in DataFrame\")\n        return df[[self.columns]]\n\n    def get_embedding(self, df: pd.DataFrame, fit: bool = True) -> np.ndarray:\n        \"\"\"\n        Get SmoleBART embeddings for SMILES strings.\n\n        Args:\n            df: Input DataFrame\n            fit: Whether to store embeddings in dictionary (ignored as model is pre-trained)\n\n        Returns:\n            Array of molecular embeddings\n        \"\"\"\n        smiles_list = df[self.columns].tolist()\n        embeddings = self._featurize_batch(smiles_list)\n\n        # Generate feature names if not already created\n        if self.feature_names is None:\n            self.feature_names = [f'SmoleBart_{i}' for i in range(embeddings.shape[1])]\n\n        self._is_fitted = True\n\n        # Store embeddings for each SMILES string\n        for idx, smiles in enumerate(smiles_list):\n            self.embedding_dict[smiles] = embeddings[idx]\n\n        return embeddings\n\n    def _featurize_batch(self, smiles_list: List[str]) -> np.ndarray:\n        \"\"\"\n        Generate embeddings for a batch of SMILES strings.\n\n        Args:\n            smiles_list: List of SMILES strings to encode\n\n        Returns:\n            Array of molecular embeddings\n\n        Raises:\n            ValueError: If embedder type is invalid\n        \"\"\"\n        inputs = self.tokenizer(\n            smiles_list,\n            padding=self.padding,\n            return_tensors=\"pt\"\n        )\n\n        with torch.no_grad():\n            outputs = self.model(**inputs)\n\n        # Select appropriate output based on embedder type\n        if self.embedder == \"encoder\":\n            hidden_states = outputs.encoder_last_hidden_state\n        elif self.embedder == \"decoder\":\n            hidden_states = outputs.last_hidden_state\n        else:\n            raise ValueError(\"embedder must be either 'encoder' or 'decoder'\")\n\n        # Apply mean pooling and convert to numpy\n        embeddings = hidden_states.mean(dim=1).cpu().numpy()\n        return embeddings\n\n    @property\n    def embedding_size(self) -> int:\n        \"\"\"\n        Get the size of the embedding vector.\n\n        Returns:\n            Size of the embedding vector\n\n        Raises:\n            ValueError: If embedding size is not yet set\n        \"\"\"\n        if self._embedding_size is None:\n            raise ValueError(\"Embedding size not set. Call get_embedding first.\")\n        return self._embedding_size\n\n    def get_feature_names(self) -> List[str]:\n        \"\"\"\n        Get names of the embedding features.\n\n        Returns:\n            List of feature names\n\n        Raises:\n            ValueError: If feature names are not yet available\n        \"\"\"\n        if self.feature_names is None:\n            raise ValueError(\"Feature names not available. Call get_embedding first.\")\n        return self.feature_names","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-20T17:59:39.592782Z","iopub.execute_input":"2025-02-20T17:59:39.593095Z","iopub.status.idle":"2025-02-20T17:59:39.623825Z","shell.execute_reply.started":"2025-02-20T17:59:39.593056Z","shell.execute_reply":"2025-02-20T17:59:39.622904Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **SelfFormer Embedding**","metadata":{}},{"cell_type":"code","source":"from transformers import RobertaConfig, RobertaModel, RobertaTokenizer\nfrom pandarallel import pandarallel\n\nclass SelfFormerEmbedding(BaseEmbedding):\n    \"\"\"\n    SELFormer embedding implementation of the BaseEmbedding class.\n    Provides molecular embeddings using the SELFormer transformer model with batch processing.\n    \"\"\"\n    def __init__(self,\n                 columns: str = 'SELFIES',\n                 name: str = \"selformer\",\n                 model_name: str = None,\n                 tokenizer_path: Optional[str] = None,\n                 padding: bool = True,\n                 batch_size: int = 32,\n                 nb_workers: int = 5,\n                 device: Optional[str] = None):\n        \"\"\"\n        Initialize SelfFormerEmbedding.\n\n        Args:\n            columns: Column name containing SELFIES strings\n            name: Identifier for this embedding\n            model_name: Path or identifier for the pretrained SELFormer model\n            tokenizer_path: Optional path to load tokenizer from\n            padding: Whether to use padding in tokenization\n            batch_size: Number of samples to process in a single batch\n            nb_workers: Number of parallel workers for pandarallel\n            device: Device to run the model on (\"cuda\" or \"cpu\")\n        \"\"\"\n        if model_name is None:\n            raise ValueError(\"model_name must be provided\")\n            \n        super().__init__(name=name)\n        \n        # Disable parallelism warnings\n        os.environ[\"TOKENIZERS_PARALLELISM\"] = \"false\"\n        os.environ[\"WANDB_DISABLED\"] = \"true\"\n        \n        self.columns = columns\n        self.model_name = model_name\n        self.padding = padding\n        self.batch_size = batch_size\n        self.device = device if device else (\"cuda\" if torch.cuda.is_available() else \"cpu\")\n        \n        # Load model with hidden states enabled\n        config = RobertaConfig.from_pretrained(model_name)\n        config.output_hidden_states = True\n        self.model = RobertaModel.from_pretrained(model_name, config=config).to(self.device)\n        self.model.eval()\n        \n        # Load tokenizer\n        self.tokenizer = RobertaTokenizer.from_pretrained(tokenizer_path or model_name)\n        \n        # Initialize parallel processing\n        pandarallel.initialize(nb_workers=nb_workers, progress_bar=True)\n        \n        self._embedding_size = self.model.config.hidden_size\n        self.feature_names = None\n        self._is_fitted = False\n\n    def preprocess(self, df: pd.DataFrame) -> pd.DataFrame:\n        \"\"\"\n        Select relevant columns for encoding.\n\n        Args:\n            df: Input DataFrame\n\n        Returns:\n            DataFrame with selected SELFIES column\n        \"\"\"\n        if self.columns not in df.columns:\n            raise ValueError(f\"Column {self.columns} not found in DataFrame\")\n        return df[[self.columns]]\n\n    def get_embedding(self, df: pd.DataFrame, fit: bool = True) -> np.ndarray:\n        \"\"\"\n        Get SELFormer embeddings for SELFIES strings using batch processing.\n\n        Args:\n            df: Input DataFrame\n            fit: Whether to store embeddings in dictionary (ignored as model is pre-trained)\n\n        Returns:\n            Array of molecular embeddings\n        \"\"\"\n        selfies_list = df[self.columns].tolist()\n        \n        try:\n            embeddings = self._featurize_batch(selfies_list)\n            \n            # Verify embedding dimensions\n            if embeddings.shape[1] != self._embedding_size:\n                raise ValueError(f\"Embedding dimension mismatch. Expected {self._embedding_size}, \"\n                               f\"got {embeddings.shape[1]}\")\n\n            # Generate feature names if not already created\n            if self.feature_names is None:\n                self.feature_names = [f'SELFormer_{i}' for i in range(embeddings.shape[1])]\n\n            self._is_fitted = True\n\n            # Store embeddings for each SELFIES string\n            for idx, selfie in enumerate(selfies_list):\n                self.embedding_dict[selfie] = embeddings[idx]\n\n            return embeddings\n            \n        except Exception as e:\n            print(f\"Error during embedding generation: {str(e)}\")\n            print(f\"DataFrame shape: {df.shape}\")\n            print(f\"Expected embedding size: {self._embedding_size}\")\n            raise\n\n    def _featurize_batch(self, selfies_list: List[str]) -> np.ndarray:\n        \"\"\"\n        Generate embeddings for batches of SELFIES strings with explicit shape handling.\n\n        Args:\n            selfies_list: List of SELFIES strings to encode\n\n        Returns:\n            Array of molecular embeddings\n        \"\"\"\n        all_embeddings = []\n\n        for i in range(0, len(selfies_list), self.batch_size):\n            batch = selfies_list[i:i + self.batch_size]\n            \n            try:\n                # Tokenize batch with explicit padding\n                encoded_input = self.tokenizer(\n                    batch,\n                    add_special_tokens=True,\n                    max_length=512,\n                    padding='max_length',  # Force padding to max_length\n                    truncation=True,\n                    return_tensors=\"pt\"\n                ).to(self.device)\n\n                # Generate embeddings\n                with torch.no_grad():\n                    output = self.model(**encoded_input)\n                    \n                # Get the last hidden state\n                last_hidden_state = output.last_hidden_state  # Shape: [batch_size, seq_len, hidden_size]\n                \n                # Create attention mask for proper mean pooling\n                attention_mask = encoded_input['attention_mask']  # Shape: [batch_size, seq_len]\n                \n                # Expand attention mask to 3D\n                attention_mask_expanded = attention_mask.unsqueeze(-1)  # Shape: [batch_size, seq_len, 1]\n                \n                # Apply attention mask before mean pooling\n                masked_hidden_states = last_hidden_state * attention_mask_expanded\n                \n                # Sum tokens and divide by actual sequence length\n                sum_hidden_states = torch.sum(masked_hidden_states, dim=1)  # Shape: [batch_size, hidden_size]\n                sequence_lengths = torch.sum(attention_mask, dim=1, keepdim=True)  # Shape: [batch_size, 1]\n                \n                # Compute mean while handling padding\n                embeddings = sum_hidden_states / sequence_lengths\n                \n                all_embeddings.append(embeddings.cpu())\n\n            except Exception as e:\n                print(f\"Error processing batch {i}-{i+self.batch_size}: {str(e)}\")\n                print(f\"Batch shapes - Input: {encoded_input['input_ids'].shape}, \"\n                      f\"Attention mask: {encoded_input['attention_mask'].shape}, \"\n                      f\"Hidden states: {last_hidden_state.shape}\")\n                raise\n\n        if not all_embeddings:\n            raise ValueError(\"No embeddings were generated successfully\")\n\n        # Concatenate all batches and convert to numpy\n        final_embeddings = torch.cat(all_embeddings, dim=0).numpy()\n        \n        # Verify final shape\n        expected_size = (len(selfies_list), self._embedding_size)\n        if final_embeddings.shape != expected_size:\n            raise ValueError(f\"Unexpected embedding shape. Expected {expected_size}, got {final_embeddings.shape}\")\n            \n        return final_embeddings\n\n\n    @property\n    def embedding_size(self) -> int:\n        \"\"\"\n        Get the size of the embedding vector.\n\n        Returns:\n            Size of the embedding vector\n        \"\"\"\n        return self._embedding_size\n\n    def get_feature_names(self) -> List[str]:\n        \"\"\"\n        Get names of the embedding features.\n\n        Returns:\n            List of feature names\n\n        Raises:\n            ValueError: If feature names are not yet available\n        \"\"\"\n        if self.feature_names is None:\n            raise ValueError(\"Feature names not available. Call get_embedding first.\")\n        return self.feature_names","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-20T17:59:39.624717Z","iopub.execute_input":"2025-02-20T17:59:39.624966Z","iopub.status.idle":"2025-02-20T17:59:39.91107Z","shell.execute_reply.started":"2025-02-20T17:59:39.624939Z","shell.execute_reply":"2025-02-20T17:59:39.910412Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **Embedding Combiner**","metadata":{}},{"cell_type":"code","source":"class MultiEmbedding(BaseEmbedding):\n    \"\"\"\n    A class that combines multiple embedding types into a single embedding.\n    \n    This allows for flexible experimentation with different embedding combinations\n    while maintaining a consistent interface.\n    \"\"\"\n    \n    def __init__(\n        self,\n        embeddings: List[BaseEmbedding],\n        name: str = \"multi_embedding\",\n        combination_method: Literal[\"concat\", \"sum\", \"average\", \"weighted\"] = \"concat\",\n        weights: Optional[List[float]] = None\n    ):\n        \"\"\"\n        Initialize the multi-embedding class.\n        \n        Args:\n            embeddings (List[BaseEmbedding]): List of embedding objects to combine\n            name (str): Identifier for this embedding type\n            combination_method (str): Method to combine embeddings (\"concat\", \"sum\", \"average\", or \"weighted\")\n            weights (Optional[List[float]]): Weights for weighted combination method\n        \"\"\"\n        super().__init__(name=name)\n        \n        if not embeddings:\n            raise ValueError(\"At least one embedding must be provided\")\n        \n        self.embeddings = embeddings\n        self.combination_method = combination_method\n        \n        # Handle weights for weighted combination\n        if combination_method == \"weighted\":\n            if weights is None:\n                self.weights = [1.0 / len(embeddings)] * len(embeddings)\n            elif len(weights) != len(embeddings):\n                raise ValueError(f\"Number of weights ({len(weights)}) must match number of embeddings ({len(embeddings)})\")\n            else:\n                # Normalize weights to sum to 1\n                total = sum(weights)\n                self.weights = [w / total for w in weights]\n        else:\n            self.weights = None\n        \n        # Check if embeddings are fitted\n        for i, emb in enumerate(embeddings):\n            if not emb._is_fitted:\n                print(f\"Warning: Embedding {i} ({emb.__class__.__name__}) is not fitted yet\")\n        \n        # Update metadata\n        self.metadata = {\n            'num_embeddings': len(embeddings),\n            'embedding_types': [emb.__class__.__name__ for emb in embeddings],\n            # 'embedding_sizes': [emb.embedding_size for emb in embeddings],\n            'combination_method': combination_method,\n            'weights': self.weights if self.weights else None\n        }\n    \n    def preprocess(self, df: pd.DataFrame) -> pd.DataFrame:\n        \"\"\"\n        Preprocess the input DataFrame using all embedding models.\n        \n        Args:\n            df (pd.DataFrame): Input DataFrame\n            \n        Returns:\n            pd.DataFrame: Preprocessed DataFrame\n        \"\"\"\n        # Preprocess with all embeddings\n        result_df = df.copy()\n        for emb in self.embeddings:\n            result_df = emb.preprocess(result_df.copy())\n        return result_df\n    \n    def get_embedding(self, df: pd.DataFrame, fit: bool = True) -> np.ndarray:\n        \"\"\"\n        Get combined embeddings from all embedding models.\n        \n        Args:\n            df (pd.DataFrame): Input DataFrame\n            fit (bool): Whether to fit the embedding models if not already fitted\n            \n        Returns:\n            np.ndarray: Combined embeddings\n        \"\"\"\n        # Get embeddings from all models\n        all_embeddings = [emb.get_embedding(df, fit=fit) for emb in self.embeddings]\n        \n        # Combine embeddings based on the specified method\n        if self.combination_method == \"concat\":\n            return np.concatenate(all_embeddings, axis=1)\n        \n        elif self.combination_method in [\"sum\", \"average\", \"weighted\"]:\n            # Check if all embeddings have the same shape\n            if not all(emb.shape[1] == all_embeddings[0].shape[1] for emb in all_embeddings):\n                raise ValueError(f\"Cannot {self.combination_method} embeddings of different dimensions\")\n            \n            if self.combination_method == \"sum\":\n                return sum(all_embeddings)\n            \n            elif self.combination_method == \"average\":\n                return sum(all_embeddings) / len(all_embeddings)\n            \n            elif self.combination_method == \"weighted\":\n                weighted_sum = sum(w * emb for w, emb in zip(self.weights, all_embeddings))\n                return weighted_sum\n        \n        else:\n            raise ValueError(f\"Unsupported combination method: {self.combination_method}\")\n    \n    def compute_and_store_embeddings(\n        self,\n        df: pd.DataFrame,\n        entity_column: str,\n        batch_size: Optional[int] = None\n    ) -> None:\n        \"\"\"\n        Compute and store combined embeddings.\n        \n        Args:\n            df (pd.DataFrame): Input DataFrame\n            entity_column (str): Column name containing entities\n            batch_size (Optional[int]): Batch size for processing\n        \"\"\"\n        # First ensure all underlying embeddings are computed\n        for emb in self.embeddings:\n            if not emb._is_fitted:\n                emb.compute_and_store_embeddings(df, entity_column, batch_size)\n        \n        # Then compute combined embeddings\n        super().compute_and_store_embeddings(df, entity_column, batch_size)\n    \n    def get_entity_embedding(\n        self,\n        entity_name: str,\n        allow_missing: bool = True\n    ) -> np.ndarray:\n        \"\"\"\n        Get combined embedding for a specific entity.\n        \n        Args:\n            entity_name (str): Name of the entity\n            allow_missing (bool): If True, compute embedding for missing entities\n            \n        Returns:\n            np.ndarray: Combined embedding vector\n        \"\"\"\n        if entity_name in self.embedding_dict:\n            return self.embedding_dict[entity_name]\n        \n        # Get embeddings from all models\n        all_embeddings = [\n            emb.get_entity_embedding(entity_name, allow_missing) \n            for emb in self.embeddings\n        ]\n        \n        # Combine embeddings based on the specified method\n        if self.combination_method == \"concat\":\n            combined = np.concatenate(all_embeddings)\n        \n        elif self.combination_method in [\"sum\", \"average\", \"weighted\"]:\n            # Check if all embeddings have the same shape\n            if not all(emb.shape == all_embeddings[0].shape for emb in all_embeddings):\n                raise ValueError(f\"Cannot {self.combination_method} embeddings of different dimensions\")\n            \n            if self.combination_method == \"sum\":\n                combined = sum(all_embeddings)\n            \n            elif self.combination_method == \"average\":\n                combined = sum(all_embeddings) / len(all_embeddings)\n            \n            elif self.combination_method == \"weighted\":\n                combined = sum(w * emb for w, emb in zip(self.weights, all_embeddings))\n        \n        else:\n            raise ValueError(f\"Unsupported combination method: {self.combination_method}\")\n        \n        # Store for future use\n        self.embedding_dict[entity_name] = combined\n        return combined\n    \n    @property\n    def embedding_size(self) -> int:\n        \"\"\"\n        Returns the size of the combined embedding vector.\n        \n        Returns:\n            int: Size of the combined embedding vector\n            \n        Raises:\n            ValueError: If embedding sizes are incompatible with the combination method\n        \"\"\"\n        if self.combination_method == \"concat\":\n            return sum(emb.embedding_size for emb in self.embeddings)\n        \n        elif self.combination_method in [\"sum\", \"average\", \"weighted\"]:\n            # Check if all embeddings have the same size\n            sizes = [emb.embedding_size for emb in self.embeddings]\n            if not all(size == sizes[0] for size in sizes):\n                raise ValueError(\n                    f\"Cannot {self.combination_method} embeddings of different dimensions: {sizes}\"\n                )\n            return sizes[0]\n        \n        else:\n            raise ValueError(f\"Unsupported combination method: {self.combination_method}\")\n    \n    def save_embeddings(\n        self,\n        filepath: str,\n        metadata: Optional[Dict] = None\n    ) -> None:\n        \"\"\"\n        Save embeddings and metadata to disk.\n        \n        Args:\n            filepath (str): Path to save the embeddings\n            metadata (Optional[Dict]): Additional metadata to save\n        \"\"\"\n        combined_metadata = self.metadata.copy()\n        if metadata:\n            combined_metadata.update(metadata)\n        super().save_embeddings(filepath, combined_metadata)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-20T17:59:39.911859Z","iopub.execute_input":"2025-02-20T17:59:39.912084Z","iopub.status.idle":"2025-02-20T17:59:39.928507Z","shell.execute_reply.started":"2025-02-20T17:59:39.912065Z","shell.execute_reply":"2025-02-20T17:59:39.927257Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def select_top_variable_genes(df, k=256, exclude_controls=True, controls=[\"Belinostat\", \"Dabrafenib\"]):\n    \"\"\"\n    Selects the top k most variable genes based on standard deviation.\n\n    Args:\n        df (pd.DataFrame): Gene expression dataframe with multi-index (\"cell_type\", \"sm_name\").\n        k (int): Number of top variable genes to select.\n        exclude_controls (bool): Whether to exclude control drugs.\n        controls (list): List of control drugs to exclude.\n\n    Returns:\n        List of top k gene names.\n    \"\"\"\n    if exclude_controls:\n        df = df.loc[~df.index.get_level_values(\"sm_name\").isin(controls)]\n    \n    # Compute standard deviation per gene\n    gene_variability = df.iloc[:, 3:].std(axis=0)  # Avoid non-numeric columns\n    \n    # Select top k most variable genes\n    top_genes = gene_variability.sort_values(ascending=False).head(k).index.tolist()\n    \n    return top_genes","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-20T17:59:39.929606Z","iopub.execute_input":"2025-02-20T17:59:39.92986Z","iopub.status.idle":"2025-02-20T17:59:39.948181Z","shell.execute_reply.started":"2025-02-20T17:59:39.929837Z","shell.execute_reply":"2025-02-20T17:59:39.947322Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **Benchmarking Models**","metadata":{}},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\nfrom abc import ABC, abstractmethod\nfrom sklearn.metrics import mean_squared_error, r2_score\nfrom typing import Dict, List, Tuple, Any, Callable, Optional, Union, Type\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom dataclasses import dataclass\nimport time\nfrom torch.utils.data import DataLoader, TensorDataset\nfrom sklearn.base import BaseEstimator\nfrom sklearn.metrics import explained_variance_score\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\ndef evaluate_evs_score(y_true, y_pred):\n    return explained_variance_score(y_true, y_pred, multioutput='variance_weighted')\n\n@dataclass\nclass BenchmarkResult:\n    \"\"\"Store results for a single model with multiple runs.\"\"\"\n    model_name: str\n    embedding_name: str\n    mse_mean: float\n    mse_std: float\n    mrrmse_mean: float\n    mrrmse_std: float\n    r2_mean: float\n    r2_std: float\n    evs_mean: float\n    evs_std: float\n    training_time: float\n    inference_time: float\n    embedding_size: int\n    memory_usage: float\n\n\nclass ModelFactory(ABC):\n    \"\"\"Abstract factory for creating models with specific input sizes.\"\"\"\n    \n    @abstractmethod\n    def create_model(self, input_size: int) -> Any:\n        \"\"\"Create a new model instance with the specified input size.\"\"\"\n        pass\n    \n    @property\n    @abstractmethod\n    def name(self) -> str:\n        \"\"\"Model name.\"\"\"\n        pass\n\nclass BaseModel(ABC):\n    \"\"\"Abstract base class for all models (both PyTorch and sklearn).\"\"\"\n    \n    @abstractmethod\n    def fit(self, X: np.ndarray, y: np.ndarray) -> None:\n        \"\"\"Train the model.\"\"\"\n        pass\n    \n    @abstractmethod\n    def predict(self, X: np.ndarray) -> np.ndarray:\n        \"\"\"Make predictions.\"\"\"\n        pass\n    \n    @property\n    @abstractmethod\n    def name(self) -> str:\n        \"\"\"Model name.\"\"\"\n        pass\n\nclass PyTorchModelFactory(ModelFactory):\n    \"\"\"Factory for creating PyTorch models with specific architectures.\"\"\"\n    \n    def __init__(\n        self,\n        model_class: Type[nn.Module],\n        model_params: Dict[str, Any],\n        criterion: Callable,\n        optimizer_class: Type[torch.optim.Optimizer],\n        optimizer_params: Dict[str, Any],\n        batch_size: int = 32,\n        num_epochs: int = 50,\n        early_stopping: bool = True,\n        early_stopping_patience: int = 5,\n        device: str = \"cuda\" if torch.cuda.is_available() else \"cpu\",\n        name: str = \"PyTorch Model\"\n    ):\n        self.model_class = model_class\n        self.model_params = model_params\n        self.criterion = criterion\n        self.optimizer_class = optimizer_class\n        self.optimizer_params = optimizer_params\n        self.batch_size = batch_size\n        self.num_epochs = num_epochs\n        self.early_stopping = early_stopping\n        self.early_stopping_patience = early_stopping_patience\n        self.device = device\n        self._name = name\n    \n    def create_model(self, input_size: int) -> 'PyTorchModel':\n        \"\"\"Create a new PyTorch model with the specified input size.\"\"\"\n        model_params = {**self.model_params, 'input_size': input_size}\n        model = self.model_class(**model_params)\n        \n        return PyTorchModel(\n            model=model,\n            criterion=self.criterion,\n            optimizer_class=self.optimizer_class,\n            optimizer_params=self.optimizer_params,\n            batch_size=self.batch_size,\n            num_epochs=self.num_epochs,\n            early_stopping = self.early_stopping,\n            early_stopping_patience=self.early_stopping_patience,\n            device=self.device,\n            name=self._name\n        )\n    \n    @property\n    def name(self) -> str:\n        return self._name\n\nclass SklearnModelFactory(ModelFactory):\n    \"\"\"Factory for creating sklearn models.\"\"\"\n    \n    def __init__(\n        self,\n        model_class: Type[BaseEstimator],\n        model_params: Dict[str, Any],\n        name: str = \"Sklearn Model\"\n    ):\n        self.model_class = model_class\n        self.model_params = model_params\n        self._name = name\n    \n    def create_model(self, input_size: int) -> 'SklearnModel':\n        \"\"\"Create a new sklearn model (input_size is ignored for most sklearn models).\"\"\"\n        model = self.model_class(**self.model_params)\n        return SklearnModel(model=model, name=self._name)\n    \n    @property\n    def name(self) -> str:\n        return self._name\n\nclass PyTorchModel(BaseModel):\n    \"\"\"Wrapper for PyTorch models to conform to the BaseModel interface.\"\"\"\n    \n    def __init__(\n        self,\n        model: nn.Module,\n        criterion: Callable,\n        optimizer_class: torch.optim.Optimizer,\n        optimizer_params: Dict[str, Any],\n        batch_size: int = 32,\n        num_epochs: int = 50,\n        early_stopping: bool = True,\n        early_stopping_patience: int = 5,\n        device: str = \"cuda\" if torch.cuda.is_available() else \"cpu\",\n        name: str = \"PyTorch Model\"\n    ):\n        self.model = model.to(device)\n        self.criterion = criterion\n        self.optimizer_class = optimizer_class\n        self.optimizer_params = optimizer_params\n        self.batch_size = batch_size\n        self.num_epochs = num_epochs\n        self.early_stopping = early_stopping\n        self.early_stopping_patience = early_stopping_patience\n        self.device = device\n        self._name = name\n    \n    def fit(self, X: np.ndarray, y: np.ndarray) -> None:\n        # Create data loaders\n        train_data = TensorDataset(\n            torch.tensor(X, dtype=torch.float32),\n            torch.tensor(y, dtype=torch.float32)\n        )\n        train_loader = DataLoader(\n            train_data, \n            batch_size=self.batch_size, \n            shuffle=True\n        )\n        \n        optimizer = self.optimizer_class(\n            self.model.parameters(),\n            **self.optimizer_params\n        )\n        \n        best_loss = float('inf')\n        patience_counter = 0\n        \n        for epoch in range(self.num_epochs):\n            self.model.train()\n            epoch_loss = 0.0\n            \n            for batch_X, batch_y in train_loader:\n                batch_X = batch_X.to(self.device)\n                batch_y = batch_y.to(self.device)\n                \n                optimizer.zero_grad()\n                outputs = self.model(batch_X)\n                loss = self.criterion(outputs, batch_y)\n                loss.backward()\n                optimizer.step()\n                \n                epoch_loss += loss.item()\n            \n            # Early stopping\n            if epoch_loss < best_loss:\n                best_loss = epoch_loss\n                patience_counter = 0\n            else:\n                patience_counter += 1\n                if self.early_stopping and patience_counter >= self.early_stopping_patience:\n                    break\n    \n    def predict(self, X: np.ndarray) -> np.ndarray:\n        self.model.eval()\n        X_tensor = torch.tensor(X, dtype=torch.float32).to(self.device)\n        with torch.no_grad():\n            return self.model(X_tensor).cpu().numpy()\n    \n    @property\n    def name(self) -> str:\n        return self._name\n\nclass SklearnModel(BaseModel):\n    \"\"\"Wrapper for sklearn models to conform to the BaseModel interface.\"\"\"\n    \n    def __init__(self, model: BaseEstimator, name: str = \"Sklearn Model\"):\n        self.model = model\n        self._name = name\n    \n    def fit(self, X: np.ndarray, y: np.ndarray) -> None:\n        self.model.fit(X, y)\n    \n    def predict(self, X: np.ndarray) -> np.ndarray:\n        return self.model.predict(X)\n    \n    @property\n    def name(self) -> str:\n        return self._name","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-20T17:59:39.94907Z","iopub.execute_input":"2025-02-20T17:59:39.949432Z","iopub.status.idle":"2025-02-20T17:59:39.973255Z","shell.execute_reply.started":"2025-02-20T17:59:39.949386Z","shell.execute_reply":"2025-02-20T17:59:39.9723Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **Sklearn Models**","metadata":{}},{"cell_type":"code","source":"from sklearn.base import BaseEstimator\nfrom sklearn.linear_model import LinearRegression\nfrom sklearn.tree import DecisionTreeRegressor\nfrom sklearn.neighbors import KNeighborsRegressor\nfrom sklearn.neural_network import MLPRegressor\n\nclass LinearRegressionModel(SklearnModel):\n    \"\"\"Linear Regression wrapper.\"\"\"\n    def __init__(self):\n        super().__init__(\n            model=LinearRegression(),\n            name=\"Linear Regression\"\n        )\n\nclass DecisionTreeModel(SklearnModel):\n    \"\"\"Decision Tree wrapper.\"\"\"\n    def __init__(\n        self,\n        max_depth: int = 7,\n        min_samples_split: int = 10\n    ):\n        super().__init__(\n            model=DecisionTreeRegressor(\n                max_depth=max_depth,\n                min_samples_split=min_samples_split,\n                random_state=42\n            ),\n            name=\"Decision Tree\"\n        )\n\nclass KNeighborsModel(SklearnModel):\n    \"\"\"K-Nearest Neighbors wrapper.\"\"\"\n    def __init__(\n        self,\n        n_neighbors: int = 7,\n        weights: str = 'uniform'\n    ):\n        super().__init__(\n            model=KNeighborsRegressor(\n                n_neighbors=n_neighbors,\n                weights=weights\n            ),\n            name=\"KNN\"\n        )\n\nclass MLPRegressorModel(SklearnModel):\n    \"\"\"Neural Network wrapper.\"\"\"\n    def __init__(\n        self,\n        hidden_layer_sizes: tuple = (64, 32),\n        alpha: float = 0.01,\n        max_iter: int = 1000\n    ):\n        super().__init__(\n            model=MLPRegressor(\n                hidden_layer_sizes=hidden_layer_sizes,\n                alpha=alpha,\n                max_iter=max_iter,\n                random_state=42\n            ),\n            name=\"MLP\"\n        )\n\n# Model Factories\nclass LinearRegressionFactory(ModelFactory):\n    \"\"\"Factory for Linear Regression.\"\"\"\n    def __init__(self):\n        self._name = \"Linear Regression\"\n    \n    def create_model(self, input_size: int) -> LinearRegressionModel:\n        return LinearRegressionModel()\n    \n    @property\n    def name(self) -> str:\n        return self._name\n\nclass DecisionTreeFactory(ModelFactory):\n    \"\"\"Factory for Decision Tree.\"\"\"\n    def __init__(\n        self,\n        max_depth: int = 7,\n        min_samples_split: int = 10\n    ):\n        self.max_depth = max_depth\n        self.min_samples_split = min_samples_split\n        self._name = \"Decision Tree\"\n    \n    def create_model(self, input_size: int) -> DecisionTreeModel:\n        return DecisionTreeModel(\n            max_depth=self.max_depth,\n            min_samples_split=self.min_samples_split\n        )\n    \n    @property\n    def name(self) -> str:\n        return self._name\n\nclass KNeighborsFactory(ModelFactory):\n    \"\"\"Factory for KNN.\"\"\"\n    def __init__(\n        self,\n        n_neighbors: int = 7,\n        weights: str = 'uniform'\n    ):\n        self.n_neighbors = n_neighbors\n        self.weights = weights\n        self._name = \"KNN\"\n    \n    def create_model(self, input_size: int) -> KNeighborsModel:\n        return KNeighborsModel(\n            n_neighbors=self.n_neighbors,\n            weights=self.weights\n        )\n    \n    @property\n    def name(self) -> str:\n        return self._name\n\nclass MLPRegressorFactory(ModelFactory):\n    \"\"\"Factory for MLP.\"\"\"\n    def __init__(\n        self,\n        hidden_layer_sizes: tuple = (64, 32),\n        alpha: float = 0.01,\n        max_iter: int = 1000\n    ):\n        self.hidden_layer_sizes = hidden_layer_sizes\n        self.alpha = alpha\n        self.max_iter = max_iter\n        self._name = \"MLP\"\n    \n    def create_model(self, input_size: int) -> MLPRegressorModel:\n        return MLPRegressorModel(\n            hidden_layer_sizes=self.hidden_layer_sizes,\n            alpha=self.alpha,\n            max_iter=self.max_iter\n        )\n    \n    @property\n    def name(self) -> str:\n        return self._name","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-20T17:59:39.974314Z","iopub.execute_input":"2025-02-20T17:59:39.974724Z","iopub.status.idle":"2025-02-20T17:59:40.009673Z","shell.execute_reply.started":"2025-02-20T17:59:39.974689Z","shell.execute_reply":"2025-02-20T17:59:40.008921Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **Sklearn Models With TSVD**","metadata":{}},{"cell_type":"code","source":"from sklearn.base import BaseEstimator, TransformerMixin\nfrom sklearn.linear_model import LinearRegression\nfrom sklearn.tree import DecisionTreeRegressor\nfrom sklearn.neighbors import KNeighborsRegressor\nfrom sklearn.neural_network import MLPRegressor\nfrom sklearn.decomposition import TruncatedSVD\n\nclass TSVDBaseModel(BaseModel):\n    \"\"\"Base class for sklearn models with TSVD dimensionality reduction.\"\"\"\n    \n    def __init__(\n        self,\n        model: BaseEstimator,\n        n_components: int = 50,\n        name: str = \"TSVD Base Model\"\n    ):\n        self.model = model\n        self.n_components = n_components\n        self._name = name\n        self.tsvd = TruncatedSVD(n_components=n_components)\n        \n    def fit(self, X: np.ndarray, y: np.ndarray) -> None:\n        \"\"\"Fit both TSVD and the underlying model.\"\"\"\n        # Apply TSVD to target variables\n        y_reduced = self.tsvd.fit_transform(y)\n        # Fit the model with reduced targets\n        self.model.fit(X, y_reduced)\n        \n    def predict(self, X: np.ndarray) -> np.ndarray:\n        \"\"\"Make predictions and inverse transform them back to original space.\"\"\"\n        # Get predictions in reduced space\n        y_pred_reduced = self.model.predict(X)\n        # Transform back to original space\n        return self.tsvd.inverse_transform(y_pred_reduced)\n    \n    @property\n    def name(self) -> str:\n        return self._name\n\nclass LinearRegressionTSVD(TSVDBaseModel):\n    \"\"\"Linear Regression with TSVD.\"\"\"\n    def __init__(self, n_components: int = 50):\n        super().__init__(\n            model=LinearRegression(),\n            n_components=n_components,\n            name=\"Linear Regression TSVD\"\n        )\n\nclass DecisionTreeTSVD(TSVDBaseModel):\n    \"\"\"Decision Tree with TSVD.\"\"\"\n    def __init__(\n        self,\n        n_components: int = 50,\n        max_depth: int = 7,\n        min_samples_split: int = 10\n    ):\n        super().__init__(\n            model=DecisionTreeRegressor(\n                max_depth=max_depth,\n                min_samples_split=min_samples_split,\n                random_state=42\n            ),\n            n_components=n_components,\n            name=\"Decision Tree TSVD\"\n        )\n\nclass KNeighborsTSVD(TSVDBaseModel):\n    \"\"\"K-Nearest Neighbors with TSVD.\"\"\"\n    def __init__(\n        self,\n        n_components: int = 50,\n        n_neighbors: int = 7,\n        weights: str = 'uniform'\n    ):\n        super().__init__(\n            model=KNeighborsRegressor(\n                n_neighbors=n_neighbors,\n                weights=weights\n            ),\n            n_components=n_components,\n            name=\"KNN TSVD\"\n        )\n\nclass MLPRegressorTSVD(TSVDBaseModel):\n    \"\"\"Neural Network with TSVD.\"\"\"\n    def __init__(\n        self,\n        n_components: int = 50,\n        hidden_layer_sizes: tuple = (64, 32),\n        alpha: float = 0.01,\n        max_iter: int = 1000\n    ):\n        super().__init__(\n            model=MLPRegressor(\n                hidden_layer_sizes=hidden_layer_sizes,\n                alpha=alpha,\n                max_iter=max_iter,\n                random_state=42\n            ),\n            n_components=n_components,\n            name=\"MLP TSVD\"\n        )\n\n# Model Factories\nclass LinearRegressionTSVDFactory(ModelFactory):\n    \"\"\"Factory for Linear Regression with TSVD.\"\"\"\n    def __init__(self, n_components: int = 50):\n        self.n_components = n_components\n        self._name = \"Linear Regression TSVD\"\n    \n    def create_model(self, input_size: int) -> LinearRegressionTSVD:\n        return LinearRegressionTSVD(n_components=self.n_components)\n    \n    @property\n    def name(self) -> str:\n        return self._name\n\nclass DecisionTreeTSVDFactory(ModelFactory):\n    \"\"\"Factory for Decision Tree with TSVD.\"\"\"\n    def __init__(\n        self,\n        n_components: int = 50,\n        max_depth: int = 7,\n        min_samples_split: int = 10\n    ):\n        self.n_components = n_components\n        self.max_depth = max_depth\n        self.min_samples_split = min_samples_split\n        self._name = \"Decision Tree TSVD\"\n    \n    def create_model(self, input_size: int) -> DecisionTreeTSVD:\n        return DecisionTreeTSVD(\n            n_components=self.n_components,\n            max_depth=self.max_depth,\n            min_samples_split=self.min_samples_split\n        )\n    \n    @property\n    def name(self) -> str:\n        return self._name\n\nclass KNeighborsTSVDFactory(ModelFactory):\n    \"\"\"Factory for KNN with TSVD.\"\"\"\n    def __init__(\n        self,\n        n_components: int = 50,\n        n_neighbors: int = 7,\n        weights: str = 'uniform'\n    ):\n        self.n_components = n_components\n        self.n_neighbors = n_neighbors\n        self.weights = weights\n        self._name = \"KNN TSVD\"\n    \n    def create_model(self, input_size: int) -> KNeighborsTSVD:\n        return KNeighborsTSVD(\n            n_components=self.n_components,\n            n_neighbors=self.n_neighbors,\n            weights=self.weights\n        )\n    \n    @property\n    def name(self) -> str:\n        return self._name\n\nclass MLPRegressorTSVDFactory(ModelFactory):\n    \"\"\"Factory for MLP with TSVD.\"\"\"\n    def __init__(\n        self,\n        n_components: int = 50,\n        hidden_layer_sizes: tuple = (64, 32),\n        alpha: float = 0.01,\n        max_iter: int = 1000\n    ):\n        self.n_components = n_components\n        self.hidden_layer_sizes = hidden_layer_sizes\n        self.alpha = alpha\n        self.max_iter = max_iter\n        self._name = \"MLP TSVD\"\n    \n    def create_model(self, input_size: int) -> MLPRegressorTSVD:\n        return MLPRegressorTSVD(\n            n_components=self.n_components,\n            hidden_layer_sizes=self.hidden_layer_sizes,\n            alpha=self.alpha,\n            max_iter=self.max_iter\n        )\n    \n    @property\n    def name(self) -> str:\n        return self._name","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-20T17:59:40.01066Z","iopub.execute_input":"2025-02-20T17:59:40.010983Z","iopub.status.idle":"2025-02-20T17:59:40.027888Z","shell.execute_reply.started":"2025-02-20T17:59:40.010951Z","shell.execute_reply":"2025-02-20T17:59:40.026911Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **LSTM Model**","metadata":{}},{"cell_type":"code","source":"class LSTMPredictionModel(nn.Module):\n    \"\"\"LSTM-based prediction model for gene expression.\"\"\"\n    def __init__(self, input_size: int, hidden_size: int = 128, output_size: int = 18211):\n        super().__init__()\n        self.lstm = nn.LSTM(input_size, hidden_size, num_layers=2, batch_first=True)\n        self.linear = nn.Sequential(\n            nn.Linear(hidden_size, 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.head = 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.head(out)\n        return out","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-20T17:59:40.028851Z","iopub.execute_input":"2025-02-20T17:59:40.029111Z","iopub.status.idle":"2025-02-20T17:59:40.047393Z","shell.execute_reply.started":"2025-02-20T17:59:40.02909Z","shell.execute_reply":"2025-02-20T17:59:40.046235Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **Benchmarker**","metadata":{}},{"cell_type":"code","source":"class EmbeddingBenchmark:\n    def __init__(\n        self,\n        embedding_methods: Dict[str, BaseEmbedding],\n        model_factories: List[ModelFactory],\n        n_runs: int = 5,\n        only_significant = False\n    ):\n        self.embedding_methods = embedding_methods\n        self.model_factories = model_factories\n        self.n_runs = n_runs\n        self.results: List[BenchmarkResult] = []\n        self.gene_cols = self.get_significant_gene_columns(df) if only_significant else self.get_gene_columns(df) \n        self.combined_embedding_size = 0\n        self.embedding_ranges = {}\n        self.feature_map = {}\n        self.only_significant = only_significant\n\n       \n        \n    def prepare_data(\n        self,\n        df: pd.DataFrame,\n        test_size: float = 0.2,\n        random_state: int = 42\n    ) -> Tuple[pd.DataFrame, pd.DataFrame, np.ndarray, np.ndarray]:\n        \"\"\"\n        Prepare data splits for benchmarking.\n        Ensures proper index alignment for all embedding methods.\n        \"\"\"\n        df['SELFIES'] = df['SMILES'].apply(smiles_to_selfies)\n        \n        # Split data\n        train_df, test_df = train_test_split(\n            df, test_size=test_size, random_state=random_state\n        )\n        \n        # Reset indices to ensure proper alignment\n        train_df = train_df.reset_index(drop=True)\n        test_df = test_df.reset_index(drop=True)\n        \n        # Extract target values\n        y_train = train_df[self.gene_cols].values\n        y_test = test_df[self.gene_cols].values\n        \n        return train_df, test_df, y_train, y_test\n        \n    def get_gene_columns(self, df: pd.DataFrame) -> List[str]:\n        \"\"\"Get list of gene columns by excluding known non-gene columns.\"\"\"\n        return [col for col in df.columns if col not in NON_GENE_COLUMNS]\n\n    def get_significant_gene_columns(self, df: pd.DataFrame) -> List[str]:\n        \"\"\"Get list of gene columns by excluding non-significant gene columns and non-gene columns.\"\"\"\n        # return [col for col in df.columns if col not in NON_GENE_COLUMNS]\n        target_cols = ['sm_lincs_id','SMILES','control',]\n        targets = df.copy()\n        targets.drop(columns=target_cols, inplace=True)\n        targets.set_index([\"cell_type\", \"sm_name\"], inplace=True)\n        # Select top 256 variable genes\n        top_genes = select_top_variable_genes(targets, k=256, exclude_controls=True)\n        return top_genes\n    \n    def evaluate(\n        self,\n        model_factory: ModelFactory,\n        embedding_name: str,\n        embedding_method: BaseEmbedding,\n        train_df: pd.DataFrame,\n        test_df: pd.DataFrame,\n        y_train: np.ndarray,\n        y_test: np.ndarray\n    ) -> BenchmarkResult:\n        \"\"\"Evaluate a single model with a specific embedding method.\"\"\"\n        run_results = []\n        \n        # Compute embeddings once\n        start_time = time.time()\n        train_embeddings = embedding_method.get_embedding(train_df, fit=True)\n        test_embeddings = embedding_method.get_embedding(test_df, fit=False)\n        embedding_time = time.time() - start_time\n        \n        # Create model with correct input size\n        model = model_factory.create_model(embedding_method.embedding_size)\n        \n        for run in range(self.n_runs):\n            print(f\"Run {run + 1}/{self.n_runs}\")\n            \n            # Train model\n            start_time = time.time()\n            model.fit(train_embeddings, y_train)\n            training_time = time.time() - start_time\n            \n            # Evaluate\n            start_time = time.time()\n            y_pred = model.predict(test_embeddings)\n            inference_time = time.time() - start_time\n            \n            # Compute metrics\n            mse = mean_squared_error(y_test, y_pred)\n            mrrmse = RMSE_rowwise_loss_numpy(y_pred, y_test)\n            r2 = r2_score(y_test, y_pred)\n            evs = evaluate_evs_score(y_test, y_pred)\n            \n            run_results.append({\n                'mse': mse,\n                'mrrmse': mrrmse,\n                'r2': r2,\n                'evs': evs,\n                'training_time': training_time,\n                'inference_time': inference_time,\n                'memory_usage': (train_embeddings.nbytes + test_embeddings.nbytes) / (1024 * 1024)\n            })\n        \n        # Compute statistics\n        return BenchmarkResult(\n            model_name=model_factory.name,\n            embedding_name=embedding_name,\n            mse_mean=np.mean([r['mse'] for r in run_results]),\n            mse_std=np.std([r['mse'] for r in run_results]),\n            mrrmse_mean=np.mean([r['mrrmse'] for r in run_results]),\n            mrrmse_std=np.std([r['mrrmse'] for r in run_results]),\n            r2_mean=np.mean([r['r2'] for r in run_results]),\n            r2_std=np.std([r['r2'] for r in run_results]),\n            evs_mean=np.mean([r['evs'] for r in run_results]),\n            evs_std=np.std([r['evs'] for r in run_results]),\n            training_time=np.mean([r['training_time'] for r in run_results]),\n            inference_time=np.mean([r['inference_time'] for r in run_results]),\n            embedding_size=embedding_method.embedding_size,\n            memory_usage=run_results[0]['memory_usage']\n        )\n        \n    \n    def run_benchmark(self, df: pd.DataFrame) -> pd.DataFrame:\n        \"\"\"Run benchmark for all combinations of embeddings and models.\"\"\"\n        df = df.copy()\n        train_df, test_df, y_train, y_test = self.prepare_data(df)\n        \n        for name, method in self.embedding_methods.items():\n            print(f\"\\nEvaluating {name} embedding...\")\n            \n            for model_factory in self.model_factories:\n                print(f\"Testing with {model_factory.name}...\")\n                try:\n                    result = self.evaluate(\n                        model_factory, name, method,\n                        train_df, test_df,\n                        y_train, y_test\n                    )\n                    self.results.append(result)\n                except Exception as e:\n                    print(f\"Error evaluating {name} embedding with {model_factory.name}: {str(e)}\")\n        \n        return pd.DataFrame([vars(r) for r in self.results])\n\n    \n    def visualize_results(self) -> None:\n        \"\"\"Create detailed visualizations of benchmark results using clear bar charts.\"\"\"\n        if not self.results:\n            raise ValueError(\"No benchmark results available\")\n        \n        results_df = pd.DataFrame([vars(r) for r in self.results])\n        \n        # Set style\n        plt.style.use('seaborn')\n        \n        # Colors for different models\n        model_colors = sns.color_palette(\"Set2\", n_colors=len(results_df['model_name'].unique()))\n        \n        # 1. Individual Performance Metrics (separate by metric)\n        metrics = [\n            ('mrrmse_mean', 'mrrmse_std', 'MRRMSE'),\n            ('mse_mean', 'mse_std', 'MSE'),\n            ('r2_mean', 'r2_std', 'R²'),\n            ('evs_mean', 'evs_std', 'Explained Variance Score (EVS)')\n        ]\n        \n        for metric_mean, metric_std, metric_name in metrics:\n            plt.figure(figsize=(15, 6))\n            \n            # Create positions for bars\n            embeddings = results_df['embedding_name'].unique()\n            models = results_df['model_name'].unique()\n            x = np.arange(len(embeddings))\n            width = 0.8 / len(models)\n            \n            # Plot bars for each model\n            for i, model in enumerate(models):\n                model_data = results_df[results_df['model_name'] == model]\n                plt.bar(x + i*width - width*len(models)/2 + width/2, \n                       model_data[metric_mean],\n                       width,\n                       label=model,\n                       color=model_colors[i],\n                       yerr=model_data[metric_std],\n                       capsize=3)\n            \n            plt.xlabel('Embedding Method')\n            plt.ylabel(metric_name)\n            plt.title(f'{metric_name} by Embedding Method and Model')\n            plt.xticks(x, embeddings, rotation=45)\n            plt.legend(bbox_to_anchor=(1.05, 1), loc='upper left')\n            plt.tight_layout()\n            plt.show()\n        \n        # 2. Computational Metrics\n        comp_metrics = [\n            ('training_time', 'Training Time (s)'),\n            ('inference_time', 'Inference Time (s)'),\n            ('memory_usage', 'Memory Usage (MB)')\n        ]\n        \n        for metric, metric_label in comp_metrics:\n            plt.figure(figsize=(15, 6))\n            \n            # Create positions for bars\n            x = np.arange(len(embeddings))\n            width = 0.8 / len(models)\n            \n            # Plot bars for each model\n            for i, model in enumerate(models):\n                model_data = results_df[results_df['model_name'] == model]\n                plt.bar(x + i*width - width*len(models)/2 + width/2,\n                       model_data[metric],\n                       width,\n                       label=model,\n                       color=model_colors[i])\n            \n            plt.xlabel('Embedding Method')\n            plt.ylabel(metric_label)\n            plt.title(f'{metric_label} by Embedding Method and Model')\n            plt.xticks(x, embeddings, rotation=45)\n            plt.legend(bbox_to_anchor=(1.05, 1), loc='upper left')\n            plt.tight_layout()\n            plt.show()\n        \n        # 3. Model-centric view\n        for metric_mean, metric_std, metric_name in metrics:\n            plt.figure(figsize=(15, 6))\n            \n            # Create positions for bars\n            x = np.arange(len(models))\n            width = 0.8 / len(embeddings)\n            embedding_colors = sns.color_palette(\"husl\", n_colors=len(embeddings))\n            \n            # Plot bars for each embedding\n            for i, embedding in enumerate(embeddings):\n                embedding_data = results_df[results_df['embedding_name'] == embedding]\n                plt.bar(x + i*width - width*len(embeddings)/2 + width/2,\n                       embedding_data[metric_mean],\n                       width,\n                       label=embedding,\n                       color=embedding_colors[i],\n                       yerr=embedding_data[metric_std],\n                       capsize=3)\n            \n            plt.xlabel('Model Type')\n            plt.ylabel(metric_name)\n            plt.title(f'{metric_name} by Model Type and Embedding Method')\n            plt.xticks(x, models, rotation=45)\n            plt.legend(bbox_to_anchor=(1.05, 1), loc='upper left')\n            plt.tight_layout()\n            plt.show()\n    \n    def generate_report(self) -> str:\n        \"\"\"Generate a comprehensive benchmark report with detailed cross-analysis.\"\"\"\n        if not self.results:\n            return \"No benchmark results available\"\n            \n        results_df = pd.DataFrame([vars(r) for r in self.results])\n        \n        report = []\n        report.append(\"Comprehensive Benchmark Analysis Report\")\n        report.append(\"=\" * 50 + \"\\n\")\n        \n        # 1. Overall Best Performers\n        report.append(\"Overall Best Performers\")\n        report.append(\"-\" * 30)\n        \n        metrics = {\n            'MRRMSE': ('mrrmse_mean', 'mrrmse_std', 'min'),\n            'MSE': ('mse_mean', 'mse_std', 'min'),\n            'R²': ('r2_mean', 'r2_std', 'max'),\n            'Explained Variance Score (EVS)': ('evs_mean', 'evs_std', 'max'),\n        }\n        \n        for metric_name, (mean_col, std_col, best_func) in metrics.items():\n            report.append(f\"\\nBest {metric_name} Performance:\")\n            if best_func == 'min':\n                idx = results_df[mean_col].idxmin()\n            else:\n                idx = results_df[mean_col].idxmax()\n                \n            row = results_df.loc[idx]\n            report.append(f\"Model: {row['model_name']}\")\n            report.append(f\"Embedding: {row['embedding_name']}\")\n            report.append(f\"Score: {row[mean_col]:.4f} ± {row[std_col]:.4f}\")\n        \n        # 2. Detailed Cross-Analysis\n        report.append(\"\\nDetailed Cross-Analysis\")\n        report.append(\"-\" * 30)\n        \n        # For each model, analyze performance with different embeddings\n        for model in results_df['model_name'].unique():\n            report.append(f\"\\nModel: {model}\")\n            model_data = results_df[results_df['model_name'] == model]\n            \n            # Sort embeddings by MRRMSE performance\n            embedding_performance = model_data.sort_values('mrrmse_mean')\n            \n            report.append(\"\\nEmbedding Performance Ranking:\")\n            for idx, row in embedding_performance.iterrows():\n                report.append(f\"\\n{row['embedding_name']}:\")\n                report.append(f\"  MRRMSE: {row['mrrmse_mean']:.4f} ± {row['mrrmse_std']:.4f}\")\n                report.append(f\"  MSE: {row['mse_mean']:.4f} ± {row['mse_std']:.4f}\")\n                report.append(f\"  R²: {row['r2_mean']:.4f} ± {row['r2_std']:.4f}\")\n                report.append(f\"  Training Time: {row['training_time']:.2f}s\")\n                report.append(f\"  Memory Usage: {row['memory_usage']:.2f}MB\")\n        \n        # 3. Embedding Method Analysis\n        report.append(\"\\nEmbedding Method Analysis\")\n        report.append(\"-\" * 30)\n        \n        for embedding in results_df['embedding_name'].unique():\n            report.append(f\"\\nEmbedding: {embedding}\")\n            embedding_data = results_df[results_df['embedding_name'] == embedding]\n            \n            # Average performance across models\n            report.append(\"\\nAverage Performance:\")\n            report.append(f\"MRRMSE: {embedding_data['mrrmse_mean'].mean():.4f} ± {embedding_data['mrrmse_std'].mean():.4f}\")\n            report.append(f\"MSE: {embedding_data['mse_mean'].mean():.4f} ± {embedding_data['mse_std'].mean():.4f}\")\n            report.append(f\"R²: {embedding_data['r2_mean'].mean():.4f} ± {embedding_data['r2_std'].mean():.4f}\")\n            \n            # Best model for this embedding\n            best_idx = embedding_data['mrrmse_mean'].idxmin()\n            report.append(\"\\nBest Model Performance:\")\n            report.append(f\"Model: {embedding_data.loc[best_idx, 'model_name']}\")\n            report.append(f\"MRRMSE: {embedding_data.loc[best_idx, 'mrrmse_mean']:.4f} ± {embedding_data.loc[best_idx, 'mrrmse_std']:.4f}\")\n            report.append(f\"R²: {embedding_data.loc[best_idx, 'r2_mean']:.4f} ± {embedding_data.loc[best_idx, 'r2_std']:.4f}\")\n            \n            # Technical details\n            report.append(\"\\nTechnical Details:\")\n            report.append(f\"Embedding Size: {embedding_data['embedding_size'].iloc[0]}\")\n            report.append(f\"Average Memory Usage: {embedding_data['memory_usage'].mean():.2f}MB\")\n            report.append(f\"Average Training Time: {embedding_data['training_time'].mean():.2f}s\")\n        \n        return \"\\n\".join(report)\n\n\n    def get_combined_embeddings(\n        self,\n        df: pd.DataFrame,\n        fit: bool = True\n    ) -> np.ndarray:\n        \"\"\"\n        Get combined embeddings from all embedding methods.\n        Also creates a mapping of features to their source embeddings.\n        \n        Args:\n            df: Input DataFrame\n            fit: Whether to fit the embeddings or use pre-fitted\n            \n        Returns:\n            Combined embedding array\n        \"\"\"\n        all_embeddings = []\n        current_position = 0\n        \n        for name, method in self.embedding_methods.items():\n            # Get embeddings for this method\n            embedding = method.get_embedding(df, fit=fit)\n            all_embeddings.append(embedding)\n            \n            if fit:\n                # Store the range for this embedding\n                embedding_size = embedding.shape[1]\n                end_position = current_position + embedding_size\n                self.embedding_ranges[name] = (current_position, end_position)\n                \n                # Get feature names if available\n                try:\n                    feature_names = method.get_feature_names()\n                except (AttributeError, NotImplementedError):\n                    feature_names = [f\"{name}_feature_{i}\" for i in range(embedding_size)]\n                \n                # Map each feature to its source embedding\n                for i, feature in enumerate(feature_names):\n                    self.feature_map[current_position + i] = {\n                        'embedding': name,\n                        'feature': feature\n                    }\n                \n                current_position = end_position\n                self.combined_embedding_size = current_position\n        \n        return np.hstack(all_embeddings)\n\n    def analyze_combined_feature_importance(\n            self,\n            df: pd.DataFrame,\n            k: int = 10,\n            methods: List[str] = ['integrated_gradients', 'shapley']\n        ) -> Dict[str, Any]:\n            \"\"\"\n            Analyze feature importance using combined embeddings from all methods.\n            \n            Args:\n                df: Input DataFrame\n                k: Number of top features to analyze per embedding method\n                methods: List of interpretability methods to use\n                \n            Returns:\n                Dictionary containing feature importance analysis results\n            \"\"\"\n            import torch\n            from captum.attr import (\n                IntegratedGradients,\n                Saliency,\n                DeepLift,\n                ShapleyValueSampling\n            )\n            \n            # Prepare data\n            train_df, test_df, y_train, y_test = self.prepare_data(df)\n            \n            # Get combined embeddings\n            train_embeddings = self.get_combined_embeddings(train_df, fit=True)\n            test_embeddings = self.get_combined_embeddings(test_df, fit=False)\n            \n            results = {}\n            \n            for model_factory in self.model_factories:\n                print(f\"\\nAnalyzing with {model_factory.name}...\")\n                \n                # Create and check model\n                model = model_factory.create_model(self.combined_embedding_size)\n                # if not isinstance(model, torch.nn.Module):\n                #     print(f\"Skipping {model_factory.name} - not a PyTorch model\")\n                #     continue\n                    \n                # Convert data to PyTorch tensors\n                X_train = torch.FloatTensor(train_embeddings)\n                y_train = torch.FloatTensor(y_train)\n                X_test = torch.FloatTensor(test_embeddings)\n                \n                # Train model\n                model.fit(X_train, y_train)\n                model.model.eval()\n                \n                model_results = {\n                    'model_name': model_factory.name,\n                    'methods': {}\n                }\n                \n                # Create subplot grid based on number of methods\n                n_methods = len(methods)\n                fig = plt.figure(figsize=(15, 6 * n_methods))\n                gs = plt.GridSpec(n_methods, 1, height_ratios=[1] * n_methods)\n                \n                for idx, method in enumerate(methods):\n                    # Compute attributions\n                    if method == 'integrated_gradients':\n                        explainer = IntegratedGradients(model.model)\n                    elif method == 'shapley':\n                        explainer = ShapleyValueSampling(model.model)\n                    elif method == 'deeplift':\n                        explainer = DeepLift(model.model)\n                    elif method == 'saliency':\n                        explainer = Saliency(model.model)\n                    \n                    baseline = torch.zeros_like(X_test)\n                    attributions = explainer.attribute(X_test, baseline)\n                    attr_mean = attributions.mean(dim=0).abs().cpu().detach().numpy()\n                    \n                    # Process results per embedding method\n                    embedding_importance = {}\n                    for embedding_name, (start, end) in self.embedding_ranges.items():\n                        embedding_attr = attr_mean[start:end]\n                        top_k_idx = np.argsort(embedding_attr)[-k:][::-1]\n                        \n                        embedding_importance[embedding_name] = {\n                            'features': [self.feature_map[start + i]['feature'] for i in top_k_idx],\n                            'importance': embedding_attr[top_k_idx].tolist(),\n                            'all_importance': embedding_attr.tolist()\n                        }\n                    \n                    model_results['methods'][method] = embedding_importance\n                    \n                    # Create visualization\n                    ax = plt.subplot(gs[idx])\n                    \n                    # Prepare data for plotting\n                    plot_data = []\n                    for emb_name, importance in embedding_importance.items():\n                        for feat, imp in zip(importance['features'], importance['importance']):\n                            plot_data.append({\n                                'Embedding': emb_name,\n                                'Feature': feat,\n                                'Importance': imp\n                            })\n                    \n                    plot_df = pd.DataFrame(plot_data)\n                    \n                    # Create grouped bar plot\n                    sns.barplot(\n                        data=plot_df,\n                        x='Importance',\n                        y='Feature',\n                        hue='Embedding',\n                        palette='Set2',\n                        ax=ax\n                    )\n                    \n                    ax.set_title(f'{method.replace(\"_\", \" \").title()} - Feature Importance by Embedding')\n                    ax.set_xlabel('Absolute Importance')\n                    \n                    # Adjust legend\n                    ax.legend(title='Embedding Type', bbox_to_anchor=(1.05, 1), loc='upper left')\n                    \n                plt.tight_layout()\n                plt.show()\n                \n                results[model_factory.name] = model_results\n            \n            return results","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-20T17:59:40.04856Z","iopub.execute_input":"2025-02-20T17:59:40.048901Z","iopub.status.idle":"2025-02-20T17:59:40.094081Z","shell.execute_reply.started":"2025-02-20T17:59:40.048875Z","shell.execute_reply":"2025-02-20T17:59:40.093199Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **ANOVA Benchmarker**","metadata":{}},{"cell_type":"code","source":"from scipy import stats\nfrom statsmodels.stats.multicomp import pairwise_tukeyhsd\nfrom tqdm.auto import tqdm\n\nclass EnhancedEmbeddingBenchmark(EmbeddingBenchmark):\n    def __init__(self, *args, **kwargs):\n        super().__init__(*args, **kwargs)\n        self.run_results = []  # Store individual run results\n        \n    def evaluate(\n        self,\n        model_factory: ModelFactory,\n        embedding_name: str,\n        embedding_method: BaseEmbedding,\n        train_df: pd.DataFrame,\n        test_df: pd.DataFrame,\n        y_train: np.ndarray,\n        y_test: np.ndarray\n    ) -> Tuple[BenchmarkResult, List[Dict]]:\n        \"\"\"Extended evaluate method to store individual run results\"\"\"\n        individual_runs = []\n        \n        # Compute embeddings once\n        with tqdm(total=2, desc=f\"Computing {embedding_name} embeddings\", leave=False) as pbar:\n            start_time = time.time()\n            train_embeddings = embedding_method.get_embedding(train_df, fit=True)\n            pbar.update(1)\n            test_embeddings = embedding_method.get_embedding(test_df, fit=False)\n            pbar.update(1)\n            embedding_time = time.time() - start_time\n        \n        # Create model with correct input size\n        model = model_factory.create_model(embedding_method.embedding_size)\n\n        # Progress bar for runs\n        run_pbar = tqdm(\n            range(self.n_runs), \n            desc=f\"{model_factory.name} with {embedding_name}\",\n            leave=False\n        )\n        \n        for run in range(self.n_runs):\n            print(f\"Run {run + 1}/{self.n_runs}\")\n            \n            # Train model\n            start_time = time.time()\n            model.fit(train_embeddings, y_train)\n            training_time = time.time() - start_time\n            \n            # Evaluate\n            start_time = time.time()\n            y_pred = model.predict(test_embeddings)\n            inference_time = time.time() - start_time\n            \n            # Compute metrics\n            mse = mean_squared_error(y_test, y_pred)\n            mrrmse = RMSE_rowwise_loss_numpy(y_pred, y_test)\n            r2 = r2_score(y_test, y_pred)\n            evs = evaluate_evs_score(y_test, y_pred)\n            \n            run_result = {\n                'run': run,\n                'model_name': model_factory.name,\n                'embedding_name': embedding_name,\n                'mse': mse,\n                'mrrmse': mrrmse,\n                'r2': r2,\n                'evs': evs,\n                'training_time': training_time,\n                'inference_time': inference_time,\n                'memory_usage': (train_embeddings.nbytes + test_embeddings.nbytes) / (1024 * 1024)\n            }\n\n            run_pbar.set_postfix({\n                'MSE': f\"{mse:.4f}\",\n                'R²': f\"{r2:.4f}\",\n                'MRRMSE': f\"{mrrmse:.4f}\",\n                'EVS': f\"{evs:.4f}\"\n            })\n            \n            individual_runs.append(run_result)\n        \n        # Calculate average metrics for backward compatibility\n        avg_result = BenchmarkResult(\n            model_name=model_factory.name,\n            embedding_name=embedding_name,\n            mse_mean=np.mean([r['mse'] for r in individual_runs]),\n            mse_std=np.std([r['mse'] for r in individual_runs]),\n            mrrmse_mean=np.mean([r['mrrmse'] for r in individual_runs]),\n            mrrmse_std=np.std([r['mrrmse'] for r in individual_runs]),\n            r2_mean=np.mean([r['r2'] for r in individual_runs]),\n            r2_std=np.std([r['r2'] for r in individual_runs]),\n            evs_mean=np.mean([r['evs'] for r in individual_runs]),\n            evs_std=np.std([r['evs'] for r in individual_runs]),\n            training_time=np.mean([r['training_time'] for r in individual_runs]),\n            inference_time=np.mean([r['inference_time'] for r in individual_runs]),\n            embedding_size=embedding_method.embedding_size,\n            memory_usage=individual_runs[0]['memory_usage']\n        )\n        print(\"Run Compelelted !\")\n        return avg_result, individual_runs\n\n    def run_benchmark(self, df: pd.DataFrame) -> pd.DataFrame:\n        \"\"\"Enhanced benchmark runner that preserves individual run results\"\"\"\n        df = df.copy()\n        train_df, test_df, y_train, y_test = self.prepare_data(df)\n        \n        all_individual_runs = []\n        aggregated_results = []\n        \n        for name, method in self.embedding_methods.items():\n            print(f\"\\nEvaluating {name} embedding...\")\n            \n            for model_factory in self.model_factories:\n                print(f\"Testing with {model_factory.name}...\")\n                try:\n                    avg_result, individual_runs = self.evaluate(\n                        model_factory, name, method,\n                        train_df, test_df,\n                        y_train, y_test\n                    )\n                    aggregated_results.append(avg_result)\n                    all_individual_runs.extend(individual_runs)\n                except Exception as e:\n                    print(f\"Error evaluating {name} embedding with {model_factory.name}: {str(e)}\")\n        \n        self.results = aggregated_results\n        self.run_results = all_individual_runs\n        \n        # Perform statistical analysis\n        print(\"\\nPerforming Statistical Analysis...\")\n        self._analyze_results()\n        \n        return pd.DataFrame([vars(r) for r in self.results])\n\n    def _analyze_results(self):\n        \"\"\"Perform statistical analysis on the run results\"\"\"\n        if not self.run_results:\n            print(\"No results available for analysis\")\n            return\n        \n        runs_df = pd.DataFrame(self.run_results)\n        metrics = ['mrrmse', 'mse', 'r2', 'evs']\n        \n        for metric in metrics:\n            print(f\"\\nAnalyzing {metric}:\")\n            \n            # Perform one-way ANOVA\n            embeddings = runs_df['embedding_name'].unique()\n            embedding_groups = [\n                runs_df[runs_df['embedding_name'] == emb][metric].values \n                for emb in embeddings\n            ]\n            \n            f_statistic, p_value = stats.f_oneway(*embedding_groups)\n            \n            print(f\"ANOVA Results:\")\n            print(f\"F-statistic: {f_statistic:.4f}\")\n            print(f\"p-value: {p_value:.4f}\")\n            \n            # Visualize distribution of results\n            plt.figure(figsize=(12, 6))\n            sns.boxplot(data=runs_df, x='embedding_name', y=metric)\n            plt.title(f'Distribution of {metric} by Embedding Method\\nANOVA p-value: {p_value:.4f}')\n            plt.xticks(rotation=45)\n            plt.tight_layout()\n            plt.show()\n            \n            # Perform Tukey's HSD test\n            try:\n                metrics_df = runs_df.copy()\n                metrics_df[metric] = pd.to_numeric(metrics_df[metric], errors='coerce')\n                metrics_df = metrics_df.dropna(subset=[metric])\n                tukey = pairwise_tukeyhsd(\n                    endog=metrics_df[metric],  \n                    groups=metrics_df['embedding_name'],  \n                    alpha=0.05\n                )\n\n                \n                # Print Tukey's test summary\n                print(\"\\nTukey HSD Test Results:\")\n                print(tukey)\n            \n                # Convert summary to DataFrame for better visualization\n                tukey_df = pd.DataFrame(data=tukey.summary().data[1:], columns=tukey.summary().data[0])\n                print(\"\\nFormatted Tukey HSD Test Results:\")\n                print(tukey_df)\n\n            except Exception as e:\n                print(f\"Could not perform Tukey's HSD test: {str(e)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-20T17:59:40.095171Z","iopub.execute_input":"2025-02-20T17:59:40.095626Z","iopub.status.idle":"2025-02-20T17:59:40.197218Z","shell.execute_reply.started":"2025-02-20T17:59:40.095599Z","shell.execute_reply":"2025-02-20T17:59:40.196089Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **Configurations**","metadata":{}},{"cell_type":"markdown","source":"## **Embedding Configs**","metadata":{}},{"cell_type":"code","source":"embeddings = {\n    \"OneHot\": OneHotEmbedding(columns=['cell_type', 'sm_name']),\n    \"SMILES\": SMILESEmbedding(),\n    \"Target\": TargetEmbedding(emb_size=256),\n    \"MorganFP\": MorganFingerPrintEmbedding(hidden_size=128),\n    \"Mol2Vec\": Mol2VecEmbedding(model_path=f\"/kaggle/input/mol2vec/pytorch/default/1/model_300dim.pkl\"),\n    \"Chemberta\": ChemBERTaEmbedding(\n        model_name=\"DeepChem/ChemBERTa-77M-MTR\",\n        embedding_type='mean_pooling',\n    ),\n    \"Molformer\": MolformerEmbedding(\n        columns='SMILES',\n        padding=True\n    ),\n    \"Smolebart\": SmoleBartEmbedding(\n        columns='SMILES',\n        embedder='encoder',\n        padding=True\n    ),\n    \"Selfformer\": SelfFormerEmbedding(\n        columns='SMILES',\n        model_name='/kaggle/input/selfformer/transformers/default/1/SELFormer',\n        batch_size=32,\n        nb_workers=5\n    )\n}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-20T17:59:40.201485Z","iopub.execute_input":"2025-02-20T17:59:40.201768Z","iopub.status.idle":"2025-02-20T18:00:00.352016Z","shell.execute_reply.started":"2025-02-20T17:59:40.201746Z","shell.execute_reply":"2025-02-20T18:00:00.350889Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **Model Configs**","metadata":{}},{"cell_type":"code","source":"def get_vanilla_sklearn():\n    return [\n        # LinearRegressionFactory(),\n        DecisionTreeFactory(\n            max_depth=7,\n            min_samples_split=10\n        ),\n        KNeighborsFactory(\n            n_neighbors=7,\n            weights='uniform'\n        ),\n        MLPRegressorFactory(\n            hidden_layer_sizes=(64, 32),\n            alpha=0.01,\n            max_iter=1000\n        )\n    ]\n\ndef get_tsvd_sklearn():\n    return [\n        # LinearRegressionTSVDFactory(n_components=50),\n        DecisionTreeTSVDFactory(\n            n_components=50,\n            max_depth=7,\n            min_samples_split=10\n        ),\n        KNeighborsTSVDFactory(\n            n_components=50,\n            n_neighbors=7,\n            weights='uniform'\n        ),\n        MLPRegressorTSVDFactory(\n            n_components=50,\n            hidden_layer_sizes=(64, 32),\n            alpha=0.01,\n            max_iter=1000\n        )\n    ]\n\ndef get_lstm(output_size: int = 18211):\n    return PyTorchModelFactory(\n        model_class=LSTMPredictionModel,\n        model_params={'hidden_size': 128, 'output_size': output_size},\n        num_epochs=50,\n        early_stopping=False,\n        criterion=RMSE_rowwise_loss,\n        optimizer_class=torch.optim.Adam,\n        optimizer_params={'lr': 1e-3},\n        name=\"LSTM\"\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-20T18:01:10.442323Z","iopub.execute_input":"2025-02-20T18:01:10.442657Z","iopub.status.idle":"2025-02-20T18:01:10.449497Z","shell.execute_reply.started":"2025-02-20T18:01:10.442634Z","shell.execute_reply":"2025-02-20T18:01:10.448501Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **Experiments with Entire Target**","metadata":{}},{"cell_type":"markdown","source":"## **Vanilla Embeddings with TSVD**","metadata":{}},{"cell_type":"code","source":"model_factories = get_tsvd_sklearn()\nmodel_factories.append(get_lstm())\n\n# benchmark = EmbeddingBenchmark(embeddings, model_factories = [lstm_factory],  n_runs = 1)\nbenchmark = EnhancedEmbeddingBenchmark(embeddings, model_factories = model_factories, n_runs = 10)\n\n# Run benchmark\nresults_df = benchmark.run_benchmark(df)\n\n# Visualize results\nbenchmark.visualize_results()\n\n# Generate detailed report\nreport = benchmark.generate_report()\nprint(report)\nprint(\"Completed Vanilla Embeddings with TSVD (Entire Target)!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-20T18:01:13.493778Z","iopub.execute_input":"2025-02-20T18:01:13.494079Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **Embeddings With Target as Base**","metadata":{}},{"cell_type":"code","source":"combined_embedding = {}\n\nfor name, emb in embeddings.items():\n    # Skip combining Target with itself\n    if name == \"Target\":\n        continue\n    \n    # Create combined embedding (Target + current embedding)\n    combined_name = f\"Target+{name}\"\n    combined_embedding[combined_name] = MultiEmbedding(\n        embeddings=[embeddings[\"Target\"], emb],\n        name=combined_name,\n        combination_method=\"concat\"\n    )\n    \n    # Compute and store the combined embeddings\n    # print(f\"Computing combined embeddings for {combined_name}...\")\n    # combinations[combined_name].compute_and_store_embeddings(df, entity_column, batch_size)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-20T18:00:55.263503Z","iopub.status.idle":"2025-02-20T18:00:55.263801Z","shell.execute_reply":"2025-02-20T18:00:55.263687Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model_factories = get_tsvd_sklearn()\nmodel_factories.append(get_lstm())\n\n# Create benchmark instance\nbenchmark = EnhancedEmbeddingBenchmark(combined_embedding, model_factories = model_factories,  n_runs = 10)\n\n# Run benchmark\nresults_df = benchmark.run_benchmark(df)\n\n# Visualize results\nbenchmark.visualize_results()\n\n# Generate detailed report\nreport = benchmark.generate_report()\nprint(report)\nprint(\"Completed Embeddings With Target as Base (Entire Target) !\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-20T18:00:55.26467Z","iopub.status.idle":"2025-02-20T18:00:55.265036Z","shell.execute_reply":"2025-02-20T18:00:55.264877Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **Experiments with Trimmed Targets**","metadata":{}},{"cell_type":"markdown","source":"## **Vanilla Embeddings with TSVD**","metadata":{}},{"cell_type":"code","source":"model_factories = get_vanilla_sklearn()\nmodel_factories.append(get_lstm(256))\n\n# benchmark = EmbeddingBenchmark(embeddings, model_factories = [lstm_factory],  n_runs = 1)\nbenchmark = EnhancedEmbeddingBenchmark(embeddings, model_factories = model_factories, n_runs = 10, only_significant = True)\n\n# Run benchmark\nresults_df = benchmark.run_benchmark(df)\n\n# Visualize results\nbenchmark.visualize_results()\n\n# Generate detailed report\nreport = benchmark.generate_report()\nprint(report)\nprint(\"Completed Vanilla Embeddings with TSVD (Trimmed Targets) !\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-20T18:00:55.265789Z","iopub.status.idle":"2025-02-20T18:00:55.266187Z","shell.execute_reply":"2025-02-20T18:00:55.265999Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **Embeddings With Target as Base**","metadata":{}},{"cell_type":"code","source":"model_factories = get_vanilla_sklearn()\nmodel_factories.append(get_lstm(256))\n\n# Create benchmark instance\nbenchmark = EnhancedEmbeddingBenchmark(combined_embedding, model_factories = model_factories,  n_runs = 10, only_significant = True)\n\n# Run benchmark\nresults_df = benchmark.run_benchmark(df)\n\n# Visualize results\nbenchmark.visualize_results()\n\n# Generate detailed report\nreport = benchmark.generate_report()\nprint(report)\nprint(\"Completed Embeddings With Target as Base (Trimmed Targets) !\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-20T18:00:55.266905Z","iopub.status.idle":"2025-02-20T18:00:55.267302Z","shell.execute_reply":"2025-02-20T18:00:55.267111Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **Feature Importance**","metadata":{}},{"cell_type":"code","source":"class Combined_Benchmarker():\n    def __init__(\n        self,\n        embedding_methods: Dict[str, BaseEmbedding],\n        model_factories: List[ModelFactory],\n        n_runs: int = 5\n    ):\n        \"\"\"\n        Initialize with dictionary of embedding methods and model factories.\n        Also creates a combined embedding handler.\n        \"\"\"\n        self.embedding_methods = embedding_methods\n        self.model_factories = model_factories\n        self.n_runs = n_runs\n        self.results: List[BenchmarkResult] = []\n        \n        # Add combined embedding functionality\n        self.combined_embedding_size = 0\n        self.embedding_ranges = {}\n        self.feature_map = {}\n        \n    def get_combined_embeddings(\n        self,\n        df: pd.DataFrame,\n        fit: bool = True\n    ) -> np.ndarray:\n        \"\"\"\n        Get combined embeddings from all embedding methods.\n        Also creates a mapping of features to their source embeddings.\n        \n        Args:\n            df: Input DataFrame\n            fit: Whether to fit the embeddings or use pre-fitted\n            \n        Returns:\n            Combined embedding array\n        \"\"\"\n        all_embeddings = []\n        current_position = 0\n        \n        for name, method in self.embedding_methods.items():\n            # Get embeddings for this method\n            embedding = method.get_embedding(df, fit=fit)\n            all_embeddings.append(embedding)\n            \n            if fit:\n                # Store the range for this embedding\n                embedding_size = embedding.shape[1]\n                end_position = current_position + embedding_size\n                self.embedding_ranges[name] = (current_position, end_position)\n                \n                # Get feature names if available\n                try:\n                    feature_names = method.get_feature_names()\n                except (AttributeError, NotImplementedError):\n                    feature_names = [f\"{name}_feature_{i}\" for i in range(embedding_size)]\n                \n                # Map each feature to its source embedding\n                for i, feature in enumerate(feature_names):\n                    self.feature_map[current_position + i] = {\n                        'embedding': name,\n                        'feature': feature\n                    }\n                \n                current_position = end_position\n                self.combined_embedding_size = current_position\n        \n        return np.hstack(all_embeddings)\n    \n    def analyze_combined_feature_importance(\n        self,\n        df: pd.DataFrame,\n        k: int = 10,\n        methods: List[str] = ['integrated_gradients', 'shapley']\n    ) -> Dict[str, Any]:\n        \"\"\"\n        Analyze feature importance using combined embeddings from all methods.\n        \n        Args:\n            df: Input DataFrame\n            k: Number of top features to analyze per embedding method\n            methods: List of interpretability methods to use\n            \n        Returns:\n            Dictionary containing feature importance analysis results\n        \"\"\"\n        import torch\n        from captum.attr import (\n            IntegratedGradients,\n            Saliency,\n            DeepLift,\n            ShapleyValueSampling\n        )\n        \n        # Prepare data\n        train_df, test_df, y_train, y_test = self.prepare_data(df)\n        \n        # Get combined embeddings\n        train_embeddings = self.get_combined_embeddings(train_df, fit=True)\n        test_embeddings = self.get_combined_embeddings(test_df, fit=False)\n        \n        results = {}\n        \n        for model_factory in self.model_factories:\n            print(f\"\\nAnalyzing with {model_factory.name}...\")\n            \n            # Create and check model\n            model = model_factory.create_model(self.combined_embedding_size)\n            if not isinstance(model, torch.nn.Module):\n                print(f\"Skipping {model_factory.name} - not a PyTorch model\")\n                continue\n                \n            # Convert data to PyTorch tensors\n            X_train = torch.FloatTensor(train_embeddings)\n            y_train = torch.FloatTensor(y_train)\n            X_test = torch.FloatTensor(test_embeddings)\n            \n            # Train model\n            model.fit(X_train, y_train)\n            model.eval()\n            \n            model_results = {\n                'model_name': model_factory.name,\n                'methods': {}\n            }\n            \n            # Create subplot grid based on number of methods\n            n_methods = len(methods)\n            fig = plt.figure(figsize=(15, 6 * n_methods))\n            gs = plt.GridSpec(n_methods, 1, height_ratios=[1] * n_methods)\n            \n            for idx, method in enumerate(methods):\n                # Compute attributions\n                if method == 'integrated_gradients':\n                    explainer = IntegratedGradients(model)\n                elif method == 'shapley':\n                    explainer = ShapleyValueSampling(model)\n                elif method == 'deeplift':\n                    explainer = DeepLift(model)\n                elif method == 'saliency':\n                    explainer = Saliency(model)\n                \n                baseline = torch.zeros_like(X_test)\n                attributions = explainer.attribute(X_test, baseline)\n                attr_mean = attributions.mean(dim=0).abs().cpu().detach().numpy()\n                \n                # Process results per embedding method\n                embedding_importance = {}\n                for embedding_name, (start, end) in self.embedding_ranges.items():\n                    embedding_attr = attr_mean[start:end]\n                    top_k_idx = np.argsort(embedding_attr)[-k:][::-1]\n                    \n                    embedding_importance[embedding_name] = {\n                        'features': [self.feature_map[start + i]['feature'] for i in top_k_idx],\n                        'importance': embedding_attr[top_k_idx].tolist(),\n                        'all_importance': embedding_attr.tolist()\n                    }\n                \n                model_results['methods'][method] = embedding_importance\n                \n                # Create visualization\n                ax = plt.subplot(gs[idx])\n                \n                # Prepare data for plotting\n                plot_data = []\n                for emb_name, importance in embedding_importance.items():\n                    for feat, imp in zip(importance['features'], importance['importance']):\n                        plot_data.append({\n                            'Embedding': emb_name,\n                            'Feature': feat,\n                            'Importance': imp\n                        })\n                \n                plot_df = pd.DataFrame(plot_data)\n                \n                # Create grouped bar plot\n                sns.barplot(\n                    data=plot_df,\n                    x='Importance',\n                    y='Feature',\n                    hue='Embedding',\n                    palette='Set2',\n                    ax=ax\n                )\n                \n                ax.set_title(f'{method.replace(\"_\", \" \").title()} - Feature Importance by Embedding')\n                ax.set_xlabel('Absolute Importance')\n                \n                # Adjust legend\n                ax.legend(title='Embedding Type', bbox_to_anchor=(1.05, 1), loc='upper left')\n                \n            plt.tight_layout()\n            plt.show()\n            \n            results[model_factory.name] = model_results\n        \n        return results","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-20T18:00:55.26807Z","iopub.status.idle":"2025-02-20T18:00:55.268363Z","shell.execute_reply":"2025-02-20T18:00:55.268259Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"lstm_factory.device = 'cpu'\ncb = EmbeddingBenchmark(embeddings, model_factories = [lstm_factory],  n_runs = 1)\n# results = cb.analyze_combined_feature_importance(df)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-20T18:00:55.268872Z","iopub.status.idle":"2025-02-20T18:00:55.269108Z","shell.execute_reply":"2025-02-20T18:00:55.269007Z"}},"outputs":[],"execution_count":null}]}