{"cells":[{"cell_type":"code","execution_count":null,"id":"34992ce7","metadata":{"language":"python"},"outputs":[],"source":"from pathlib import Path\n\nfrom IPython.display import Image, Markdown, display\n\n\n\nSTATIC_IMAGE_DIR = Path(\"/kaggle/input/notebook-static-images-public\")\n\ndisplay(Markdown(\"## Saved analysis figures\"))\n\nfor image_path in sorted(STATIC_IMAGE_DIR.glob(\"perturbations-*.png\")):\n\n    display(Markdown(f\"**{image_path.stem}**\"))\n\n    display(Image(filename=str(image_path)))"},{"cell_type":"code","execution_count":14,"id":"90b9fe9c-dcf1-4d10-a215-939f94f77a56","metadata":{},"outputs":[],"source":"import pandas as pd\nimport anndata as ad\nimport scipy.sparse as sp\n\nbase = \"/home/Data_Drive_8TB/kykim/7. Kaggle/Single-Cell_Perturbations/baseline\"\n\nobs = pd.read_csv(f\"{base}/multiome_obs_meta.csv\")\nvar = pd.read_csv(f\"{base}/multiome_var_meta.csv\")\nX = pd.read_parquet(f\"{base}/multiome_train.parquet\")\n\nobs.index = obs.index.astype(str)\n\ncell_map = pd.Series(\n    range(len(obs)),\n    index=obs[\"obs_id\"]\n)\n\nvar.index = var[\"location\"].astype(str)\n\nfeature_map = pd.Series(\n    range(len(var)),\n    index=var[\"location\"]\n)\n\nrows = cell_map.loc[X[\"obs_id\"]].values\ncols = feature_map.loc[X[\"location\"]].values\nvalues = X[\"normalized_count\"].values\n\nmatrix = sp.coo_matrix(\n    (values, (rows, cols)),\n    shape=(len(obs), len(var))\n).tocsr()\n\nadata = ad.AnnData(\n    X=matrix,\n    obs=obs,\n    var=var\n)\n\nprint(adata)"},{"cell_type":"code","execution_count":17,"id":"cf5ef1d2-dbb9-4a76-af7f-ba952716a9e2","metadata":{},"outputs":[],"source":"save_path = \"/home/Data_Drive_8TB/kykim/7. Kaggle/Single-Cell_Perturbations/baseline/multiome_adata.h5ad\"\n\nadata.write_h5ad(\n    save_path,\n    compression=\"gzip\")"},{"cell_type":"code","execution_count":19,"id":"c0bc1548-e684-426e-81ad-93c47fe899de","metadata":{},"outputs":[],"source":"adata.obs"},{"cell_type":"code","execution_count":23,"id":"6e81bf8d-a9dc-46be-8cdc-b48819216bae","metadata":{},"outputs":[],"source":"import pandas as pd\n\nde_train = pd.read_parquet(\"/home/Data_Drive_8TB/kykim/7. Kaggle/Single-Cell_Perturbations/baseline/de_train.parquet\")"},{"cell_type":"code","execution_count":27,"id":"8010de02-90b8-463e-a772-bd77247a9c11","metadata":{},"outputs":[],"source":"de_train[de_train['cell_type'] == 'B cells']"},{"cell_type":"code","execution_count":28,"id":"86b29aab-dfcf-4be6-9894-845dd4c167c2","metadata":{},"outputs":[],"source":"import pandas as pd\nimport pyranges as pr\n\n# Build peak_matrix and peak_annotation for downstream ATAC regulatory mapping.\nif \"adata\" not in globals():\n    raise NameError(\"Run cell 1 first to create adata.\")\n\npeak_var = adata.var[adata.var[\"feature_type\"] != \"Gene Expression\"].copy()\nif peak_var.empty:\n    raise ValueError(\"No ATAC peak features found in adata.var['feature_type'].\")\n\npeak_var[\"location\"] = peak_var[\"location\"].astype(str)\n\n# Parse genomic interval from location (e.g., chr1:100-200).\ncoord = peak_var[\"location\"].str.extract(r\"^([^:]+):(\\d+)-(\\d+)$\")\ncoord.columns = [\"chr\", \"start\", \"end\"]\npeak_var[\"chr\"] = coord[\"chr\"]\npeak_var[\"start\"] = pd.to_numeric(coord[\"start\"], errors=\"coerce\")\npeak_var[\"end\"] = pd.to_numeric(coord[\"end\"], errors=\"coerce\")\npeak_var = peak_var.dropna(subset=[\"chr\", \"start\", \"end\"]).copy()\npeak_var[\"start\"] = peak_var[\"start\"].astype(int)\npeak_var[\"end\"] = peak_var[\"end\"].astype(int)\n\npromoter_bed = \"/home/Data_Drive_8TB/kykim/7. Kaggle/Single-Cell_Perturbations/baseline/promoter.bed\"\nenhancer_bed = \"/home/Data_Drive_8TB/kykim/7. Kaggle/Single-Cell_Perturbations/baseline/enhancer.bed\"\n\npeak_var[\"peak_id\"] = peak_var.index.astype(str)\npeak_var[\"region_type\"] = \"other\"\n\npeak_gr = pr.PyRanges(\n    peak_var[[\"chr\", \"start\", \"end\", \"peak_id\"]].rename(\n        columns={\"chr\": \"Chromosome\", \"start\": \"Start\", \"end\": \"End\", \"peak_id\": \"PeakID\"}\n    )\n)\n\nprom_df = pd.read_csv(\n    promoter_bed,\n    sep=\"\\t\",\n    header=None,\n    usecols=[0, 1, 2],\n    names=[\"Chromosome\", \"Start\", \"End\"]\n)\nenh_df = pd.read_csv(\n    enhancer_bed,\n    sep=\"\\t\",\n    header=None,\n    usecols=[0, 1, 2],\n    names=[\"Chromosome\", \"Start\", \"End\"]\n)\n\nprom_hit = peak_gr.join(pr.PyRanges(prom_df)).df\nenh_hit = peak_gr.join(pr.PyRanges(enh_df)).df\n\nif not enh_hit.empty:\n    enh_ids = enh_hit[\"PeakID\"].astype(str).unique()\n    peak_var.loc[peak_var[\"peak_id\"].isin(enh_ids), \"region_type\"] = \"enhancer\"\nif not prom_hit.empty:\n    prom_ids = prom_hit[\"PeakID\"].astype(str).unique()\n    peak_var.loc[peak_var[\"peak_id\"].isin(prom_ids), \"region_type\"] = \"promoter\"\n\npeak_annotation = peak_var.copy()\npeak_locations = peak_annotation[\"location\"].tolist()\n\npeak_matrix = pd.DataFrame.sparse.from_spmatrix(\n    adata[:, peak_locations].X,\n    index=adata.obs_names,\n    columns=peak_locations\n)\n\nprint(f\"peak_matrix: {peak_matrix.shape}\")\nprint(\"peak_annotation region_type counts:\")\nprint(peak_annotation[\"region_type\"].value_counts())"},{"cell_type":"markdown","id":"85af61ab-dc85-4de2-83de-5ba259251b7e","metadata":{},"source":"# Multiome - ATAC"},{"cell_type":"code","execution_count":38,"id":"f9e99fc5-c087-4e0d-9278-ea46b389be39","metadata":{},"outputs":[],"source":"import pandas as pd\nimport pyranges as pr\n\n\ndef peak_to_gene_regulatory_matrix(\n    peak_matrix,\n    peak_annotation,\n    gtf_file\n):\n\n    gtf = pr.read_gtf(gtf_file)\n\n    genes = (\n        gtf[gtf.Feature == \"gene\"]\n        .df[[\"Chromosome\", \"Start\", \"End\", \"gene_name\"]]\n        .copy()\n    )\n\n    peak = peak_annotation[\n        [\"chr\", \"start\", \"end\", \"region_type\"]\n    ].copy()\n\n    peak[\"Chromosome\"] = peak[\"chr\"].astype(str)\n    peak[\"Start\"] = peak[\"start\"].astype(int)\n    peak[\"End\"] = peak[\"end\"].astype(int)\n\n    peak_gr = pr.PyRanges(\n        peak[[\"Chromosome\", \"Start\", \"End\"]]\n    )\n\n    gene_gr = pr.PyRanges(genes)\n\n    nearest = peak_gr.nearest(gene_gr).df\n\n    nearest[\"peak_key\"] = (\n        nearest[\"Chromosome\"].astype(str)\n        + \":\"\n        + nearest[\"Start\"].astype(str)\n        + \"-\"\n        + nearest[\"End\"].astype(str)\n    )\n\n    peak[\"peak_key\"] = (\n        peak[\"Chromosome\"]\n        + \":\"\n        + peak[\"Start\"].astype(str)\n        + \"-\"\n        + peak[\"End\"].astype(str)\n    )\n\n    peak_gene = peak.merge(\n        nearest[[\"peak_key\", \"gene_name\"]],\n        on=\"peak_key\",\n        how=\"left\"\n    )\n\n    peak_gene = peak_gene.dropna(\n        subset=[\"gene_name\"]\n    )\n\n    peak_gene[\"feature\"] = (\n        peak_gene[\"gene_name\"].astype(str)\n        + \"_\"\n        + peak_gene[\"region_type\"].astype(str)\n    )\n\n    peak_to_feature = dict(\n        zip(\n            peak_gene[\"peak_key\"],\n            peak_gene[\"feature\"]\n        )\n    )\n\n    features = [\n        peak_to_feature.get(\n            col.split(\"_\", 1)[0]\n        )\n        for col in peak_matrix.columns\n    ]\n\n    valid = [\n        i\n        for i, feature in enumerate(features)\n        if feature is not None\n    ]\n\n    matrix = peak_matrix.iloc[:, valid].copy()\n\n    matrix.columns = [\n        features[i]\n        for i in valid\n    ]\n\n    gene_reg_matrix = (\n        matrix.T\n        .groupby(level=0)\n        .sum()\n        .T\n    )\n\n    return gene_reg_matrix, peak_gene"},{"cell_type":"code","execution_count":44,"id":"c4ddbe92-4b17-4dc6-9277-aa2832b48297","metadata":{},"outputs":[],"source":"import pandas as pd\nimport pyranges as pr\n\ngtf_file = \"/home/Data_Drive_8TB/kykim/7. Kaggle/Single-Cell_Perturbations/baseline/gencode.v44.annotation.gtf.gz\"\n\ngtf = pr.read_gtf(gtf_file)\n\ngenes = (\n    gtf[gtf.Feature == \"gene\"]\n    .df[[\"Chromosome\", \"Start\", \"End\", \"gene_name\"]]\n    .copy()\n)\n\npeak = peak_annotation[\n    [\"chr\", \"start\", \"end\", \"region_type\"]\n].copy()\n\npeak[\"Chromosome\"] = peak[\"chr\"].astype(str)\npeak[\"Start\"] = peak[\"start\"].astype(int)\npeak[\"End\"] = peak[\"end\"].astype(int)\n\npeak_gr = pr.PyRanges(\n    peak[[\"Chromosome\", \"Start\", \"End\"]]\n)\n\ngene_gr = pr.PyRanges(genes)\n\nnearest = peak_gr.nearest(gene_gr).df"},{"cell_type":"code","execution_count":45,"id":"d42b181a-b437-493b-a673-dc65eb045984","metadata":{},"outputs":[],"source":"nearest[\"peak_key\"] = (\n    nearest[\"Chromosome\"].astype(str)\n    + \":\"\n    + nearest[\"Start\"].astype(str)\n    + \"-\"\n    + nearest[\"End\"].astype(str)\n)\n\npeak[\"peak_key\"] = (\n    peak[\"Chromosome\"]\n    + \":\"\n    + peak[\"Start\"].astype(str)\n    + \"-\"\n    + peak[\"End\"].astype(str)\n)\n\npeak_gene = peak.merge(\n    nearest[[\"peak_key\", \"gene_name\"]],\n    on=\"peak_key\",\n    how=\"left\"\n)\n\npeak_gene = peak_gene.dropna(\n    subset=[\"gene_name\"]\n)\n\npeak_gene[\"feature\"] = (\n    peak_gene[\"gene_name\"].astype(str)\n    + \"_\"\n    + peak_gene[\"region_type\"].astype(str)\n)"},{"cell_type":"code","execution_count":49,"id":"053afaf0-6c90-447b-a4ea-a29441023ec9","metadata":{},"outputs":[],"source":"peak_to_feature = dict(\n    zip(\n        peak_gene[\"peak_key\"],\n        peak_gene[\"feature\"]\n    )\n)\n\nfeatures = [\n    peak_to_feature.get(\n        col.split(\"_\", 1)[0]\n    )\n    for col in peak_matrix.columns\n]\n\nvalid = [\n    i\n    for i, x in enumerate(features)\n    if x is not None\n]\n\nmatrix = peak_matrix.iloc[:, valid].copy()\n\nmatrix.columns = [\n    features[i]\n    for i in valid\n]\n\n###\n\nfrom scipy.sparse import csr_matrix\n\nfeature_names = pd.Index(\n    pd.unique([\n        features[i]\n        for i in valid\n    ])\n)\n\nfeature_codes = pd.Categorical(\n    [\n        features[i]\n        for i in valid\n    ],\n    categories=feature_names\n).codes\n\nX = matrix.sparse.to_coo().tocsr()\n\nmapping = csr_matrix(\n    (\n        np.ones(len(feature_codes), dtype=np.float32),\n        (\n            np.arange(len(feature_codes)),\n            feature_codes\n        )\n    ),\n    shape=(len(feature_codes), len(feature_names))\n)\n\nX_gene = X @ mapping\n\ngene_reg_matrix = pd.DataFrame.sparse.from_spmatrix(\n    X_gene,\n    index=matrix.index,\n    columns=feature_names\n)"},{"cell_type":"code","execution_count":51,"id":"e9de834c-1ee9-4bd4-ab1b-e9009fd3268c","metadata":{},"outputs":[],"source":"gene_reg_matrix = gene_reg_matrix.loc[\n    :,\n    ~gene_reg_matrix.columns.str.endswith(\"_other\")\n]\n\ngene_reg_matrix = gene_reg_matrix.loc[\n    :,\n    gene_reg_matrix.sum(axis=0) > 0\n]"},{"cell_type":"code","execution_count":62,"id":"b0e68583-d63e-4360-9750-23a01366bc7a","metadata":{},"outputs":[],"source":"genes = adata.var[\"location\"].astype(str)\n\ngene_reg_matrix = gene_reg_matrix.loc[\n    :,\n    [\n        col\n        for gene in genes\n        for col in (f\"{gene}_promoter\", f\"{gene}_enhancer\")\n        if col in gene_reg_matrix.columns\n    ]\n].copy()"},{"cell_type":"code","execution_count":65,"id":"2e9cbd8f-f552-4f97-995f-ae46570b31c8","metadata":{},"outputs":[],"source":"promoter_matrix = gene_reg_matrix.loc[\n    :,\n    gene_reg_matrix.columns.str.endswith(\"_promoter\")\n].copy()\n\nenhancer_matrix = gene_reg_matrix.loc[\n    :,\n    gene_reg_matrix.columns.str.endswith(\"_enhancer\")\n].copy()"},{"cell_type":"code","execution_count":66,"id":"be69b272-b476-4827-81a2-962a7e1effc5","metadata":{},"outputs":[],"source":"promoter_matrix"},{"cell_type":"code","execution_count":null,"id":"5c51be67-e6ad-4fdb-81ad-5fa07a4e36da","metadata":{},"outputs":[],"source":""},{"cell_type":"markdown","id":"735ee862-92bc-4071-a29f-ac5882ad37c1","metadata":{},"source":"# Multiome - RNA"},{"cell_type":"code","execution_count":56,"id":"5e73cc25-4184-47e2-9b3e-2347f1fdb8ae","metadata":{},"outputs":[],"source":"import pandas as pd\nfrom scipy import sparse\n\n\nrna_idx = adata.var[\"feature_type\"] == \"Gene Expression\"\n\n\nrna_adata = adata[:, rna_idx]\n\n\nX = rna_adata.X\n\nif sparse.issparse(X):\n    X = X.toarray()\n\n\nrna_matrix = pd.DataFrame(\n    X,\n    index=rna_adata.obs_names,\n    columns=rna_adata.var_names\n)"},{"cell_type":"code","execution_count":null,"id":"afdf9c5b-e8cb-47c5-9b88-a5f9bfa208e4","metadata":{},"outputs":[],"source":""},{"cell_type":"markdown","id":"1e3b15a3-a94c-4a07-ba26-677703230eb0","metadata":{},"source":"# Merge adata"},{"cell_type":"code","execution_count":73,"id":"d0ee15bf-d45b-4347-b508-f5efb692b147","metadata":{},"outputs":[],"source":"import numpy as np\nimport pandas as pd\nimport anndata as ad\n\ngenes = rna_matrix.columns.astype(str)\ncells = rna_matrix.index\n\npromoter_columns = [f\"{gene}_promoter\" for gene in genes]\nenhancer_columns = [f\"{gene}_enhancer\" for gene in genes]\n\npromoter_matrix = promoter_matrix.reindex(\n    index=cells,\n    columns=promoter_columns,\n    fill_value=0\n)\n\nenhancer_matrix = enhancer_matrix.reindex(\n    index=cells,\n    columns=enhancer_columns,\n    fill_value=0\n)\n\npromoter_matrix.columns = genes\nenhancer_matrix.columns = genes\n\nrna_matrix = rna_matrix.loc[cells, genes]\npromoter_matrix = promoter_matrix.loc[cells, genes]\nenhancer_matrix = enhancer_matrix.loc[cells, genes]\n\nobs = adata.obs.reindex(cells).copy()\n\nadata_multi = ad.AnnData(\n    X=rna_matrix.astype(np.float32).values,\n    obs=obs,\n    var=pd.DataFrame(index=genes)\n)\n\nadata_multi.layers[\"rna_matrix\"] = rna_matrix.astype(np.float32).values\nadata_multi.layers[\"promoter_matrix\"] = promoter_matrix.astype(np.float32).values\nadata_multi.layers[\"enhancer_matrix\"] = enhancer_matrix.astype(np.float32).values"},{"cell_type":"code","execution_count":74,"id":"adc525ba-35f2-42d4-b48c-2e59dc92b82a","metadata":{},"outputs":[],"source":"adata_multi"},{"cell_type":"code","execution_count":75,"id":"a255d632-3518-4944-bb39-457fbbadd4f9","metadata":{},"outputs":[],"source":"output_path = \"/home/Data_Drive_8TB/kykim/7. Kaggle/Single-Cell_Perturbations/adata_multi.h5ad\"\n\nadata_multi.write_h5ad(\n    output_path,\n    compression=\"gzip\")"},{"cell_type":"markdown","id":"64f39212-e810-4af1-bd1f-5ec0d964ad99","metadata":{},"source":"# Multihead attention"},{"cell_type":"code","execution_count":1,"id":"90a7cd6f-de54-4b1a-bdad-52988ac6a442","metadata":{},"outputs":[],"source":"import scanpy as sc\n\ninput_path = \"/home/Data_Drive_8TB/kykim/7. Kaggle/Single-Cell_Perturbations/adata_multi.h5ad\"\n\nadata_multi = sc.read_h5ad(\n    input_path\n)"},{"cell_type":"code","execution_count":2,"id":"b417af4c-92b5-46d1-9c25-3214a0028007","metadata":{},"outputs":[],"source":"import numpy as np\nimport torch\nimport torch.nn as nn\n\n\ndevice = torch.device(\n    \"cuda\" if torch.cuda.is_available() else \"cpu\"\n)\n\nbatch_size = 32\n\nmodality_dim = 32\ntoken_dim = 64\nnum_tokens = 256\n\nattention_heads = 4\ngene_transformer_layers = 2\n\nglobal_hidden_dim = 512\nlatent_dim = 128\n\nlearning_rate = 1e-4\n\n\nrna = adata_multi.layers[\"rna_matrix\"]\npromoter = adata_multi.layers[\"promoter_matrix\"]\nenhancer = adata_multi.layers[\"enhancer_matrix\"]\n\ncells = adata_multi.obs_names.to_numpy()\ngenes = adata_multi.var_names.to_numpy()\n\ncell_types = (\n    adata_multi.obs[\"cell_type\"]\n    .astype(str)\n    .to_numpy()\n)\n\ngene_count = len(genes)\n\nassert rna.shape == promoter.shape\nassert rna.shape == enhancer.shape\nassert rna.shape[1] == gene_count\n\n\nclass ModalityEncoder(nn.Module):\n\n    def __init__(self, latent_dim):\n\n        super().__init__()\n\n        self.encoder = nn.Sequential(\n            nn.Linear(1, 32),\n            nn.GELU(),\n            nn.Linear(32, latent_dim)\n        )\n\n    def forward(self, x):\n\n        return self.encoder(\n            x.unsqueeze(-1)\n        )\n\n\nclass GeneRegulatoryAttention(nn.Module):\n\n    def __init__(\n        self,\n        dim,\n        heads\n    ):\n\n        super().__init__()\n\n        self.q_proj = nn.Linear(\n            dim,\n            dim\n        )\n\n        self.k_proj = nn.Linear(\n            dim,\n            dim\n        )\n\n        self.v_proj = nn.Linear(\n            dim,\n            dim\n        )\n\n        self.scale = dim ** -0.5\n\n        self.out_proj = nn.Linear(\n            dim,\n            dim\n        )\n\n        self.norm = nn.LayerNorm(\n            dim\n        )\n\n        self.ffn = nn.Sequential(\n            nn.Linear(\n                dim,\n                dim * 2\n            ),\n            nn.GELU(),\n            nn.Linear(\n                dim * 2,\n                dim\n            )\n        )\n\n        self.ffn_norm = nn.LayerNorm(\n            dim\n        )\n\n    def forward(\n        self,\n        rna,\n        promoter,\n        enhancer\n    ):\n\n        q_input = rna * enhancer\n\n        k_input = rna * promoter\n\n        v_input = rna\n\n        q = self.q_proj(\n            q_input\n        )\n\n        k = self.k_proj(\n            k_input\n        )\n\n        v = self.v_proj(\n            v_input\n        )\n\n        score = (\n            q * k\n        ).sum(\n            dim=-1,\n            keepdim=True\n        ) * self.scale\n\n        attention = torch.sigmoid(\n            score\n        )\n\n        x = v * attention\n\n        x = self.out_proj(\n            x\n        )\n\n        x = self.norm(\n            x + rna\n        )\n\n        x = self.ffn_norm(\n            x + self.ffn(x)\n        )\n\n        return x, attention.squeeze(-1)\n\n\nclass MultiomeGeneEncoder(nn.Module):\n\n    def __init__(\n        self,\n        gene_count,\n        modality_dim,\n        token_dim,\n        num_tokens,\n        attention_heads,\n        gene_transformer_layers,\n        global_hidden_dim,\n        latent_dim\n    ):\n\n        super().__init__()\n\n        self.gene_count = gene_count\n        self.num_tokens = num_tokens\n\n        self.rna_encoder = ModalityEncoder(\n            modality_dim\n        )\n\n        self.promoter_encoder = ModalityEncoder(\n            modality_dim\n        )\n\n        self.enhancer_encoder = ModalityEncoder(\n            modality_dim\n        )\n\n        self.regulatory_attention = GeneRegulatoryAttention(\n            modality_dim,\n            attention_heads\n        )\n\n        self.fusion = nn.Sequential(\n            nn.Linear(\n                modality_dim * 2,\n                token_dim\n            ),\n            nn.GELU()\n        )\n\n        kernel_size = max(\n            1,\n            gene_count // num_tokens\n        )\n\n        self.gene_tokenizer = nn.Conv1d(\n            in_channels=token_dim,\n            out_channels=token_dim,\n            kernel_size=kernel_size,\n            stride=kernel_size\n        )\n\n        self.token_count = (\n            gene_count - kernel_size\n        ) // kernel_size + 1\n\n        encoder_layer = nn.TransformerEncoderLayer(\n            d_model=token_dim,\n            nhead=attention_heads,\n            dim_feedforward=token_dim * 4,\n            dropout=0.1,\n            activation=\"gelu\",\n            batch_first=True,\n            norm_first=True\n        )\n\n        self.gene_transformer = nn.TransformerEncoder(\n            encoder_layer,\n            num_layers=gene_transformer_layers\n        )\n\n        self.token_norm = nn.LayerNorm(\n            token_dim\n        )\n\n        self.global_encoder = nn.Sequential(\n            nn.Flatten(),\n            nn.Linear(\n                self.token_count * token_dim,\n                global_hidden_dim\n            ),\n            nn.GELU(),\n            nn.Linear(\n                global_hidden_dim,\n                latent_dim\n            )\n        )\n\n    def encode_pseudobulk(\n        self,\n        rna,\n        promoter,\n        enhancer\n    ):\n\n        rna_z = self.rna_encoder(\n            rna\n        )\n\n        promoter_z = self.promoter_encoder(\n            promoter\n        )\n\n        enhancer_z = self.enhancer_encoder(\n            enhancer\n        )\n\n        regulatory_z, attention = (\n            self.regulatory_attention(\n                rna=rna_z,\n                promoter=promoter_z,\n                enhancer=enhancer_z\n            )\n        )\n\n        x = torch.cat(\n            [\n                rna_z,\n                regulatory_z\n            ],\n            dim=-1\n        )\n\n        x = self.fusion(\n            x\n        )\n\n        x = x.transpose(\n            1,\n            2\n        )\n\n        x = self.gene_tokenizer(\n            x\n        )\n\n        x = x.transpose(\n            1,\n            2\n        )\n\n        x = self.token_norm(\n            x\n        )\n\n        x = self.gene_transformer(\n            x\n        )\n\n        return x, attention\n\n    def encode_baseline(\n        self,\n        cell_tokens\n    ):\n\n        return self.global_encoder(\n            cell_tokens\n        )\n\n\nmodel = MultiomeGeneEncoder(\n    gene_count=gene_count,\n    modality_dim=modality_dim,\n    token_dim=token_dim,\n    num_tokens=num_tokens,\n    attention_heads=attention_heads,\n    gene_transformer_layers=gene_transformer_layers,\n    global_hidden_dim=global_hidden_dim,\n    latent_dim=latent_dim\n).to(device)\n\n\nprint(\"Device:\", device)\n\nif torch.cuda.is_available():\n\n    print(\n        \"GPU:\",\n        torch.cuda.get_device_name(0)\n    )\n\n    print(\n        \"GPU memory:\",\n        round(\n            torch.cuda.memory_allocated() / 1024**3,\n            2\n        ),\n        \"GB\"\n    )"},{"cell_type":"code","execution_count":3,"id":"15b71609-b868-4db7-9038-85d823fe8b56","metadata":{},"outputs":[],"source":"cell_type_to_indices = {}\n\nfor i, ct in enumerate(cell_types):\n\n    cell_type_to_indices.setdefault(\n        ct,\n        []\n    ).append(i)\n\n\npseudobulk_rna = {}\npseudobulk_promoter = {}\npseudobulk_enhancer = {}\n\n\nfor ct, indices in cell_type_to_indices.items():\n\n    rna_ct = rna[indices].mean(\n        axis=0\n    )\n\n    promoter_ct = promoter[indices].mean(\n        axis=0\n    )\n\n    enhancer_ct = enhancer[indices].mean(\n        axis=0\n    )\n\n    if hasattr(rna_ct, \"A1\"):\n        rna_ct = rna_ct.A1\n    else:\n        rna_ct = np.asarray(\n            rna_ct\n        ).reshape(-1)\n\n    if hasattr(promoter_ct, \"A1\"):\n        promoter_ct = promoter_ct.A1\n    else:\n        promoter_ct = np.asarray(\n            promoter_ct\n        ).reshape(-1)\n\n    if hasattr(enhancer_ct, \"A1\"):\n        enhancer_ct = enhancer_ct.A1\n    else:\n        enhancer_ct = np.asarray(\n            enhancer_ct\n        ).reshape(-1)\n\n    pseudobulk_rna[ct] = rna_ct.astype(\n        np.float32,\n        copy=False\n    )\n\n    pseudobulk_promoter[ct] = promoter_ct.astype(\n        np.float32,\n        copy=False\n    )\n\n    pseudobulk_enhancer[ct] = enhancer_ct.astype(\n        np.float32,\n        copy=False\n    )\n\n\ncell_type_embeddings = {}\nbaseline_latents = {}\nattention_weights = {}\n\n\nmodel.eval()\n\n\nfor ct in cell_type_to_indices:\n\n    rna_input = torch.from_numpy(\n        pseudobulk_rna[ct]\n    ).unsqueeze(0).to(device)\n\n    promoter_input = torch.from_numpy(\n        pseudobulk_promoter[ct]\n    ).unsqueeze(0).to(device)\n\n    enhancer_input = torch.from_numpy(\n        pseudobulk_enhancer[ct]\n    ).unsqueeze(0).to(device)\n\n    with torch.no_grad():\n\n        tokens, attention = (\n            model.encode_pseudobulk(\n                rna_input,\n                promoter_input,\n                enhancer_input\n            )\n        )\n\n        baseline_z = (\n            model.encode_baseline(\n                tokens\n            )\n        )\n\n    cell_type_embeddings[ct] = (\n        tokens.squeeze(0).cpu()\n    )\n\n    baseline_latents[ct] = (\n        baseline_z.squeeze(0).cpu()\n    )\n\n    attention_weights[ct] = (\n        attention.squeeze(0).cpu()\n    )\n\n    del (\n        rna_input,\n        promoter_input,\n        enhancer_input,\n        tokens,\n        attention,\n        baseline_z\n    )\n\n    if torch.cuda.is_available():\n\n        torch.cuda.empty_cache()\n\n\ncell_type_order = list(\n    cell_type_to_indices.keys()\n)\n\n\ncell_type_embeddings = torch.stack(\n    [\n        cell_type_embeddings[ct]\n        for ct in cell_type_order\n    ]\n)\n\n\nbaseline_latents = torch.stack(\n    [\n        baseline_latents[ct]\n        for ct in cell_type_order\n    ]\n)\n\n\nattention_weights = torch.stack(\n    [\n        attention_weights[ct]\n        for ct in cell_type_order\n    ]\n)\n\n\nprint(\n    \"Number of cell types:\",\n    len(cell_type_order)\n)\n\nprint(\n    \"Cell-type token representation:\",\n    cell_type_embeddings.shape\n)\n\nprint(\n    \"Baseline latent:\",\n    baseline_latents.shape\n)\n\nprint(\n    \"Gene regulatory attention:\",\n    attention_weights.shape\n)"},{"cell_type":"code","execution_count":4,"id":"e0ab1933-b5a0-463e-adff-42971cf47bc4","metadata":{},"outputs":[],"source":"import os\nimport torch\n\nsave_dir = \"/home/Data_Drive_8TB/kykim/7. Kaggle/Single-Cell_Perturbations/attention\"\n\nos.makedirs(save_dir, exist_ok=True)\n\ntorch.save(\n    cell_type_embeddings,\n    os.path.join(\n        save_dir,\n        \"cell_type_embeddings.pt\"\n    )\n)\n\ntorch.save(\n    baseline_latents,\n    os.path.join(\n        save_dir,\n        \"baseline_latents.pt\"\n    )\n)\n\ntorch.save(\n    attention_weights,\n    os.path.join(\n        save_dir,\n        \"attention_weights.pt\"\n    )\n\n)\n\nnp.save(\n    os.path.join(\n        save_dir,\n        \"cell_type_order.npy\"\n    ),\n    np.array(\n        cell_type_order,\n        dtype=str\n    )\n)\n\nprint(\"Saved:\")\nprint(\n    \"cell_type_embeddings:\",\n    cell_type_embeddings.shape\n)\nprint(\n    \"baseline_latents:\",\n    baseline_latents.shape\n)\nprint(\n    \"attention_weights:\",\n    len(attention_weights)\n)\nprint(\n    \"cell_type_order:\",\n    cell_type_order\n)"},{"cell_type":"code","execution_count":23,"id":"e5b49266-9d52-4a74-bc54-cc59059b6e7d","metadata":{},"outputs":[],"source":"import matplotlib.pyplot as plt\nfrom sklearn.decomposition import PCA\n\nZ = baseline_latents.detach().cpu().numpy()\n\nmax_components = min(\n    Z.shape[0],\n    Z.shape[1]\n)\n\npca = PCA(\n    n_components=max_components\n)\n\npca.fit(Z)\n\nexplained_variance = pca.explained_variance_ratio_\ncumulative_variance = np.cumsum(\n    explained_variance\n)\n\nplt.figure(figsize=(6, 5))\n\nplt.plot(\n    range(1, max_components + 1),\n    explained_variance,\n    marker=\"o\"\n)\n\nplt.xlabel(\"Number of Principal Components\")\nplt.ylabel(\"Explained Variance Ratio\")\nplt.title(\"PCA Elbow Plot\")\nplt.xticks(\n    range(1, max_components + 1)\n)\nplt.tight_layout()\nplt.show()"},{"cell_type":"code","execution_count":5,"id":"f411b590-f3f5-493f-9436-29a1f371a7d1","metadata":{},"outputs":[],"source":"import matplotlib.pyplot as plt\nfrom sklearn.decomposition import PCA\n\nZ = baseline_latents.detach().cpu().numpy()\n\npca = PCA(n_components=5)\n\nZ_pca = pca.fit_transform(Z)\n\nplt.figure(figsize=(6, 5))\n\nfor i, ct in enumerate(cell_type_order):\n    plt.scatter(\n        Z_pca[i, 0],\n        Z_pca[i, 1],\n        s=100,\n        label=ct\n    )\n\n    plt.text(\n        Z_pca[i, 0],\n        Z_pca[i, 1],\n        ct\n    )\n\nplt.xlabel(\"PC1\")\nplt.ylabel(\"PC2\")\nplt.title(\"Baseline latent space\")\nplt.legend()\nplt.tight_layout()\nplt.show()"},{"cell_type":"code","execution_count":6,"id":"0a6a05ea-4db3-4270-9aa4-15898b57dd5b","metadata":{},"outputs":[],"source":"latent = baseline_latents.detach().cpu().numpy()\n\ncorr_matrix = np.corrcoef(latent)\n\nprint(corr_matrix)"},{"cell_type":"code","execution_count":7,"id":"b288dd74-6adc-458c-9244-37a118217479","metadata":{},"outputs":[],"source":"latent = baseline_latents.detach().cpu().numpy()\n\nlatent_corr = np.corrcoef(latent.T)\n\nprint(latent_corr.shape)\nprint(latent_corr)"},{"cell_type":"code","execution_count":8,"id":"07110cbf-a8df-464d-a96f-121abcdb9716","metadata":{},"outputs":[],"source":"import pandas as pd\n\nbaseline_latents_df = pd.DataFrame(\n    baseline_latents.detach().cpu().numpy(),\n    index=cell_type_order\n)\n\nbaseline_latents_df.index.name = \"cell_type\"\n\nbaseline_latents_df.columns = [\n    f\"latent_{i+1}\"\n    for i in range(baseline_latents_df.shape[1])\n]\n\nprint(baseline_latents_df)"},{"cell_type":"markdown","id":"8cecae8a-5cc2-4f3b-a117-8e383e4299b1","metadata":{},"source":"# Drug latent by MLP"},{"cell_type":"code","execution_count":9,"id":"6734ed2a-8c17-42e3-9dab-ec6f08cf3ff4","metadata":{},"outputs":[],"source":"import pandas as pd\n\nde_train = pd.read_parquet(\n    \"/home/Data_Drive_8TB/kykim/7. Kaggle/Single-Cell_Perturbations/baseline/de_train.parquet\")"},{"cell_type":"code","execution_count":10,"id":"00a06c51-eabc-4dcc-9386-9497b0beb49c","metadata":{},"outputs":[],"source":"de_train = de_train[\n    de_train[\"control\"] != True\n].copy()\n\nde_train_no_smiles = de_train.drop(\n    columns=[\"sm_lincs_id\", \"SMILES\",\"control\"]\n)"},{"cell_type":"code","execution_count":11,"id":"d0435aad-b18a-4097-a64b-20684af3d2ce","metadata":{},"outputs":[],"source":"import os\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import TensorDataset, DataLoader\n\n\ndevice = torch.device(\n    \"cuda\" if torch.cuda.is_available() else \"cpu\"\n)\n\n\ntraining_cell_types = [\n    \"T cells CD4+\",\n    \"NK cells\",\n    \"T regulatory cells\",\n    \"T cells CD8+\"\n]\n\n\ngene_columns = [\n    col\n    for col in de_train_no_smiles.columns\n    if col not in [\n        \"cell_type\",\n        \"sm_name\",\n        \"control\"\n    ]\n]\n\n\ncell_type_order = list(\n    cell_type_to_indices.keys()\n)\n\n\nbaseline_latents_cpu = (\n    baseline_latents\n    .detach()\n    .cpu()\n    .float()\n)\n\n\nlatent_dict = {\n    ct: baseline_latents_cpu[i]\n    for i, ct in enumerate(cell_type_order)\n}\n\n\nprint(\n    \"Baseline latent:\",\n    baseline_latents_cpu.shape\n)\n\n\nprint(\n    \"Training cell types:\",\n    training_cell_types\n)\n\n\nfor ct in training_cell_types:\n\n    if ct not in latent_dict:\n        raise ValueError(\n            f\"{ct} not found in baseline_latents\"\n        )\n\n\ndrug_names = (\n    de_train_no_smiles[\"sm_name\"]\n    .drop_duplicates()\n    .tolist()\n)\n\n\ndrug_to_index = {\n    drug: i\n    for i, drug in enumerate(drug_names)\n}\n\n\nnum_drugs = len(drug_names)\n\n\nX_list = []\ndrug_index_list = []\ny_list = []\nsample_info = []\n\n\nfor _, row in de_train_no_smiles.iterrows():\n\n    ct = row[\"cell_type\"]\n    drug = row[\"sm_name\"]\n\n    if ct not in training_cell_types:\n        continue\n\n    cell_latent = latent_dict[ct]\n\n    y = torch.tensor(\n        row[gene_columns].to_numpy(\n            dtype=np.float32\n        ),\n        dtype=torch.float32\n    )\n\n    X_list.append(\n        cell_latent\n    )\n\n    drug_index_list.append(\n        drug_to_index[drug]\n    )\n\n    y_list.append(\n        y\n    )\n\n    sample_info.append(\n        {\n            \"cell_type\": ct,\n            \"sm_name\": drug\n        }\n    )\n\n\nX = torch.stack(\n    X_list\n)\n\ndrug_indices = torch.tensor(\n    drug_index_list,\n    dtype=torch.long\n)\n\ny = torch.stack(\n    y_list\n)\n\nsample_info = pd.DataFrame(\n    sample_info\n)\n\n\nprint(\n    \"X:\",\n    X.shape\n)\n\nprint(\n    \"Drug indices:\",\n    drug_indices.shape\n)\n\nprint(\n    \"y:\",\n    y.shape\n)\n\nprint(\n    \"Number of drugs:\",\n    num_drugs\n)\n\nprint(\n    \"Number of genes:\",\n    len(gene_columns)\n)\n\n\nclass DrugConditionedDEModel(nn.Module):\n\n    def __init__(\n        self,\n        latent_input_dim,\n        num_drugs,\n        drug_embedding_dim,\n        output_dim\n    ):\n\n        super().__init__()\n\n        self.cell_encoder = nn.Sequential(\n\n            nn.Linear(\n                latent_input_dim,\n                1024\n            ),\n\n            nn.GELU(),\n\n            nn.Linear(\n                1024,\n                512\n            ),\n\n            nn.GELU(),\n\n            nn.Linear(\n                512,\n                256\n            ),\n\n            nn.GELU()\n        )\n\n        self.drug_embedding = nn.Embedding(\n            num_drugs,\n            drug_embedding_dim\n        )\n\n        self.de_head = nn.Sequential(\n\n            nn.Linear(\n                256 + drug_embedding_dim,\n                512\n            ),\n\n            nn.GELU(),\n\n            nn.Linear(\n                512,\n                512\n            ),\n\n            nn.GELU(),\n\n            nn.Linear(\n                512,\n                output_dim\n            )\n        )\n\n\n    def forward(\n        self,\n        cell_latent,\n        drug_index\n    ):\n\n        cell_hidden = self.cell_encoder(\n            cell_latent\n        )\n\n        drug_hidden = self.drug_embedding(\n            drug_index\n        )\n\n        combined = torch.cat(\n            [\n                cell_hidden,\n                drug_hidden\n            ],\n            dim=1\n        )\n\n        output = self.de_head(\n            combined\n        )\n\n        return output\n\n\nmodel = DrugConditionedDEModel(\n    latent_input_dim=X.shape[1],\n    num_drugs=num_drugs,\n    drug_embedding_dim=128,\n    output_dim=y.shape[1]\n).to(device)\n\n\ndataset = TensorDataset(\n    X,\n    drug_indices,\n    y\n)\n\n\nloader = DataLoader(\n    dataset,\n    batch_size=32,\n    shuffle=True,\n    pin_memory=torch.cuda.is_available()\n)\n\n\ncriterion = nn.MSELoss()\n\n\noptimizer = torch.optim.AdamW(\n    model.parameters(),\n    lr=1e-4,\n    weight_decay=1e-5\n)\n\n\nepochs = 500\nearly_stopping_patience = 30\n\n\nprint(\n    \"Model device:\",\n    device\n)\n\nprint(\n    \"Trainable parameters:\",\n    sum(\n        p.numel()\n        for p in model.parameters()\n        if p.requires_grad\n    )\n)\n\n\nloss_history = []\n\nbest_loss = float(\"inf\")\nbest_epoch = 0\npatience_counter = 0\n\nbest_state_dict = None\n\n\nfor epoch in range(epochs):\n\n    model.train()\n\n    epoch_loss = 0.0\n\n    for X_batch, drug_batch, y_batch in loader:\n\n        X_batch = X_batch.to(\n            device,\n            non_blocking=True\n        )\n\n        drug_batch = drug_batch.to(\n            device,\n            non_blocking=True\n        )\n\n        y_batch = y_batch.to(\n            device,\n            non_blocking=True\n        )\n\n        optimizer.zero_grad(\n            set_to_none=True\n        )\n\n        prediction = model(\n            X_batch,\n            drug_batch\n        )\n\n        loss = criterion(\n            prediction,\n            y_batch\n        )\n\n        loss.backward()\n\n        optimizer.step()\n\n        epoch_loss += (\n            loss.item()\n            * X_batch.size(0)\n        )\n\n    epoch_loss /= len(dataset)\n\n    loss_history.append(\n        epoch_loss\n    )\n\n    if epoch_loss < best_loss:\n\n        best_loss = epoch_loss\n        best_epoch = epoch + 1\n        patience_counter = 0\n\n        best_state_dict = {\n            k: v.detach().cpu().clone()\n            for k, v in model.state_dict().items()\n        }\n\n    else:\n\n        patience_counter += 1\n\n\n    if (\n        epoch + 1\n    ) % 10 == 0:\n\n        print(\n            f\"Epoch {epoch + 1}/{epochs} \"\n            f\"Loss: {epoch_loss:.6f} \"\n            f\"Best: {best_loss:.6f} \"\n            f\"Patience: {patience_counter}/{early_stopping_patience}\"\n        )\n\n\n    if patience_counter >= early_stopping_patience:\n\n        print(\n            f\"Early stopping at epoch {epoch + 1}\"\n        )\n\n        print(\n            f\"Best epoch: {best_epoch}\"\n        )\n\n        print(\n            f\"Best loss: {best_loss:.6f}\"\n        )\n\n        break\n\n\nif best_state_dict is not None:\n\n    model.load_state_dict(\n        best_state_dict\n    )\n\n\nmodel = model.to(device)\nmodel.eval()\n\n\nwith torch.no_grad():\n\n    predictions = []\n\n    for X_batch, drug_batch, _ in loader:\n\n        X_batch = X_batch.to(\n            device,\n            non_blocking=True\n        )\n\n        drug_batch = drug_batch.to(\n            device,\n            non_blocking=True\n        )\n\n        prediction = model(\n            X_batch,\n            drug_batch\n        )\n\n        predictions.append(\n            prediction.cpu()\n        )\n\n\npredictions = torch.cat(\n    predictions,\n    dim=0\n)\n\n\nprint(\n    \"Prediction:\",\n    predictions.shape\n)\n\n\nprint(\n    \"Best epoch:\",\n    best_epoch\n)\n\nprint(\n    \"Best loss:\",\n    best_loss\n)\n\n\nplt.figure(\n    figsize=(8, 6)\n)\n\nplt.plot(\n    range(\n        1,\n        len(loss_history) + 1\n    ),\n    loss_history\n)\n\nplt.axvline(\n    best_epoch,\n    linestyle=\"--\",\n    label=f\"Best epoch: {best_epoch}\"\n)\n\nplt.xlabel(\n    \"Epoch\"\n)\n\nplt.ylabel(\n    \"MSE Loss\"\n)\n\nplt.title(\n    \"Training Loss\"\n)\n\nplt.legend()\n\nplt.tight_layout()\n\nplt.show()\n\n\nsave_dir = (\n    \"/home/Data_Drive_8TB/kykim/\"\n    \"7. Kaggle/Single-Cell_Perturbations/\"\n    \"baseline/model\"\n)\n\n\nos.makedirs(\n    save_dir,\n    exist_ok=True\n)\n\n\ntorch.save(\n    model.state_dict(),\n    os.path.join(\n        save_dir,\n        \"drug_conditioned_de_model.pt\"\n    )\n)\n\n\ntorch.save(\n    {\n        \"model_state_dict\": model.state_dict(),\n        \"drug_to_index\": drug_to_index,\n        \"training_cell_types\": training_cell_types,\n        \"gene_columns\": gene_columns,\n        \"loss_history\": loss_history,\n        \"best_epoch\": best_epoch,\n        \"best_loss\": best_loss,\n        \"latent_input_dim\": X.shape[1],\n        \"drug_embedding_dim\": 128\n    },\n    os.path.join(\n        save_dir,\n        \"drug_conditioned_de_checkpoint.pt\"\n    )\n)\n\n\nnp.save(\n    os.path.join(\n        save_dir,\n        \"loss_history.npy\"\n    ),\n    np.asarray(\n        loss_history,\n        dtype=np.float32\n    )\n)\n\n\nprint(\n    \"Model saved:\",\n    save_dir\n)"},{"cell_type":"code","execution_count":12,"id":"54ffa88c-812d-4c82-b783-2b1bcbe3686e","metadata":{},"outputs":[],"source":"import numpy as np\nimport pandas as pd\nimport torch\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\n\nmodel.eval()\n\nX_device = X.to(device)\ndrug_indices_device = drug_indices.to(device)\n\nwith torch.no_grad():\n\n    predictions = model(\n        X_device,\n        drug_indices_device\n    ).cpu()\n\ntrue_values = y.cpu()\n\nresiduals = (\n    true_values - predictions\n).numpy()\n\nresidual_df = pd.DataFrame(\n    residuals,\n    columns=gene_columns\n)\n\nresidual_df[\"cell_type\"] = sample_info[\n    \"cell_type\"\n].values\n\nresidual_df[\"sm_name\"] = sample_info[\n    \"sm_name\"\n].values\n\n\ndef gaussian_pooling(\n    matrix,\n    output_size=1000,\n    sigma=10.0\n):\n\n    n_drugs, n_genes = matrix.shape\n\n    pooled = np.zeros(\n        (n_drugs, output_size),\n        dtype=np.float32\n    )\n\n    source_positions = np.linspace(\n        0,\n        n_genes - 1,\n        n_genes\n    )\n\n    target_positions = np.linspace(\n        0,\n        n_genes - 1,\n        output_size\n    )\n\n    for i, target in enumerate(\n        target_positions\n    ):\n\n        distance = (\n            source_positions - target\n        )\n\n        weights = np.exp(\n            -0.5 *\n            (distance / sigma) ** 2\n        )\n\n        weights /= weights.sum()\n\n        pooled[:, i] = (\n            matrix @ weights\n        )\n\n    return pooled\n\n\nfig, axes = plt.subplots(\n    2,\n    2,\n    figsize=(30, 20)\n)\n\naxes = axes.flatten()\n\n\nfor ax, ct in zip(\n    axes,\n    training_cell_types\n):\n\n    ct_mask = (\n        residual_df[\"cell_type\"] == ct\n    )\n\n    heatmap_data = residual_df.loc[\n        ct_mask,\n        gene_columns\n    ].to_numpy(\n        dtype=np.float32\n    )\n\n    pooled_residual = gaussian_pooling(\n        heatmap_data,\n        output_size=1000,\n        sigma=10.0\n    )\n\n    sns.heatmap(\n        pooled_residual,\n        ax=ax,\n        cmap=\"coolwarm\",\n        center=0,\n        xticklabels=False,\n        yticklabels=False,\n        cbar=True,\n        cbar_kws={\n            \"label\": \"Residual\"\n        }\n    )\n\n    ax.set_xlabel(\n        \"Gene\",\n        fontsize=18\n    )\n\n    ax.set_ylabel(\n        \"Drug\",\n        fontsize=18\n    )\n\n    ax.set_title(\n        ct,\n        fontsize=26,\n        fontweight=\"bold\",\n        pad=12\n    )\n\n\nplt.tight_layout(\n    pad=3\n)\n\nplt.show()\n\nplt.close()"},{"cell_type":"markdown","id":"f8be0ccf-f5cb-411a-80c2-b120895c99a7","metadata":{},"source":"# Final tunning (Myeloid, B cells)"},{"cell_type":"code","execution_count":13,"id":"18eba75d-63a1-4f44-8267-4aa7ab0b582f","metadata":{},"outputs":[],"source":"import os\nimport copy\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import TensorDataset, DataLoader\nimport matplotlib.pyplot as plt\n\n\ndevice = torch.device(\n    \"cuda\" if torch.cuda.is_available() else \"cpu\"\n)\n\n\ntarget_cell_types = [\n    \"B cells\",\n    \"Myeloid cells\"\n]\n\n\ngene_columns = [\n    col\n    for col in de_train_no_smiles.columns\n    if col not in [\n        \"cell_type\",\n        \"sm_name\",\n        \"control\"\n    ]\n]\n\n\nlatent_columns = [\n    col\n    for col in baseline_latents_df.columns\n    if str(col).startswith(\"latent_\")\n]\n\n\nif len(latent_columns) == 0:\n\n    latent_columns = [\n        col\n        for col in baseline_latents_df.columns\n        if pd.api.types.is_numeric_dtype(\n            baseline_latents_df[col]\n        )\n    ]\n\n\nprint(\n    \"Latent dimension:\",\n    len(latent_columns)\n)\n\nprint(\n    \"Number of genes:\",\n    len(gene_columns)\n)\n\n\ntarget_df = de_train_no_smiles[\n    de_train_no_smiles[\"cell_type\"].isin(\n        target_cell_types\n    )\n].copy()\n\n\nprint(\n    \"\\nFine-tuning data:\"\n)\n\nprint(\n    target_df[\"cell_type\"].value_counts()\n)\n\n\ndrug_names = list(\n    drug_to_index.keys()\n)\n\n\nX_list = []\ndrug_index_list = []\ny_list = []\n\n\nfor _, row in target_df.iterrows():\n\n    ct = row[\"cell_type\"]\n    drug = row[\"sm_name\"]\n\n    if ct not in baseline_latents_df.index:\n        continue\n\n    if drug not in drug_to_index:\n        continue\n\n    latent = torch.tensor(\n        baseline_latents_df.loc[\n            ct,\n            latent_columns\n        ].to_numpy(\n            dtype=np.float32\n        ),\n        dtype=torch.float32\n    )\n\n    target = torch.tensor(\n        row[\n            gene_columns\n        ].to_numpy(\n            dtype=np.float32\n        ),\n        dtype=torch.float32\n    )\n\n    X_list.append(\n        latent\n    )\n\n    drug_index_list.append(\n        drug_to_index[drug]\n    )\n\n    y_list.append(\n        target\n    )\n\n\nX_ft = torch.stack(\n    X_list\n)\n\ndrug_indices_ft = torch.tensor(\n    drug_index_list,\n    dtype=torch.long\n)\n\ny_ft = torch.stack(\n    y_list\n)\n\n\nprint(\n    \"\\nFine-tuning tensors:\"\n)\n\nprint(\n    \"X:\",\n    X_ft.shape\n)\n\nprint(\n    \"Drug:\",\n    drug_indices_ft.shape\n)\n\nprint(\n    \"Y:\",\n    y_ft.shape\n)\n\n\nclass DrugConditionedDEModel(nn.Module):\n\n    def __init__(\n        self,\n        latent_input_dim,\n        num_drugs,\n        drug_embedding_dim,\n        output_dim\n    ):\n\n        super().__init__()\n\n        self.cell_encoder = nn.Sequential(\n\n            nn.Linear(\n                latent_input_dim,\n                1024\n            ),\n\n            nn.GELU(),\n\n            nn.Linear(\n                1024,\n                512\n            ),\n\n            nn.GELU(),\n\n            nn.Linear(\n                512,\n                256\n            ),\n\n            nn.GELU()\n        )\n\n\n        self.drug_embedding = nn.Embedding(\n            num_drugs,\n            drug_embedding_dim\n        )\n\n\n        self.de_head = nn.Sequential(\n\n            nn.Linear(\n                256 + drug_embedding_dim,\n                512\n            ),\n\n            nn.GELU(),\n\n            nn.Linear(\n                512,\n                512\n            ),\n\n            nn.GELU(),\n\n            nn.Linear(\n                512,\n                output_dim\n            )\n        )\n\n\n    def forward(\n        self,\n        cell_latent,\n        drug_index\n    ):\n\n        cell_hidden = self.cell_encoder(\n            cell_latent\n        )\n\n        drug_hidden = self.drug_embedding(\n            drug_index\n        )\n\n        combined = torch.cat(\n            [\n                cell_hidden,\n                drug_hidden\n            ],\n            dim=1\n        )\n\n        return self.de_head(\n            combined\n        )\n\n\nmodel_ft = DrugConditionedDEModel(\n    latent_input_dim=len(latent_columns),\n    num_drugs=len(drug_names),\n    drug_embedding_dim=128,\n    output_dim=len(gene_columns)\n).to(device)\n\n\ncheckpoint_path = (\n    \"/home/Data_Drive_8TB/kykim/\"\n    \"7. Kaggle/Single-Cell_Perturbations/\"\n    \"baseline/model/\"\n    \"drug_conditioned_de_checkpoint.pt\"\n)\n\n\ncheckpoint = torch.load(\n    checkpoint_path,\n    map_location=device,\n    weights_only=False\n)\n\n\nmodel_ft.load_state_dict(\n    checkpoint[\"model_state_dict\"]\n)\n\n\nprint(\n    \"\\nPretrained model loaded.\"\n)\n\n\nfor param in model_ft.cell_encoder.parameters():\n    param.requires_grad = False\n\n\nfor param in model_ft.drug_embedding.parameters():\n    param.requires_grad = True\n\n\nfor param in model_ft.de_head.parameters():\n    param.requires_grad = True\n\n\ntrainable_parameters = sum(\n    p.numel()\n    for p in model_ft.parameters()\n    if p.requires_grad\n)\n\n\ntotal_parameters = sum(\n    p.numel()\n    for p in model_ft.parameters()\n)\n\n\nprint(\n    \"Trainable parameters:\",\n    trainable_parameters\n)\n\nprint(\n    \"Total parameters:\",\n    total_parameters\n)\n\n\ndataset_ft = TensorDataset(\n    X_ft,\n    drug_indices_ft,\n    y_ft\n)\n\n\nloader_ft = DataLoader(\n    dataset_ft,\n    batch_size=8,\n    shuffle=True,\n    drop_last=False\n)\n\n\ncriterion = nn.MSELoss()\n\n\noptimizer = torch.optim.AdamW(\n    [\n        {\n            \"params\": model_ft.drug_embedding.parameters(),\n            \"lr\": 1e-5\n        },\n        {\n            \"params\": model_ft.de_head.parameters(),\n            \"lr\": 1e-5\n        }\n    ],\n    weight_decay=1e-5\n)\n\n\nepochs = 500\n\n\nloss_history = []\n\n\nbest_loss = float(\"inf\")\n\nbest_state = None\n\npatience = 30\n\npatience_counter = 0\n\n\nprint(\n    \"\\nFine-tuning started.\"\n)\n\n\nfor epoch in range(epochs):\n\n    model_ft.train()\n\n    epoch_loss = 0.0\n\n\n    for X_batch, drug_batch, y_batch in loader_ft:\n\n        X_batch = X_batch.to(\n            device\n        )\n\n        drug_batch = drug_batch.to(\n            device\n        )\n\n        y_batch = y_batch.to(\n            device\n        )\n\n\n        optimizer.zero_grad(\n            set_to_none=True\n        )\n\n\n        prediction = model_ft(\n            X_batch,\n            drug_batch\n        )\n\n\n        loss = criterion(\n            prediction,\n            y_batch\n        )\n\n\n        loss.backward()\n\n\n        torch.nn.utils.clip_grad_norm_(\n            model_ft.parameters(),\n            max_norm=1.0\n        )\n\n\n        optimizer.step()\n\n\n        epoch_loss += (\n            loss.item()\n            * X_batch.size(0)\n        )\n\n\n    epoch_loss /= len(dataset_ft)\n\n    loss_history.append(\n        epoch_loss\n    )\n\n\n    if epoch_loss < best_loss:\n\n        best_loss = epoch_loss\n\n        best_state = copy.deepcopy(\n            model_ft.state_dict()\n        )\n\n        patience_counter = 0\n\n    else:\n\n        patience_counter += 1\n\n\n    if (\n        epoch + 1\n    ) % 10 == 0:\n\n        print(\n            f\"Epoch {epoch + 1}/{epochs} \"\n            f\"Loss: {epoch_loss:.6f} \"\n            f\"Best: {best_loss:.6f} \"\n            f\"Patience: \"\n            f\"{patience_counter}/{patience}\"\n        )\n\n\n    if patience_counter >= patience:\n\n        print(\n            f\"\\nEarly stopping at \"\n            f\"epoch {epoch + 1}\"\n        )\n\n        break\n\n\nmodel_ft.load_state_dict(\n    best_state\n)\n\n\nprint(\n    \"\\nBest fine-tuning loss:\",\n    best_loss\n)"},{"cell_type":"code","execution_count":14,"id":"60d64dcb-2a29-4b42-bf74-38e6fe3e503c","metadata":{},"outputs":[],"source":"plt.figure(\n    figsize=(8, 5)\n)\n\nplt.plot(\n    loss_history\n)\n\nplt.xlabel(\n    \"Epoch\"\n)\n\nplt.ylabel(\n    \"MSE Loss\"\n)\n\nplt.title(\n    \"B cells + Myeloid cells Fine-tuning\"\n)\n\nplt.tight_layout()\n\nplt.show()"},{"cell_type":"code","execution_count":15,"id":"42047a78-1a3b-436f-88e5-08f7219a85f6","metadata":{},"outputs":[],"source":"model_ft.eval()\n\n\nprediction_rows = []\n\n\nfor ct in target_cell_types:\n\n    cell_latent = torch.tensor(\n        baseline_latents_df.loc[\n            ct,\n            latent_columns\n        ].to_numpy(\n            dtype=np.float32\n        ),\n        dtype=torch.float32\n    )\n\n\n    for drug in drug_names:\n\n        drug_index = torch.tensor(\n            [drug_to_index[drug]],\n            dtype=torch.long\n        )\n\n\n        with torch.no_grad():\n\n            prediction = model_ft(\n                cell_latent.unsqueeze(0).to(device),\n                drug_index.to(device)\n            )\n\n\n        prediction = (\n            prediction\n            .squeeze(0)\n            .cpu()\n            .numpy()\n        )\n\n\n        row = {\n            \"cell_type\": ct,\n            \"sm_name\": drug\n        }\n\n\n        row.update(\n            {\n                gene: value\n                for gene, value in zip(\n                    gene_columns,\n                    prediction\n                )\n            }\n        )\n\n\n        prediction_rows.append(\n            row\n        )\n\n\npredicted_de_finetuned_df = pd.DataFrame(\n    prediction_rows,\n    columns=[\n        \"cell_type\",\n        \"sm_name\"\n    ] + gene_columns\n)\n\n\nprint(\n    \"Fine-tuned prediction shape:\",\n    predicted_de_finetuned_df.shape\n)\n\n\nprint(\n    predicted_de_finetuned_df[\n        [\n            \"cell_type\",\n            \"sm_name\"\n        ] + gene_columns[:5]\n    ].head()\n)"},{"cell_type":"markdown","id":"dc7642c2-495a-4eff-b1d0-0bde88b0875f","metadata":{},"source":"# Final submission"},{"cell_type":"code","execution_count":16,"id":"34e31c10-4b36-434e-8e47-14312287106e","metadata":{},"outputs":[],"source":"import pandas as pd\n\nsample_submission = pd.read_csv(\n    \"/home/Data_Drive_8TB/kykim/7. Kaggle/Single-Cell_Perturbations/submission/sample_submission.csv\"\n)"},{"cell_type":"code","execution_count":17,"id":"d79d9b1b-d23b-40f9-87dc-e38bd4982577","metadata":{},"outputs":[],"source":"sample_submission"},{"cell_type":"code","execution_count":18,"id":"7e3978c9-7dcc-4e5e-85c0-b48b887bcd7b","metadata":{},"outputs":[],"source":"missing_columns = [\n    col\n    for col in predicted_de_finetuned_df.columns\n    if col not in sample_submission.columns\n]\n\nprint(missing_columns)"},{"cell_type":"code","execution_count":19,"id":"e15cb858-e016-423f-8aa3-e77174582c87","metadata":{},"outputs":[],"source":"id_map = pd.read_csv(\n    \"/home/Data_Drive_8TB/kykim/7. Kaggle/Single-Cell_Perturbations/submission/id_map.csv\"\n)"},{"cell_type":"code","execution_count":20,"id":"64a24d9e-ad77-406b-82e7-65a8858e7d83","metadata":{},"outputs":[],"source":"id_map"},{"cell_type":"code","execution_count":21,"id":"c77e8ae5-737d-4c08-91f4-696ef7406c7e","metadata":{},"outputs":[],"source":"import numpy as np\nimport pandas as pd\n\nid_col = \"id\"\n\ngene_columns_submission = [\n    col\n    for col in sample_submission.columns\n    if col != id_col\n]\n\nprint(\"1. Checking gene columns...\")\n\nmissing_genes = [\n    gene\n    for gene in gene_columns_submission\n    if gene not in predicted_de_finetuned_df.columns\n]\n\nprint(\n    f\"Submission genes: {len(gene_columns_submission)}\"\n)\n\nprint(\n    f\"Missing genes: {len(missing_genes)}\"\n)\n\nif missing_genes:\n    print(missing_genes[:20])\n    raise ValueError(\"Missing genes detected\")\n\n\nprint(\"2. Preparing prediction table...\")\n\nprediction_table = predicted_de_finetuned_df[\n    [\"cell_type\", \"sm_name\"] + gene_columns_submission\n].copy()\n\nprint(\n    f\"Prediction rows: {len(prediction_table)}\"\n)\n\n\nprint(\"3. Matching id_map with predictions...\")\n\noutput = (\n    id_map[\n        [\"id\", \"cell_type\", \"sm_name\"]\n    ]\n    .merge(\n        prediction_table,\n        on=[\"cell_type\", \"sm_name\"],\n        how=\"left\",\n        validate=\"one_to_one\"\n    )\n)\n\nprint(\n    f\"Matched rows: {len(output)}\"\n)\n\n\nprint(\"4. Checking prediction values...\")\n\nnan_count = (\n    output[\n        gene_columns_submission\n    ]\n    .isna()\n    .sum()\n    .sum()\n)\n\ninf_count = np.isinf(\n    output[\n        gene_columns_submission\n    ].to_numpy(\n        dtype=np.float64\n    )\n).sum()\n\nprint(\n    f\"NaN values: {nan_count}\"\n)\n\nprint(\n    f\"Inf values: {inf_count}\"\n)\n\n\nif nan_count > 0:\n    raise ValueError(\n        f\"{nan_count} NaN values found\"\n    )\n\nif inf_count > 0:\n    raise ValueError(\n        f\"{inf_count} Inf values found\"\n    )\n\n\nprint(\"5. Reordering columns...\")\n\noutput = output[\n    [\"id\"] + gene_columns_submission\n]\n\n\nprint(\n    f\"Final output shape: {output.shape}\"\n)\n\nprint(\n    \"6. Checking ID order...\"\n)\n\nif output[\"id\"].equals(\n    sample_submission[\"id\"]\n):\n\n    print(\n        \"ID order: OK\"\n    )\n\nelse:\n\n    raise ValueError(\n        \"ID order does not match sample_submission\"\n    )\n\n\nprint(\n    \"7. Submission generation complete.\"\n)"},{"cell_type":"code","execution_count":22,"id":"d811d929-f64f-4c8e-a995-4d0c86a4b475","metadata":{},"outputs":[],"source":"save_dir = \"/home/Data_Drive_8TB/kykim/7. Kaggle/Single-Cell_Perturbations/submission/final\"\n\nos.makedirs(save_dir, exist_ok=True)\n\noutput_path = os.path.join(\n    save_dir,\n    \"submission.csv\"\n)\n\noutput.to_csv(\n    output_path,\n    index=False\n)"},{"cell_type":"code","execution_count":26,"id":"d11bafb4-0c97-451e-80d1-158c5e59902d","metadata":{},"outputs":[],"source":"import matplotlib.pyplot as plt\nfrom matplotlib.patches import FancyBboxPatch, FancyArrowPatch\n\nfig, ax = plt.subplots(figsize=(12, 4))\nax.set_xlim(0, 12)\nax.set_ylim(0, 4)\nax.axis(\"off\")\n\nboxes = [\n    (0.5, 1.5, 2.2, 1.0, \"True DE\\n$y_{ig}$\"),\n    (3.2, 1.5, 2.2, 1.0, \"Predicted DE\\n$\\\\hat{y}_{ig}$\"),\n    (6.0, 1.5, 2.2, 1.0, \"Gene-wise Error\\n$(y_{ig}-\\\\hat{y}_{ig})^2$\"),\n    (8.8, 1.5, 2.2, 1.0, \"Row-wise RMSE\\n$\\\\sqrt{\\\\frac{1}{G}\\\\sum_g error}$\"),\n]\n\nfor x, y, w, h, text in boxes:\n    box = FancyBboxPatch(\n        (x, y), w, h,\n        boxstyle=\"round,pad=0.05\",\n        linewidth=1.5,\n        facecolor=\"white\"\n    )\n    ax.add_patch(box)\n    ax.text(\n        x + w / 2,\n        y + h / 2,\n        text,\n        ha=\"center\",\n        va=\"center\",\n        fontsize=12\n    )\n\nfor i in range(len(boxes) - 1):\n    x1 = boxes[i][0] + boxes[i][2]\n    x2 = boxes[i + 1][0]\n    y = boxes[i][1] + boxes[i][3] / 2\n\n    ax.add_patch(\n        FancyArrowPatch(\n            (x1, y),\n            (x2, y),\n            arrowstyle=\"->\",\n            mutation_scale=15,\n            linewidth=1.5\n        )\n    )\n\nax.text(\n    10.9,\n    1.0,\n    r\"$MRRMSE=\\frac{1}{N}\\sum_{i=1}^{N}RMSE_i$\",\n    ha=\"center\",\n    va=\"center\",\n    fontsize=13\n)\n\nax.add_patch(\n    FancyArrowPatch(\n        (9.9, 1.5),\n        (10.9, 1.2),\n        arrowstyle=\"->\",\n        mutation_scale=15,\n        linewidth=1.5\n    )\n)\n\nplt.tight_layout()\nplt.show()"}],"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.11.15"}},"nbformat":4,"nbformat_minor":5}