{"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":7086037,"sourceType":"datasetVersion","datasetId":4082310}],"dockerImageVersionId":30587,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install scanpy pyro-ppl muon tqdm ipywidgets polars scikit-misc fastcluster","metadata":{"execution":{"iopub.status.busy":"2023-12-08T03:03:33.937225Z","iopub.execute_input":"2023-12-08T03:03:33.937538Z","iopub.status.idle":"2023-12-08T03:03:48.399737Z","shell.execute_reply.started":"2023-12-08T03:03:33.937505Z","shell.execute_reply":"2023-12-08T03:03:48.398538Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# !curl https://raw.githubusercontent.com/pytorch/xla/master/contrib/scripts/env-setup.py -o pytorch-xla-env-setup.py\n# !python pytorch-xla-env-setup.py --version nightly --apt-packages libomp5 libopenblas-dev","metadata":{"execution":{"iopub.status.busy":"2023-12-08T03:03:48.401419Z","iopub.execute_input":"2023-12-08T03:03:48.401801Z","iopub.status.idle":"2023-12-08T03:03:48.407100Z","shell.execute_reply.started":"2023-12-08T03:03:48.401763Z","shell.execute_reply":"2023-12-08T03:03:48.406067Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import warnings\nwarnings.simplefilter(action='ignore', category=FutureWarning)\nimport os\nos.environ['CUDA_LAUNCH_BLOCKING']='1'\nos.environ[\"TOKENIZERS_PARALLELISM\"] = \"false\"\nimport torch\nfrom torch import nn\nimport torch.nn.functional as F\nimport pyro\nfrom pyro import distributions as dist\nfrom pyro.distributions import constraints\nimport pandas as pd\nimport polars as pl\nimport scanpy as sc\nimport anndata as ad\nimport muon as mu\nfrom muon import MuData\nimport anndata\nfrom torch.utils.data import DataLoader\nfrom sklearn.preprocessing import OneHotEncoder, LabelEncoder\nfrom anndata.experimental.pytorch import AnnLoader\nimport numpy as np\nfrom transformers import AutoModelForMaskedLM, AutoTokenizer, pipeline, RobertaModel, RobertaTokenizer, AutoModel, AutoModelWithLMHead, AutoModelForMaskedLM\nfrom pathlib import Path\nimport scipy\nimport gc\nfrom tqdm import tqdm\nfrom functools import partial\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\ninput_dir = Path(\"/kaggle/input/open-problems-single-cell-perturbations\")\noutput_dir = Path(\"/kaggle/working\")\ntmp_dir = Path(\"/kaggle/tmp\")\n\n\ntry:\n    import shutil\n    shutil.copy2(\"/kaggle/input/prepared-anndatas-opscrna/adata_prepared_multiome.h5mu\", \"/kaggle/working/adata_prepared_multiome.h5mu\")\n    shutil.copy2(\"/kaggle/input/prepared-anndatas-opscrna/adata_prepared_train.h5ad\", \"/kaggle/working/adata_prepared_train.h5ad\")\nexcept:\n    print(\"FAILED TO COPY DATA\")\n    pass\n\ntry:\n    if not torch.cuda.is_available():\n        torch.set_num_threads(os.cpu_count())\n        print(\"Using\", os.cpu_count(), \"cores\")\nexcept:\n    print(\"Failed to increase torch cores.\")\n\nUSE_ATAC = False\nATAC_PRETRAIN_EPOCHS = 1\nFINETUNE_EPOCHS = 45\nEMBEDDING_MODEL = \"chemberta-zinc\"  # \"chemberta-pubchem\"  # \"molformer\" \"chemberta-zinc\" \"chemberta-pubchem\"\nHOLD_OUT = True","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-12-08T03:03:48.408673Z","iopub.execute_input":"2023-12-08T03:03:48.409533Z","iopub.status.idle":"2023-12-08T03:04:21.685781Z","shell.execute_reply.started":"2023-12-08T03:03:48.409496Z","shell.execute_reply":"2023-12-08T03:04:21.684659Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def read_competition(filter_excluded: bool = False):\n    # Ref: https://github.com/openproblems-bio/neurips-2023-scripts/blob/main/compute_de.ipynb\n    # https://www.kaggle.com/code/jeskowagner/converting-input-files-to-anndata/notebook\n    print(\"Reading competition adata\", flush=True)\n    with pl.StringCache():\n        adata = pl.scan_parquet(input_dir / 'adata_train.parquet', low_memory=True).cast({\n            'obs_id': pl.Categorical,\n            'gene': pl.Categorical,\n            'count': pl.UInt16\n        }).select(['obs_id', 'gene', 'count'])\n        adata_meta = pl.scan_csv(input_dir / 'adata_obs_meta.csv', low_memory=True, dtypes={\n            'library_id': pl.Categorical,\n            'plate_name': pl.Categorical,\n            'well': pl.Categorical,\n            'donor_id': pl.Categorical,\n            'cell_type': pl.Categorical,\n            # TODO: Include LINCS metadata?\n            'sm_name': pl.Categorical,\n            'SMILES': pl.Utf8,\n            'control': pl.Boolean,\n            'obs_id': pl.Categorical,\n        }).select(['library_id', 'plate_name', 'well', 'donor_id', 'cell_type', 'sm_name', 'SMILES', 'control', 'obs_id'])\n\n        excluded_ids = pl.scan_csv(input_dir / 'adata_excluded_ids.csv', low_memory=True, dtypes={'obs_id': pl.Categorical, 'gene': pl.Categorical})\n\n        if filter_excluded:\n            adata = adata.join(excluded_ids, on=['obs_id', 'gene'], how='anti')  # Filter out excluded ids\n        del excluded_ids\n        gc.collect()\n\n        print(\"Building the final adata object\", flush=True)\n        return build_anndata(adata, adata_meta)\n\ndef build_anndata(adata, adata_meta_obs, adata_meta_var=None, gene_id_col='gene', count_col=\"count\", count_dtype=np.uint16):\n    obs_ids = adata.select('obs_id').unique().collect(streaming=True).get_column('obs_id').to_list()\n    obs_id_map = {obs_id: index for index, obs_id in enumerate(obs_ids)}\n\n    genes = adata.select(gene_id_col).unique().collect(streaming=True).get_column(gene_id_col).to_list()\n    gene_map = {gene: index for index, gene in enumerate(genes)}\n\n    adata = adata.with_columns(pl.col('obs_id').map_dict(obs_id_map, return_dtype=pl.UInt32).cast(pl.UInt32).alias('obs_index'))\n#     adata_meta_obs = adata_meta_obs.filter(pl.col('obs_id').is_in(obs_ids))\n    adata_meta_obs = adata_meta_obs.with_columns(pl.col('obs_id').map_dict(obs_id_map, return_dtype=pl.UInt32).cast(pl.UInt32).alias('obs_index'))\n    adata_meta_obs = adata_meta_obs.filter(pl.col('obs_index').is_not_null())\n\n    adata = adata.with_columns(pl.col(gene_id_col).map_dict(gene_map, return_dtype=pl.UInt32).cast(pl.UInt32).alias('gene_index'))\n\n    if adata_meta_var is not None:\n#         adata_meta_var = adata_meta_var.filter(pl.col('gene_id').is_in(genes))\n        adata_meta_var = adata_meta_var.with_columns(pl.col(gene_id_col).map_dict(gene_map, return_dtype=pl.UInt32).cast(pl.UInt32).alias('gene_index'))\n        adata_meta_var = adata_meta_var.filter(pl.col('gene_index').is_not_null())\n\n     # Filter 0s\n    adata = adata.filter(pl.col(count_col).is_not_null()).filter(pl.col(count_col) > 0)\n\n    # Re-order by index\n    adata = adata.sort(['obs_index', 'gene_index'])\n    adata_meta_obs = adata_meta_obs.sort('obs_index')\n    if adata_meta_var is not None:\n        adata_meta_var = adata_meta_var.sort('gene_index')\n\n    del obs_ids\n    del obs_id_map\n    del genes\n    del gene_map\n\n    gc.collect()\n    row_indices = adata.select('obs_index').collect().get_column('obs_index').to_numpy()\n    gc.collect()\n    col_indices = adata.select('gene_index').collect().get_column('gene_index').to_numpy()\n    gc.collect()\n    counts = scipy.sparse.csr_matrix((adata.select(count_col).collect().get_column(count_col).to_numpy(), (row_indices, col_indices)))\n\n    if adata_meta_var is None:\n        final_genes = adata.select(gene_id_col).collect().get_column(gene_id_col).cat.get_categories().to_numpy()\n        var_df = pd.DataFrame({gene_id_col: final_genes}, index=final_genes)\n        del final_genes\n    else:\n        var_df = adata_meta_var.collect().to_pandas()\n    gc.collect()\n\n    return ad.AnnData(\n        X=counts,\n        obs=adata_meta_obs.collect().to_pandas(),\n        var=var_df,\n        dtype=count_dtype\n    )\n\n\ndef read_multiome():\n    print(\"Reading baseline multiome data\", flush=True)\n    with pl.StringCache():\n        multiome = pl.scan_parquet(input_dir / 'multiome_train.parquet', low_memory=True).cast({\n            'obs_id': pl.Categorical,\n            'location': pl.Categorical,  # This is a feature ID. If the feature_type in multiome_var_meta.csv is Gene Expression then this is a gene symbol. If feature_type is Peaks, then this is the genomic interval of the peak.\n            'count': pl.UInt16,\n            'normalized_count': pl.Float32 # If feature_type is Peaks, then this is ATAC-seq peak counts transformed with TF-IDF using the default log(TF) * log(IDF), else library size normalized\n        }).select(['obs_id', 'location', 'count'])#, 'normalized_count'])\n        multiome_meta_obs = pl.scan_csv(input_dir / 'multiome_obs_meta.csv', low_memory=True, dtypes={\n            'donor_id': pl.Categorical,\n            'cell_type': pl.Categorical,\n            'obs_id': pl.Categorical,\n        }).select(['donor_id', 'cell_type', 'obs_id'])\n        multiome_meta_var = pl.scan_csv(input_dir / \"multiome_var_meta.csv\", low_memory=True, dtypes={\n            'location': pl.Categorical,  # This is a feature ID. If the feature_type is Gene Expression then this is a gene symbol. If feature_type is Peaks, then this is the genomic interval of the peak.\n            'gene_id': pl.Categorical,  # ENSEMBL gene id or peaks\n            'feature_type': pl.Categorical,  # Gene Expression or Peaks\n            'genome': pl.Categorical,  # GRCh38\n            'interval': pl.Categorical,  # If feature_type is Peaks, then this is the genomic interval of the peak. else its the genomic interval of the gene\n        }).select(['location', 'feature_type', 'interval'])\n\n        # Now we must split the expression data from the ATAC data\n        loc2feature_type = multiome_meta_var.select(['location', 'feature_type']).collect(streaming=True).to_pandas()\n        loc2feature_type = {k: v for (k,v) in zip(loc2feature_type.location.values.tolist(), loc2feature_type.feature_type.values.tolist())}\n\n        multiome_expr = multiome.filter(pl.col('location').map_dict(loc2feature_type, return_dtype=pl.Categorical).cast(pl.Categorical) == 'Gene Expression')\n        multiome_atac = multiome.filter(pl.col('location').map_dict(loc2feature_type, return_dtype=pl.Categorical).cast(pl.Categorical) == 'Peaks')\n\n        # Now we must split the expression metadata from the ATAC metadata\n\n        multiome_meta_var_expr = multiome_meta_var.filter(pl.col('feature_type') == 'Gene Expression')\n        multiome_meta_var_atac = multiome_meta_var.filter(pl.col('feature_type') == 'Peaks')\n\n        # Build the two AnnDatas\n        print(\"Building ATAC adata\", flush=True)\n        atac_adata = build_anndata(multiome_atac, multiome_meta_obs, multiome_meta_var_atac, gene_id_col=\"location\")#, count_col='normalized_count', count_dtype=np.float32)\n        gc.collect()\n        print(\"Building expression adata\", flush=True)\n        expression_adata = build_anndata(multiome_expr, multiome_meta_obs, multiome_meta_var_expr, gene_id_col=\"location\")\n        gc.collect()\n\n        # Now merge them\n        return MuData({\n            'rna': expression_adata,\n            'atac': atac_adata\n        })\n\n\nif not (output_dir / 'adata_prepared_train.h5ad').exists():\n    adata_train = read_competition()\n    gc.collect()\n    adata_train.write_h5ad(output_dir / 'adata_prepared_train.h5ad')\nelse:\n    adata_train = sc.read_h5ad(output_dir / 'adata_prepared_train.h5ad')\n\nif USE_ATAC:\n    if not (output_dir / 'adata_prepared_multiome.h5mu').exists():\n        adata_multiome = read_multiome()\n        gc.collect()\n        adata_multiome.write_h5mu(output_dir / 'adata_prepared_multiome.h5mu')\n    else:\n        adata_multiome = mu.read_h5mu(output_dir / 'adata_prepared_multiome.h5mu')\n    rna_genes = list(adata_train.var['gene'])\n    atac_genes = list(adata_multiome['rna'].var['location'])\n    rearranged_subset = [(atac_genes.index(g) if g in atac_genes else 0) for g in rna_genes]\n    missing_genes = [i for i,g in enumerate(rna_genes) if g not in atac_genes]\n    # print(\",\".join([rna_genes[i] for i in missing_genes]))\n    print(\"Missing\", len(missing_genes), \"genes in ATAC, filling with 0\")\n    adata_multiome = MuData({\n        'rna': adata_multiome['rna'][:, rearranged_subset],\n        'atac': adata_multiome['atac']\n    })\n    X = adata_multiome['rna'].X.tolil()\n    X[:, missing_genes] = 0\n    adata_multiome['rna'].X = X.tocsr()\n    del X\n    adata_multiome['rna'].var['location'] = adata_train.var['gene']\n    adata_multiome['rna'].var.index = adata_train.var.index\n    adata_multiome.update()\n    adata_multiome.var_names_make_unique()\n    adata_multiome.obs['size_factors'] = adata_multiome['rna'].X.astype(np.float32).sum(-1)\n    adata_multiome.update()\n    gc.collect()\n\n# Filter\n# hvgs = 5_000\n# sc.pp.filter_cells(adata_train, min_counts=100)\n# sc.pp.highly_variable_genes(adata_train, flavor='seurat_v3', n_top_genes=hvgs)\n# adata_train = adata_train[:, adata_train.var.highly_variable]\nadata_train.obs['size_factors'] = adata_train.X.astype(np.float32).sum(-1)\nadata_train.obs['negative_control'] = adata_train.obs.sm_name == 'Dimethyl Sulfoxide'\nadata_train.obs['positive_control'] = (adata_train.obs.control) & (~adata_train.obs.negative_control)\n","metadata":{"execution":{"iopub.status.busy":"2023-12-08T03:04:21.689288Z","iopub.execute_input":"2023-12-08T03:04:21.690209Z","iopub.status.idle":"2023-12-08T03:04:26.207731Z","shell.execute_reply.started":"2023-12-08T03:04:21.690176Z","shell.execute_reply":"2023-12-08T03:04:26.206878Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import scanpy as sc\n# adata_train.layers['orig'] = adata_train.X\n# adata_train.X = adata_train.X.astype(np.float32)\n# sc.pp.log1p(adata_train)\n# sc.tl.pca(adata_train, svd_solver='arpack')\n# sc.pp.neighbors(adata_train, n_neighbors=10, n_pcs=40)\n# sc.tl.umap(adata_train)\n# sc.pl.umap(adata_train, color=['cell_type', 'sm_name', 'donor_id'])\n# adata_train.X = adata_train.layers['orig']\nsc.pp.highly_variable_genes(adata_train, flavor='seurat_v3', n_top_genes=5_000)\nhvg_mask = adata_train.var.highly_variable.to_numpy()\nn_hvgs = int(hvg_mask.sum())","metadata":{"execution":{"iopub.status.busy":"2023-12-08T03:04:26.208901Z","iopub.execute_input":"2023-12-08T03:04:26.209183Z","iopub.status.idle":"2023-12-08T03:04:49.726205Z","shell.execute_reply.started":"2023-12-08T03:04:26.209158Z","shell.execute_reply":"2023-12-08T03:04:49.725367Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(adata_train.n_obs)\nadata_train.obs.sm_name.value_counts()\n#adata_multiome","metadata":{"execution":{"iopub.status.busy":"2023-12-08T03:04:49.727349Z","iopub.execute_input":"2023-12-08T03:04:49.727649Z","iopub.status.idle":"2023-12-08T03:04:49.744296Z","shell.execute_reply.started":"2023-12-08T03:04:49.727622Z","shell.execute_reply":"2023-12-08T03:04:49.743496Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"potential_perturbations = list(zip(adata_train.obs.sm_name.unique().tolist(), adata_train.obs.SMILES.unique().tolist()))\n# Create embeddings for each of the compounds\n\ndef _molformer_embeddings(perturbations):\n    model = AutoModel.from_pretrained(\"ibm/MoLFormer-XL-both-10pct\", deterministic_eval=True, trust_remote_code=True)\n    tokenizer = AutoTokenizer.from_pretrained(\"ibm/MoLFormer-XL-both-10pct\", trust_remote_code=True)\n    inputs = tokenizer(potential_perturbations, padding=True, return_tensors=\"pt\")\n    outputs = model(**inputs)\n    return outputs.pooler_output\n\ndef _mean_pooling(model_output, attention_mask):\n    token_embeddings = model_output[0] #First element of model_output contains all token embeddings\n    input_mask_expanded = attention_mask.unsqueeze(-1).expand(token_embeddings.size()).float()\n    sum_embeddings = torch.sum(token_embeddings * input_mask_expanded, 1)\n    sum_mask = torch.clamp(input_mask_expanded.sum(1), min=1e-9)\n    return sum_embeddings / sum_mask\n\n    \ndef _chemberta_embeddings(perturbations, dataset):\n    tokenizer = AutoTokenizer.from_pretrained(f\"seyonec/{dataset}\")  # (\"seyonec/ChemBERTa-zinc-base-v1\")\n    model = AutoModelForMaskedLM.from_pretrained(f\"seyonec/{dataset}\")  # (\"seyonec/ChemBERTa-zinc-base-v1\")\n    inputs = tokenizer(potential_perturbations, padding=True, return_tensors=\"pt\")\n    outputs = model(**inputs)\n    # Ref: https://github.com/seyonechithrananda/bert-loves-chemistry/blob/master/chemberta/visualization/viz_utils.py#L107\n    outputs = _mean_pooling(outputs, inputs['attention_mask']) \n    return outputs\n\n\ndef produce_embeddings(perturbations, model='chemberta'):\n    pert_names = [p[0] for p in perturbations]\n    perturbations = [p[1] for p in perturbations]\n    with torch.no_grad():\n        if model == 'molformer':\n            embeds = _molformer_embeddings(perturbations)\n        else:\n            dataset = \"PubChem10M_SMILES_BPE_450k\"\n            if '-' in model:\n                dataset = model.split(\"-\")[-1]\n                if dataset == 'pubchem':\n                    dataset = \"PubChem10M_SMILES_BPE_450k\"\n                elif dataset == 'zinc':\n                    dataset = \"ChemBERTa-zinc-base-v1\"\n            embeds = _chemberta_embeddings(perturbations, dataset)\n    return {k: v for k,v in zip(pert_names, embeds)} \n    \n\noverwrite = True\nif not (output_dir / 'compound_embeddings.pkl').exists() or overwrite:\n    embeddings = produce_embeddings(potential_perturbations, EMBEDDING_MODEL)\n    torch.save(embeddings, output_dir / \"compound_embeddings.pkl\")\nelse:\n    embeddings = torch.load(output_dir / 'compound_embeddings.pkl')\n    \n#Standardize the embeddings\n# mean = torch.stack(list(embeddings.values())).mean(0)\n# std = torch.stack(list(embeddings.values())).std(0)\n# embeddings = {k: (v - mean) / std for k,v in embeddings.items()}\n# del mean\n# del std\n# embeddings\n# Normalize with respect to DMSO since it has the null effect\ndmso = embeddings['Dimethyl Sulfoxide']\ncondition_embedding_dim = dmso.shape[-1]\nembeddings = {k: (v - dmso) / dmso for (k, v) in embeddings.items()}\ndel dmso\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-12-08T03:04:49.745870Z","iopub.execute_input":"2023-12-08T03:04:49.746378Z","iopub.status.idle":"2023-12-08T03:05:07.868169Z","shell.execute_reply.started":"2023-12-08T03:04:49.746341Z","shell.execute_reply":"2023-12-08T03:05:07.867256Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Ref: https://discourse.scverse.org/t/annloader-for-mudata/1323/2\ndef sparse_csr_to_tensor(csr: scipy.sparse.csr_matrix):\n    \"\"\"\n    Transform scipy csr matrix to pytorch sparse tensor\n    \"\"\"\n    values = csr.data\n    indices = np.vstack(csr.nonzero())\n    shape = csr.shape\n\n    i = torch.LongTensor(indices)\n    v = torch.FloatTensor(values)\n    s = torch.Size(shape)\n\n    return torch.sparse.FloatTensor(i, v, s).to_dense()\n\ndef sparse_batch_collate(batch:list):\n    \"\"\"\n    Collate function to transform anndata csr view to pytorch sparse tensor\n    \"\"\"\n    if type(batch[0]['atac'].X) == anndata._core.views.SparseCSRView:\n        atac_batch_X = sparse_csr_to_tensor(scipy.sparse.vstack([x['atac'].X for x in batch]))\n    else:\n        atac_batch_X = torch.FloatTensor(np.vstack([x['atac'].X for x in batch]))\n\n    if type(batch[0]['rna'].X) == anndata._core.views.SparseCSRView:\n        rna_batch_X = sparse_csr_to_tensor(scipy.sparse.vstack([x['rna'].X for x in batch]))\n    else:\n        rna_batch_X = torch.FloatTensor(np.vstack([x['rna'].X for x in batch]))\n    return {\n        'atac': {\n            'X': atac_batch_X\n        },\n        'rna': {\n            'X': rna_batch_X\n        },\n        'donor_id': torch.FloatTensor(np.vstack([x['rna'].obsm['donor_id_encoded'] for x in batch])),\n        'cell_type': torch.FloatTensor(np.vstack([x['rna'].obsm['cell_type_encoded'] for x in batch])),\n        'size_factors': torch.FloatTensor(np.vstack([x.obs['size_factors'] for x in batch]))\n    }\n\n\ndef one_hot_encoder(adata, obs):\n    encoder = OneHotEncoder(sparse_output=False, dtype=np.float32).fit(adata.obs[obs].unique().to_numpy().reshape(-1, 1)).transform\n    return lambda x: encoder(x.to_numpy()[:, None])\n\n\ndef convert(mudata: MuData) -> MuData:\n    mudata['rna'].obsm['donor_id_encoded'] = one_hot_encoder(mudata['rna'], 'donor_id')(mudata['rna'].obs['donor_id'])\n    mudata['rna'].obsm['cell_type_encoded'] = one_hot_encoder(mudata['rna'], 'cell_type')(mudata['rna'].obs['cell_type'])\n    mudata['rna'].X = mudata['rna'].X.astype(np.float32)\n    mudata['atac'].X = mudata['atac'].X.astype(np.float32)\n    mudata.update()\n    return mudata\n\n\ndef MuDataLoader(mudata: MuData, batch_size: int):\n    mudata.update()\n    loader = DataLoader(\n        convert(mudata),\n        batch_size=batch_size,\n        shuffle=True,\n        collate_fn=sparse_batch_collate,\n\n    )\n    return loader","metadata":{"execution":{"iopub.status.busy":"2023-12-08T03:05:07.869436Z","iopub.execute_input":"2023-12-08T03:05:07.869794Z","iopub.status.idle":"2023-12-08T03:05:07.886223Z","shell.execute_reply.started":"2023-12-08T03:05:07.869765Z","shell.execute_reply":"2023-12-08T03:05:07.885304Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"adata_train.obs['sm_index'] = LabelEncoder().fit_transform(adata_train.obs.sm_name)\nadata_train.obs['cell_index'] = np.arange(0, adata_train.n_obs)\n\nunique_celltypes = adata_train.obs['cell_type'].unique()\nunique_perturbs = adata_train.obs['sm_name'].unique()\n\ndef csr_to_tensor(x):\n    x = x.tocoo()\n    return torch.sparse.FloatTensor(\n        torch.LongTensor([x.row.tolist(), x.col.tolist()]),\n        torch.FloatTensor(x.data.astype(np.float32))\n    ) # .to_dense()\n\nencoders = {\n    'obs': {\n        'cell_type': one_hot_encoder(adata_train, 'cell_type'),\n        'sm_name': lambda sm_name: torch.stack([embeddings[n] for n in sm_name.astype(str).to_numpy()]),\n        # Convert to 0/1\n        'negative_control': lambda c: torch.tensor(c.to_numpy()).float(),\n        'positive_control': lambda c: torch.tensor(c.to_numpy()).float(),\n        'plate_name': one_hot_encoder(adata_train, 'plate_name'),\n        'well': one_hot_encoder(adata_train, 'well'),\n        'donor_id': one_hot_encoder(adata_train, 'donor_id')\n    },\n    'X': lambda x: x.astype(np.float32)\n}\n\ndef make_dataloaders(batch_size: int = 64):\n    drugs = adata_train.obs.sm_name.unique()\n    # Ignore DMSO and positive controls\n    drugs = drugs[(drugs != 'Dimethyl Sulfoxide') & (drugs != 'dabrfenib') & (drugs != 'belinostat')]\n    celltypes = adata_train.obs.cell_type.unique()\n    n_cells = adata_train.n_obs\n    drug_celltype_pairs = [(d, c) for d in drugs for c in celltypes]\n    # Randomly select 10 drug/cell type pairs for validation and test\n    val_drug_celltype_pairs = np.random.choice(np.arange(len(drug_celltype_pairs)), size=5, replace=False)\n    test_drug_celltype_pairs = np.random.choice([p for p in range(len(drug_celltype_pairs)) if p not in val_drug_celltype_pairs], size=5, replace=False)\n\n    val_drug_celltype_pairs = [drug_celltype_pairs[i] for i in val_drug_celltype_pairs]\n    test_drug_celltype_pairs = [drug_celltype_pairs[i] for i in test_drug_celltype_pairs]\n\n    print(\"Validation pairs:\", \", \".join([f\"{d}/{c}\" for d,c in val_drug_celltype_pairs]))\n    print(\"Test pairs:\", \", \".join([f\"{d}/{c}\" for d,c in test_drug_celltype_pairs]))\n\n    val_pair_df = pd.DataFrame(val_drug_celltype_pairs, columns=['sm_name', 'cell_type'])\n    test_pair_df = pd.DataFrame(test_drug_celltype_pairs, columns=['sm_name', 'cell_type'])\n\n    adata_val = adata_train[adata_train.obs.sm_name.isin(val_pair_df.sm_name) & adata_train.obs.cell_type.isin(val_pair_df.cell_type)]\n    adata_test = adata_train[adata_train.obs.sm_name.isin(test_pair_df.sm_name) & adata_train.obs.cell_type.isin(test_pair_df.cell_type)]\n    if HOLD_OUT:\n        final_adata_train = adata_train[(~adata_train.obs.cell_index.isin(adata_val.obs.cell_index)) & (~adata_train.obs.cell_index.isin(adata_test.obs.cell_index))]\n    else:\n        final_adata_train = adata_train\n#     drugs = adata_train.obs.sm_name.unique()\n#     # Ignore DMSO and positive controls\n#     celltypes = adata_train.obs.cell_type.unique()\n#     n_cells = adata_train.n_obs\n#     val_cells = test_cells = int(.1 * n_cells)\n#     val_cells = np.random.choice(np.arange(n_cells), size=val_cells, replace=False)\n#     test_cells = np.random.choice([p for p in range(n_cells) if p not in val_cells], size=test_cells, replace=False)\n\n#     adata_val = adata_train[adata_train.obs.cell_index.isin(val_cells)]\n#     adata_test = adata_train[adata_train.obs.cell_index.isin(test_cells)]\n#     if HOLD_OUT:\n#         final_adata_train = adata_train[(~adata_train.obs.cell_index.isin(adata_val.obs.cell_index)) & (~adata_train.obs.cell_index.isin(adata_test.obs.cell_index))]\n#     else:\n#         final_adata_train = adata_train\n    dataloader = AnnLoader(final_adata_train, batch_size=batch_size, shuffle=True, convert=encoders)\n    val_dataloader = AnnLoader(adata_val, batch_size=batch_size, shuffle=False, convert=encoders)\n    test_dataloader = AnnLoader(adata_test, batch_size=max(1, batch_size//128), shuffle=False, convert=encoders)\n    return dataloader, val_dataloader, test_dataloader\n\ndef train(model, guide, epochs=100, lr=5e-3, dataloader=None, val_dataloader=None, test_dataloader=None, reset=True):\n    if reset:\n        pyro.clear_param_store()\n        pyro.set_rng_seed(0)\n\n    if not dataloader:\n        dataloader, val_dataloader, test_dataloader = make_dataloaders()\n    elif isinstance(dataloader, tuple):\n        dataloader, val_dataloader, test_dataloader = dataloader\n\n    optimizer = pyro.optim.ClippedAdam({'lr': lr})\n\n    svi = pyro.infer.SVI(model, guide, optimizer, loss=pyro.infer.Trace_ELBO())\n    n_cells = dataloader.dataset.n_obs\n    n_genes = dataloader.dataset.n_vars\n    last_loss = np.nan\n    last_val_loss = np.nan\n    loss_values = []\n    loss_step_values = []\n    val_loss_values = []\n    with pyro.poutine.scale(scale=1):#/(n_genes/2)):\n        with tqdm(total=epochs, desc=\"Training....\") as pbar:\n            for epoch in range(epochs):\n                epoch_loss = 0\n                for i, batch in enumerate(dataloader):\n                    batch_n = batch['rna']['X'].shape[0] if isinstance(batch, dict) else batch.shape[0]\n                    step_loss = svi.step(batch) / batch_n  # min(svi.step(batch), 1e5)\n                    if epoch > 0 or i > 25:  # Ignore the first 25 minibatches because the loss is ridiculous and meaningless\n                        epoch_loss += step_loss * (batch_n/dataloader.dataset.n_obs)\n                        loss_step_values.append(step_loss)\n                    pbar.set_description(f\"Training {batch_n*(i+1)}/{n_cells}... [step_loss={step_loss:.4f}] [curr_epoch_loss={epoch_loss:.4f}] [last_epoch_loss={last_loss:.4f}] [last_val_loss={last_val_loss:.4f}]\", refresh=True)\n\n                if val_dataloader:\n                    val_loss = 0.\n                    pbar.set_description(f\"Running Validation Loop...\", refresh=True)\n                    for i, batch in enumerate(val_dataloader):\n                        batch_n = batch['rna']['X'].shape[0] if isinstance(batch, dict) else batch.shape[0]\n                        val_loss_step = svi.evaluate_loss(batch) / batch_n #min(svi.evaluate_loss(batch), 1e4) / val_dataloader.dataset.n_obs\n                        val_loss += val_loss_step * (batch_n/val_dataloader.dataset.n_obs)\n                        pbar.set_description(f\"Running Validation Loop {batch_n*(i+1)}/{val_dataloader.dataset.n_obs}... [val_loss={val_loss_step:.4f}]\", refresh=True)\n                    val_loss_values.append(val_loss)\n                    last_val_loss = val_loss\n\n                pbar.update(1)\n                last_loss = epoch_loss\n                loss_values.append(min(epoch_loss, 1e7))\n\n    if val_dataloader:\n        # Plot the validation loss with the training loss\n        sns.lineplot(x=range(epochs), y=loss_values, label=\"Training Loss\")\n        sns.lineplot(x=range(epochs), y=val_loss_values, label=\"Validation Loss\")\n        plt.xlabel(\"Epoch\")\n        plt.ylabel(\"Loss\")\n        plt.title(\"ELBO Loss per Epoch\")\n        plt.show()\n        plt.clf()\n        # Plot the validation loss with the training loss\n        sns.lineplot(x=range(epochs), y=loss_values, label=\"Training Loss\")\n        sns.lineplot(x=range(epochs), y=val_loss_values, label=\"Validation Loss\")\n        plt.xlabel(\"Epoch\")\n        plt.ylabel(\"Loss\")\n        plt.yscale(\"log\")\n        plt.title(\"ELBO Loss per Epoch\")\n        plt.show()\n        plt.clf()\n    else:\n        sns.lineplot(x=range(epochs), y=loss_values)\n        plt.xlabel(\"Epoch\")\n        plt.ylabel(\"Loss\")\n        plt.title(\"ELBO Loss per Epoch\")\n        plt.show()\n        plt.clf()\n        sns.lineplot(x=range(epochs), y=loss_values)\n        plt.xlabel(\"Epoch\")\n        plt.ylabel(\"Loss\")\n        plt.yscale(\"log\")\n        plt.title(\"ELBO Loss per Epoch\")\n        plt.show()\n        plt.clf()\n\n    sns.lineplot(x=range(len(loss_step_values)), y=loss_step_values)\n    plt.xlabel(\"Training Step\")\n    plt.ylabel(\"Loss\")\n    plt.yscale(\"log\")\n    plt.title(\"ELBO Loss per Training Step\")\n    plt.show()\n    plt.clf()\n    sns.lineplot(x=range(len(loss_step_values)), y=loss_step_values)\n    plt.xlabel(\"Training Step\")\n    plt.ylabel(\"Loss\")\n    plt.title(\"ELBO Loss per Training Step\")\n    plt.show()\n    plt.clf()\n\n    if test_dataloader:\n        print(\"Posterior predictive check on test data\")\n        del dataloader, val_dataloader, val_loss_values\n        gc.collect()\n\n        # We will examine the distribution for the cell expressions for each gene\n        predictive = pyro.infer.Predictive(model, guide=guide, num_samples=100, return_sites=(\"x\",))\n        differences = np.zeros((test_dataloader.dataset.n_obs, n_hvgs))\n        # percentiles = np.zeros((test_dataloader.dataset.n_obs, n_hvgs))\n        for i, batch in tqdm(enumerate(test_dataloader), desc='Analyzing...'):\n            batch_n = batch['rna']['X'].shape[0] if isinstance(batch, dict) else batch.shape[0]\n            predictive_posterior = predictive(batch, observe_x=False)\n            x_true = batch['rna']['X'].detach().cpu() if isinstance(batch, dict) else batch.X.detach().cpu() # (batch_n, n_genes)\n            x_pred = predictive_posterior['x'].detach().cpu()\n            differences[i*batch_n:(i+1)*batch_n] = (x_pred - x_true).mean(dim=0).numpy()[:, hvg_mask]  # (n_samples, batch_n, n_genes)\n            # percentiles[i*batch_n:(i+1)*batch_n] = ((x_pred > x_true).float().numpy()[:, hvg_mask].median(axis=0))\n        import scipy\n        print(f\"Summary: {scipy.stats.describe(differences.flatten())}\")\n        # Plot the distribution of the percentiles\n        sns.histplot(differences.flatten(), bins=100)\n        plt.axvline(x=0, color='red')\n        plt.yscale(\"log\")\n        plt.xlabel(\"Difference From Ground Truth\")\n        plt.ylabel(\"Count\")\n        plt.title(\"Deviation from True Expression\")\n        plt.show()\n        plt.clf()\n        sns.histplot(differences.flatten(), bins=100)\n        plt.axvline(x=0, color='red')\n        plt.xlabel(\"Difference From Ground Truth\")\n        plt.ylabel(\"Count\")\n        plt.title(\"Deviation from True Expression\")\n        plt.show()\n        plt.clf()\n        # sns.histplot(((percentiles.flatten()-0.5)*2), bins=100)\n        # plt.axvline(x=0, color='red')\n        # plt.yscale(\"log\")\n        # plt.xlabel(\"Percentage Difference From Ground Truth\")\n        # plt.ylabel(\"Count\")\n        # plt.title(\"Deviation from True Expression\")\n        # plt.show()\n        # plt.clf()\n        # sns.histplot(((percentiles.flatten()-0.5)*2), bins=100)\n        # plt.axvline(x=0, color='red')\n        # plt.xlabel(\"Percentage Difference From Ground Truth\")\n        # plt.ylabel(\"Count\")\n        # plt.title(\"Deviation from True Expression\")\n        # plt.show()\n        # plt.clf()\n        # import sys\n        # sys.setrecursionlimit(100000)  # Large data recurses too much\n        # # Heatmap of the percentiles\n        # cg = sns.clustermap(((percentiles.flatten()-0.5)*2), cmap='viridis', vmin=-1, vmax=1)\n        # cg.ax_row_dendrogram.set_visible(False) #suppress row dendrogram\n        # cg.ax_col_dendrogram.set_visible(False) #suppress column dendrogram\n        # plt.xlabel(\"Gene\")\n        # plt.ylabel(\"Cell\")\n        # plt.title(\"Mean Percentage Difference from True Expression\")\n        # plt.show()\n        # plt.clf()\n\n\n    return loss_values, loss_step_values\n","metadata":{"execution":{"iopub.status.busy":"2023-12-08T03:05:07.887586Z","iopub.execute_input":"2023-12-08T03:05:07.887902Z","iopub.status.idle":"2023-12-08T03:05:08.000449Z","shell.execute_reply.started":"2023-12-08T03:05:07.887876Z","shell.execute_reply":"2023-12-08T03:05:07.999704Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class VAEModule(nn.Module):\n    \"\"\"\n    A simple VAE module to be used with Pyro.\n    \"\"\"\n\n    def __init__(self, name: str, n_obs_in: int, n_latent: int, n_obs_out: int = None, n_hidden: tuple[int] = None, hidden_act: nn.Module = nn.ReLU(), dropout: float = 0.1):\n        super().__init__()\n        self.name = name\n        self.n_obs_in = n_obs_in\n        if not n_obs_out:\n            n_obs_out = n_obs_in\n        self.n_obs_out = n_obs_out\n        self.n_latent = n_latent\n\n        if n_hidden is None:\n            n_hidden_in = (n_obs_in // 2, n_latent * 2, n_latent)\n            n_hidden_out = (n_latent, n_latent * 2, n_obs_out // 2)\n        else:\n            n_hidden_in = n_hidden\n            n_hidden_out = n_hidden[::-1]\n\n        in_dims = [n_obs_in] + list(n_hidden_in) + [n_latent]\n        out_dims = [n_latent] + list(n_hidden_out) + [n_obs_out]\n        encoder = []\n        for i in range(len(in_dims) - 1):\n            encoder.append(nn.Linear(in_dims[i], in_dims[i+1]))\n            if i < len(in_dims) - 2:\n                encoder.append(hidden_act)\n                encoder.append(nn.Dropout(dropout))\n        self.encoder = nn.Sequential(*encoder)\n        self.mu_encoder = nn.Linear(n_latent, n_latent)\n        self.logvar_encoder = nn.Linear(n_latent, n_latent)\n\n        decoder = []\n        for i in range(len(out_dims) - 1):\n            decoder.append(nn.Linear(out_dims[i], out_dims[i+1]))\n            if i < len(out_dims) - 2:\n                decoder.append(hidden_act)\n                decoder.append(nn.Dropout(dropout))\n        self.decoder = nn.Sequential(*decoder)\n        self.dummy = nn.Parameter(torch.tensor([1/2]))\n\n    @property\n    def device(self):\n        return self.dummy.device\n\n    def forward_encoder(self, x):\n        encoded = self.encoder(x)\n        mu = self.mu_encoder(encoded)\n        lvar = self.logvar_encoder(encoded)\n        return mu, lvar\n\n    def sample_latent(self, x=None, expand_dims=None):\n        if x is not None:\n            mu, lvar = self.forward_encoder(x)\n            norm = dist.Normal(mu, lvar.mul(0.5).exp())\n        else:\n            norm = dist.Normal(torch.zeros(self.n_latent, device=self.device), torch.ones(self.n_latent, device=self.device))\n\n        if expand_dims:\n            norm = norm.expand_by(expand_dims)\n\n        return pyro.sample(f\"z_{self.name}\", norm.to_event(1))\n\n    def forward_decoder(self, z):\n        decoded = self.decoder(z)\n        return decoded\n\n# @torch.compile\ndef convert_to_indices(x_cat: torch.Tensor, x_subset_cat: torch.Tensor) -> torch.Tensor:\n    output = torch.ones_like(x_cat)\n    for i in range(x_subset_cat.shape[0]):\n        x_subset_val = x_subset_cat[..., i]\n        output = torch.where(x_cat == x_subset_val, torch.ones_like(output) * i, output)\n    return output\n\n# @torch.compile\ndef one_hot_encode(x: torch.Tensor, max_n: int) -> torch.Tensor:\n    return torch.eye(max_n, device=x.device)[x.int()]","metadata":{"execution":{"iopub.status.busy":"2023-12-08T03:05:08.001479Z","iopub.execute_input":"2023-12-08T03:05:08.001748Z","iopub.status.idle":"2023-12-08T03:05:08.019855Z","shell.execute_reply.started":"2023-12-08T03:05:08.001723Z","shell.execute_reply":"2023-12-08T03:05:08.018811Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class scRna2AtacModel(nn.Module):\n    \"\"\"\n    Simple VAE-based model to predict ATAC from RNA. \n    The intention of this model is to pre-train a rich latent encoder for scRNA-seq data.\n    \"\"\"\n    \n    def __init__(self, n_cells: int, n_genes: int, n_locations: int, n_latent: int, n_donors: int, n_celltypes: int, dropout:float = 0.01):\n        super().__init__()\n        self.n_cells = n_cells\n        self.n_genes = n_genes\n        self.n_locations = n_locations\n        self.n_latent = n_latent\n        self.n_donors = n_donors\n        self.n_celltypes = n_celltypes\n        \n        # Learn batch effects separately to not taint the cell state latent space\n        self.batch_vae = VAEModule(\"ATAC_batch_effects\", n_donors, n_latent, dropout=dropout)\n        # Cell state has cell type and expression piped in, but must reproduce all of those plus the ATAC\n        self.cell_state_vae = VAEModule(\"cell_state\", n_genes, n_latent, dropout=dropout)\n        self.celltype_decoder = nn.Sequential(\n            nn.Linear(n_latent, n_celltypes)\n        )\n        # ATAC detached decoders\n        base_atac_decoder = nn.Sequential(\n            nn.Linear(n_latent, n_latent),\n            nn.ReLU(),\n            nn.Dropout(dropout)\n        )\n        \n        self.atac_library_scaling_factor = nn.Sequential(\n            base_atac_decoder,\n            nn.Linear(n_latent, 1),\n            nn.Sigmoid()\n        )\n        \n        self.atac_decode_prob_of_accessibility = nn.Sequential(\n            base_atac_decoder,\n            nn.Linear(n_latent, n_locations)\n        )\n        self.batch2atac_shift = nn.Sequential(\n            nn.Linear(n_latent, n_locations)\n        )\n        \n        self.batch2rna_shift = nn.Sequential(\n            nn.Linear(n_latent, n_genes),\n        )\n        self.dummy = nn.Parameter(torch.tensor([1/2]))\n\n    @property\n    def device(self):\n        return self.dummy.device\n    \n    def model(self, mudata: MuData):\n        pyro.module(\"scRna2Atac\", self)\n        n_obs = mudata['rna']['X'].shape[0]\n        # Gene dispersion from the mean\n        theta = pyro.param(\"ATAC_theta\", lambda: torch.ones((self.n_genes,), device=self.device), constraint=constraints.positive)\n        # Scaling factor for regions\n        region_theta = pyro.param(\"ATAC_region_theta\", lambda: torch.ones((self.n_locations,), device=self.device)/self.n_locations, constraint=constraints.positive)\n        with pyro.plate(\"ATAC_cells1\", n_obs, dim=-1):\n            Zb = self.batch_vae.sample_latent(expand_dims=(n_obs,))\n            # Predict the batch labels\n            pred_batch = self.batch_vae.forward_decoder(Zb)\n            # Likelihood\n            pyro.sample(\"ATAC_donor\", \n                        dist.OneHotCategorical(logits=pred_batch).to_event(1),\n                        obs=mudata['donor_id'].to(self.device))\n            \n            # Cell state latent space\n            Z = self.cell_state_vae.sample_latent(expand_dims=(n_obs,))\n            x0 = self.cell_state_vae.forward_decoder(Z)\n            celltype_pred = self.celltype_decoder(Z)\n            pyro.sample(\"ATAC_celltype\", dist.OneHotCategorical(logits=celltype_pred).to_event(1), obs=mudata['cell_type'].to(self.device))\n            \n            library_size = pyro.deterministic(\"ATAC_library_size\", mudata['size_factors'].to(self.device))\n        \n        with pyro.plate(\"ATAC_cells2\", n_obs, dim=-2):\n            rna = mudata['rna']['X']\n            atac = mudata['atac']['X']\n            with pyro.plate(\"ATAC_rna\", rna.shape[1], dim=-1):\n                batch_rna_change = self.batch2rna_shift(Zb)\n                mu = F.softmax((x0 + batch_rna_change).exp2(), dim=-1)\n#                 mu = F.softmax(self.cell_state_vae.forward_decoder(Z), dim=-1)\n                # See https://github.com/pytorch/pytorch/issues/42449 for Negative Binomial parametrization\n                nb_logits = (library_size * mu + 1e-8).log() - (theta + 1e-8).log()\n                x_dist = dist.NegativeBinomial(total_count=theta, logits=nb_logits)\n                pyro.sample(\"ATAC_rna_x\", x_dist.to_event(0), obs=rna.to(self.device))\n            with pyro.plate(\"ATAC_atac\", atac.shape[1], dim=-1):\n                batch_atac_change = self.batch2atac_shift(Zb)\n                cell_atac_scaling = self.atac_library_scaling_factor(Z)\n                cell_probability = self.atac_decode_prob_of_accessibility(Z)\n                \n                atac_prob = region_theta * cell_atac_scaling * F.softmax((batch_atac_change + cell_probability).exp2(), dim=-1)\n                \n                pyro.sample('ATAC_atac_x', dist.Bernoulli(atac_prob).to_event(0), \n                            obs=(atac.to(self.device) > 0).float())\n\n    def guide(self, mudata: MuData):\n        pyro.module(\"scRna2Atac\", self)\n        celltype = mudata['cell_type'].to(self.device)\n        n_obs = mudata['rna']['X'].shape[0]\n        with pyro.plate(\"ATAC_cells1\", n_obs, dim=-1):\n            self.batch_vae.sample_latent(mudata['donor_id'].to(self.device))\n            self.cell_state_vae.sample_latent(mudata['rna']['X'].to(self.device).log1p())\n            \n        # with pyro.plate(\"cells2\", n_obs, dim=-2):\n        #     rna = mudata['rna']\n        #     atac = mudata['atac']\n        #     with pyro.plate(\"rna\", rna.n_var, dim=-1):\n        #         ...\n        #     with pyro.plate(\"atac\", atac.n_var, dim=-1):\n        #         ...\n        \n        ","metadata":{"execution":{"iopub.status.busy":"2023-12-08T03:05:08.021264Z","iopub.execute_input":"2023-12-08T03:05:08.021923Z","iopub.status.idle":"2023-12-08T03:05:08.045707Z","shell.execute_reply.started":"2023-12-08T03:05:08.021895Z","shell.execute_reply":"2023-12-08T03:05:08.044834Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if USE_ATAC:\n    atac_model = scRna2AtacModel(adata_multiome['rna'].n_obs, adata_multiome['rna'].n_vars, adata_multiome['atac'].n_vars, 256, adata_multiome['rna'].obs['donor_id'].nunique(), adata_multiome['rna'].obs['cell_type'].nunique())\n    if torch.cuda.is_available():\n        atac_model = atac_model.to('cuda')\n    train(atac_model.model, atac_model.guide, epochs=ATAC_PRETRAIN_EPOCHS, lr=2e-5, dataloader=MuDataLoader(adata_multiome, 512))\n    if torch.cuda.is_available():\n        atac_model = atac_model.to('cpu')\n    atac_latent_module = atac_model.cell_state_vae\n    del adata_multiome\n    gc.collect()\nelse:\n    print(\"SKIPPING ATAC PRE-TRAINING\")\n    atac_latent_module = None","metadata":{"execution":{"iopub.status.busy":"2023-12-08T03:05:08.046822Z","iopub.execute_input":"2023-12-08T03:05:08.047122Z","iopub.status.idle":"2023-12-08T03:05:08.059235Z","shell.execute_reply.started":"2023-12-08T03:05:08.047092Z","shell.execute_reply":"2023-12-08T03:05:08.058283Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.distributions.utils import logits_to_probs, probs_to_logits, clamp_probs\n\nclass scPerturbModel(nn.Module):\n\n    def __init__(self, n_cells: int, n_genes: int, n_conditions: int, condition_embedding_dim: int, n_latent: int, n_plates: int, n_donors: int, n_wells: int, n_celltypes: int, dropout: float = 0.01):\n        super().__init__()\n        self.n_cells = n_cells\n        self.n_genes = n_genes\n        self.n_conditions = n_conditions\n        self.condition_embedding_dim = condition_embedding_dim\n        self.n_latent = n_latent\n        self.n_plates = n_plates\n        self.n_donors = n_donors\n        self.n_wells = n_wells\n        self.n_celltypes = n_celltypes\n\n        self.batch_vae = VAEModule(\"batch_effects\", n_plates + n_donors + n_wells, n_latent, dropout=dropout)\n        self.cell_state_vae = VAEModule(\"cell_state\", n_genes, n_latent, dropout=dropout)\n        self.celltype_decoder = nn.Sequential(\n            nn.Linear(n_latent, n_celltypes)\n        )\n\n         # No bias so DMSO remains at all 0s\n        self.perturb_projection = nn.Linear(condition_embedding_dim, n_latent, bias=False)\n        self.perturb2change = nn.Sequential(\n            nn.ReLU(),\n            nn.Dropout(dropout),\n            nn.Linear(n_latent, n_genes*2)\n        )\n        self.perturb2gene_prior = nn.Sequential(\n            nn.ReLU(),\n            nn.Dropout(dropout),\n            nn.Linear(n_latent, n_genes*2),\n            nn.Softplus()\n        )\n        self.batch2shift = nn.Sequential(\n            nn.Linear(n_latent, n_genes),\n        )\n        self.dummy = nn.Parameter(torch.tensor([1/2]))\n\n        # An independent network for each cell type\n#         self.cell_type_parameterizers = nn.ModuleList([\n#             nn.Sequential(\n#                 nn.Linear(n_latent + condition_embedding_dim, n_latent*2),\n#                 nn.ReLU(),\n#                 nn.Dropout(dropout),\n#                 nn.Linear(n_latent*2, n_latent*4),\n#                 nn.ReLU(),\n#                 nn.Dropout(dropout),\n#                 nn.Linear(n_latent*4, n_genes)\n#             )\n#             for ct in range(n_celltypes)\n#         ])\n\n    @property\n    def device(self):\n        return self.dummy.device\n\n    def model(self, adata, observe_x=True):\n        pyro.module(\"scPerturb\", self)\n\n        # Gene dispersion from the mean\n        theta = pyro.param(\"theta\", lambda: torch.ones((self.n_genes,), device=self.device), constraint=constraints.positive)\n\n        sm_index = adata.obs['sm_index'].to(self.device)\n        present_perturbs = sm_index.unique()\n        with torch.no_grad():\n            all_perturbs = torch.stack(tuple(e for i,e in enumerate(embeddings.values()) if i in present_perturbs), dim=0).detach().to(self.device)\n        all_perturbs = self.perturb_projection(all_perturbs)\n\n        with pyro.plate(\"perturbations\", self.n_conditions, subsample=present_perturbs, dim=-2) as ind:\n            with pyro.plate(\"perturbations_x_genes\", self.n_genes, dim=-1):\n\n                alpha, beta = self.perturb2gene_prior(all_perturbs).chunk(2, dim=-1)\n                genewise_perturb_prior = pyro.sample(\"genewise_perturb_prior\",\n                                                    dist.Beta(alpha, beta).to_event(0))\n\n        with pyro.plate(\"cells1\", adata.n_obs, dim=-1):\n            # Get index mapping to our sampled perturbations\n            condition_idx = convert_to_indices(sm_index, present_perturbs)\n            condition_latents = self.perturb_projection(adata.obs['sm_name'].to(self.device))\n\n            negative_control = adata.obs['negative_control'].to(self.device).unsqueeze(-1).expand(-1, self.n_genes).bool()\n\n            # Batch effect latent space\n            Zb = self.batch_vae.sample_latent(expand_dims=(adata.n_obs,))\n            # Predict the batch labels\n            pred_batch = self.batch_vae.forward_decoder(Zb)\n            # Likelihoods\n            well = pyro.sample(\"well\",\n                               dist.OneHotCategorical(logits=pred_batch[..., :self.n_wells]).to_event(1),\n                               obs=adata.obs['well'].to(self.device))\n            donor = pyro.sample(\"donor\",\n                                dist.OneHotCategorical(logits=pred_batch[..., self.n_wells:self.n_wells+self.n_donors]).to_event(1),\n                                obs=adata.obs['donor_id'].to(self.device))\n            plate = pyro.sample(\"plate\",\n                                dist.OneHotCategorical(logits=pred_batch[..., self.n_wells+self.n_donors:]).to_event(1),\n                                obs=adata.obs['plate_name'].to(self.device))\n\n\n            # Cell state latent space\n            Z = self.cell_state_vae.sample_latent(expand_dims=(adata.n_obs,))\n\n            x0 = self.cell_state_vae.forward_decoder(Z)\n            celltype_pred = self.celltype_decoder(Z)\n            pyro.sample(\"celltype\", dist.OneHotCategorical(logits=celltype_pred).to_event(1), obs=adata.obs['cell_type'].to(self.device))\n\n            library_size = pyro.deterministic(\"library_size\", adata.obs['size_factors'].to(self.device).unsqueeze(-1))\n\n        with pyro.plate(\"cells2\", adata.n_obs, dim=-2):\n            with pyro.plate(\"cells_x_genes\", self.n_genes, dim=-1):\n                perturbation_mask = pyro.sample(\"perturbation_mask\",\n                                                dist.Bernoulli(logits=probs_to_logits(genewise_perturb_prior[condition_idx])).to_event(0))#,\n                                                #obs=(~negative_control).float(),\n                                                #obs_mask=negative_control)\n\n                mu, logvar = self.perturb2change(condition_latents).chunk(2, dim=-1)\n                perturb_change = pyro.sample(\"perturbation_scale\",\n                                             dist.Normal(mu, logvar.mul(0.5).exp()).to_event(0))#,\n                                             #obs=torch.zeros_like(mu),\n                                             #obs_mask=negative_control)\n\n                perturb_change = perturb_change * perturbation_mask.round()\n                batch_change = self.batch2shift(Zb)\n\n                mu = F.softmax((x0 + perturb_change + batch_change).exp2(), dim=-1).clamp(max=1e10)\n#                 mu = F.softmax(self.cell_state_vae.forward_decoder(Z), dim=-1)\n                # See https://github.com/pytorch/pytorch/issues/42449 for Negative Binomial parametrization\n                nb_logits = (library_size * mu).clamp(min=1e-7).log() - (theta).clamp(min=1e-7).log()\n                x_dist = dist.NegativeBinomial(total_count=theta, logits=nb_logits)\n                pyro.sample(\"x\", x_dist.to_event(0), obs=adata.X.to(self.device) if observe_x else None)\n\n\n    def guide(self, adata, observe_x=True):\n        pyro.module(\"scPerturb\", self)\n\n        # MAP ESTIMATES\n        genewise_perturb_prior_alpha_MAP = pyro.param(\"genewise_perturb_prior_alpha_MAP\", lambda: torch.ones((self.n_conditions, self.n_genes), device=self.device)/2, constraint=constraints.positive)\n        genewise_perturb_prior_beta_MAP = pyro.param(\"genewise_perturb_prior_beta_MAP\", lambda: torch.ones((self.n_conditions, self.n_genes), device=self.device)/2, constraint=constraints.positive)\n        perturbation_mask_MAP = pyro.param(\"perturbation_mask_MAP\", lambda: torch.ones((self.n_conditions, self.n_celltypes, self.n_genes), device=self.device)/2, constraint=constraints.unit_interval)\n        perturbation_scale_MAP = pyro.param(\"perturbation_scale_MAP\", lambda: torch.zeros((self.n_conditions, self.n_celltypes, self.n_genes), device=self.device))\n\n        # USES WAY TOO MUCH MEMORY\n#         perturbation_mask_MAP = pyro.param(\"perturbation_mask_MAP\", lambda: torch.ones((self.n_cells, self.n_genes), device='cpu')/2, constraint=constraints.unit_interval)\n\n        sm_index = adata.obs['sm_index'].to(self.device)\n        present_perturbs = sm_index.unique()\n\n        with pyro.plate(\"perturbations\", self.n_conditions, subsample=present_perturbs, dim=-2):\n            with pyro.plate(\"perturbations_x_genes\", self.n_genes, dim=-1):\n                alpha = genewise_perturb_prior_alpha_MAP[present_perturbs]\n                beta = genewise_perturb_prior_beta_MAP[present_perturbs]\n                perturb_prior = pyro.sample(\"genewise_perturb_prior\",\n                                            dist.Beta(alpha, beta, validate_args=False).to_event(0))\n\n        celltype = adata.obs['cell_type'].to(self.device)\n        with pyro.plate(\"cells1\", adata.n_obs, dim=-1):\n            # Encode the batch effects\n            Zb = self.batch_vae.sample_latent(torch.cat([adata.obs['well'].to(self.device), adata.obs['donor_id'].to(self.device), adata.obs['plate_name'].to(self.device)], dim=1))\n\n            # Encode the cell states\n            Z = self.cell_state_vae.sample_latent(adata.X.to(self.device).log1p())\n\n            condition_idx = convert_to_indices(sm_index, present_perturbs)\n\n        with pyro.plate(\"cells2\", adata.n_obs, dim=-2):\n            with pyro.plate(\"cells_x_genes\", self.n_genes, dim=-1):\n                pyro.sample(\"perturbation_mask\", dist.Bernoulli(logits=probs_to_logits(perturbation_mask_MAP[sm_index, celltype.argmax(-1)])).to_event(0))#dist.Bernoulli(perturbation_mask_MAP[adata.obs['cell_index'].cpu(), :].to(self.device)).to_event(1))\n                pyro.sample(\"perturbation_scale\", dist.Delta(perturbation_scale_MAP[sm_index, celltype.argmax(-1)]).to_event(0))\n\n    def save(self):\n        torch.save(self.state_dict(), \"model.pt\")\n        pyro.get_param_store().save(\"model_params.pt\")\n\n    def load(self):\n        saved_model = torch.load(\"model.pt\")\n        self.load_state_dict(saved_model)\n        pyro.get_param_store().load(\"model_params.pt\")","metadata":{"execution":{"iopub.status.busy":"2023-12-08T03:05:08.063275Z","iopub.execute_input":"2023-12-08T03:05:08.063597Z","iopub.status.idle":"2023-12-08T03:05:08.103997Z","shell.execute_reply.started":"2023-12-08T03:05:08.063561Z","shell.execute_reply":"2023-12-08T03:05:08.103213Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# gc.collect()\n# model = scPerturbModel(adata_train.n_obs, adata_train.n_vars, len(potential_perturbations), condition_embedding_dim, 4, adata_train.obs.plate_name.nunique(), adata_train.obs.donor_id.nunique(), adata_train.obs.well.nunique(), adata_train.obs.cell_type.nunique())\n# # Plot the plate diagram\n# # pyro.render_model(model.model, (make_dataloader(1).dataset[0],), render_distributions=True, render_params=False)\n# del model\n# gc.collect()\nmodel = scPerturbModel(adata_train.n_obs, adata_train.n_vars, len(potential_perturbations), condition_embedding_dim, 32, adata_train.obs.plate_name.nunique(), adata_train.obs.donor_id.nunique(), adata_train.obs.well.nunique(), adata_train.obs.cell_type.nunique())\n\nif atac_latent_module is not None:\n    model.cell_state_vae = atac_latent_module\n    # Remove all ATAC params to save some memory\n    del atac_model\n    for param_name in list(pyro.get_param_store().keys()):\n        if 'ATAC_' in param_name:\n            del pyro.get_param_store()[param_name]","metadata":{"execution":{"iopub.status.busy":"2023-12-08T03:05:08.105096Z","iopub.execute_input":"2023-12-08T03:05:08.105398Z","iopub.status.idle":"2023-12-08T03:05:11.760216Z","shell.execute_reply.started":"2023-12-08T03:05:08.105372Z","shell.execute_reply":"2023-12-08T03:05:11.759236Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if torch.cuda.is_available():\n    gc.collect()\n    torch.cuda.empty_cache()\n    model = model.to('cuda')\nepoch_losses, step_losses = train(model.model,\n                                  model.guide,\n                                  lr=6e-6,\n                                  epochs=FINETUNE_EPOCHS,\n                                  dataloader=make_dataloaders(1250 // (1 if torch.cuda.is_available() else 6)),\n                                  reset=not USE_ATAC)\ntorch.save(epoch_losses, \"epoch_losses.pt\")\ntorch.save(step_losses, \"step_losses.pt\")\ndel epoch_losses, step_losses","metadata":{"execution":{"iopub.status.busy":"2023-12-08T03:05:11.761336Z","iopub.execute_input":"2023-12-08T03:05:11.761635Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.save()\n# model.load()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"results = {\n    'cell_type': [],\n    'gene': [],\n    'sm_name': [],\n    'pval': [],\n    'sign': [],\n    'score': []\n}\nwith torch.no_grad():\n    param_store = pyro.get_param_store()\n    dmso_i = [i for i, perturb_name in enumerate(embeddings.keys()) if perturb_name == \"Dimethyl Sulfoxide\"][0]\n    dmso_alpha = param_store['genewise_perturb_prior_alpha_MAP'][dmso_i].detach().cpu().numpy()\n    dmso_beta = param_store['genewise_perturb_prior_beta_MAP'][dmso_i].detach().cpu().numpy()\n    for i, (perturb_name, perturb) in tqdm(enumerate(embeddings.items()), \"Collecting...\", total=len(embeddings)):\n#         alpha = param_store['genewise_perturb_prior_alpha_MAP'][i]\n#         beta = param_store['genewise_perturb_prior_beta_MAP'][i]\n#         posterior_beta = dist.Beta(alpha, beta)\n        for celltype_i in range(model.n_celltypes):\n            pvals = scipy.stats.beta.sf(param_store['perturbation_mask_MAP'][i, celltype_i].detach().cpu().numpy(), dmso_alpha, dmso_beta)\n            scales = param_store['perturbation_scale_MAP'][i, celltype_i].detach().cpu().numpy()\n            results['sm_name'] += [perturb_name]*model.n_genes\n            results['cell_type'] += [unique_celltypes[celltype_i]]*model.n_genes\n            results['gene'] += list(adata_train.var_names)\n            results['pval'] += list(pvals)\n            results['sign'] += list(np.sign(scales))\n            results['score'] += list(-1 * np.log10(np.maximum(pvals, 1e-12)) * np.sign(scales))\n\nresults = pd.DataFrame(results)\nresults.to_csv(output_dir / 'results.csv')\n","metadata":{"scrolled":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"id_map = pd.read_csv(input_dir / 'id_map.csv')\nid_map.id = id_map.id.astype(int)\nresults = pd.read_csv(output_dir / 'results.csv')\nresults = results.merge(id_map, on=['cell_type', 'sm_name'], how=\"left\")\nresults = results.dropna()  # NA for Test set FIXME\nresults = results.drop(['cell_type', 'sm_name', 'pval', 'sign', 'Unnamed: 0'], axis=1)\nresults.id = results.id.astype(int)\ndel id_map\nsample_submission = pd.read_csv(input_dir / 'sample_submission.csv')\ngene_names = sample_submission.columns.tolist()\nresults = results.pivot(index='id', columns='gene', values='score')\nresults = results.sort_values('id')\nresults = results[gene_names[1:]]  # Re-order\ngc.collect()\nresults.to_csv(output_dir / 'submission.csv')\nresults","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}