{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# JAX Autoencoder Playground\n\nThis notebook creates an autoencoder in JAX, using best practices around random keys. We read in features created from other Vendekagon Labs notebooks that enrich our small molecule information, one-hot encode our cell type and small molecule mechanism of action, and then train the autoencoder on the full dimensionality of the resulting dataset.\n\nOnce we have the autoencoder trained, we use it to train basic linear regression models. We validate that we can use these to cross-predict cell types, when trained on other cell types that are present in our dataset.\n\nWe then train linear models on the few B cell and Myeloid lineage cells provided, one for each cell type in the `de_train` data. We then use these models to make out-of-fold predictions on the remaining B cells and Myeloid cells for gene perturbation outcomes. After that, we average these together into one DataFrame, and use that as our submission.\n\n_NOTE_: this notebook deliberately avoids doing any rigorous cross-validation, hyperparameter search, or much other neural network tuning. Instead we make a best guess about how to provide a little regularization and not train too long. It's intended to be used as a starting point, showing a very simple implementation and training loop, so as to support more elaborate follow-on work.","metadata":{}},{"cell_type":"markdown","source":"For some reason, TPU instances don't always start with the parquet reading libraries available, so we do a just-in-case pip install.","metadata":{}},{"cell_type":"code","source":"%%capture\n%pip install pyarrow fastparquet","metadata":{"execution":{"iopub.status.busy":"2023-09-30T12:40:18.418577Z","iopub.execute_input":"2023-09-30T12:40:18.419459Z","iopub.status.idle":"2023-09-30T12:40:27.102149Z","shell.execute_reply.started":"2023-09-30T12:40:18.419422Z","shell.execute_reply":"2023-09-30T12:40:27.100685Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import jax\nimport flax\nimport pandas as pd\nfrom pathlib import Path\nfrom jax import numpy as jnp\nfrom jax import random\nfrom flax import linen as nn\nimport optax\nfrom flax.training import train_state","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-09-30T12:40:27.104304Z","iopub.execute_input":"2023-09-30T12:40:27.104569Z","iopub.status.idle":"2023-09-30T12:40:30.823362Z","shell.execute_reply.started":"2023-09-30T12:40:27.104543Z","shell.execute_reply":"2023-09-30T12:40:30.822563Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Supplementary Datasets\n\nWe read in supplemental compound information from these notebooks:\n\n- [Come on Chemicals!](https://www.kaggle.com/code/vendekagonlabs/come-on-chemicals-r-version/data)\n- [Supplemental Compound Info](https://www.kaggle.com/code/vendekagonlabs/supplemental-compound-info/data)\n\n_NOTE_: this notebook makes no claim these are the best way to encode our chemical features (they almost certainly aren't), but simply uses these as a starting example for how to incorporate feature engineering before using our AE model.","metadata":{}},{"cell_type":"code","source":"data_dir = Path('/kaggle/input/') \ncomp_dir = Path(data_dir / 'open-problems-single-cell-perturbations')\nde_train_path = comp_dir / 'de_train.parquet'\nfinger_features_path = '/kaggle/input/come-on-chemicals-r-version/finger_features.csv'\nsupp_compound_path = '/kaggle/input/supplemental-compound-info/compounds.tsv'\nworking_dir = Path('/kaggle/working/')","metadata":{"execution":{"iopub.status.busy":"2023-09-30T12:40:30.824297Z","iopub.execute_input":"2023-09-30T12:40:30.824633Z","iopub.status.idle":"2023-09-30T12:40:30.829144Z","shell.execute_reply.started":"2023-09-30T12:40:30.824610Z","shell.execute_reply":"2023-09-30T12:40:30.828371Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"de_train = pd.read_parquet(de_train_path)\nsupp_compound = pd.read_csv(supp_compound_path, sep='\\t')\nfinger_features = pd.read_csv(finger_features_path)","metadata":{"execution":{"iopub.status.busy":"2023-09-30T12:40:30.830752Z","iopub.execute_input":"2023-09-30T12:40:30.831012Z","iopub.status.idle":"2023-09-30T12:40:33.113783Z","shell.execute_reply.started":"2023-09-30T12:40:30.830990Z","shell.execute_reply":"2023-09-30T12:40:33.112933Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"moa = pd.get_dummies(supp_compound['moa'])\nsm_moa = supp_compound[['sm_name']].join(moa).drop_duplicates('sm_name')\nsm_moa","metadata":{"execution":{"iopub.status.busy":"2023-09-30T12:40:33.114758Z","iopub.execute_input":"2023-09-30T12:40:33.115010Z","iopub.status.idle":"2023-09-30T12:40:33.147340Z","shell.execute_reply.started":"2023-09-30T12:40:33.114987Z","shell.execute_reply":"2023-09-30T12:40:33.146515Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Preprocessing\n\n## Adding Features\n\nI've rolled up all the feature additions into one function, so that it's easy to use this as a starting point to incorporate your own feature engineering. This makes it easy to get started with an AE model by simply just changing this function, then running the rest of the notebook.","metadata":{}},{"cell_type":"code","source":"def add_features(df):\n    one_hot = df.join(pd.get_dummies(df['cell_type']), how='left')\n    add_chem  = one_hot.merge(finger_features, on='sm_name', how='left')\n    one_hot_moa = add_chem.join(pd.get_dummies(supp_compound['moa']), how='left')\n    return one_hot_moa\n\nde_train_feats = add_features(de_train)","metadata":{"execution":{"iopub.status.busy":"2023-09-30T12:40:33.148394Z","iopub.execute_input":"2023-09-30T12:40:33.148686Z","iopub.status.idle":"2023-09-30T12:40:33.266907Z","shell.execute_reply.started":"2023-09-30T12:40:33.148662Z","shell.execute_reply":"2023-09-30T12:40:33.265968Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Extracting and Scaling Training Data\n\nIn this case, we train the AE using only the four cell types which we have (mostly) complete gene perturbation information for. Later, we'll use cross-prediction across these types as a proxy task to evaluate how to set model parameters for fitting a simple model to our AE's embeddings.","metadata":{}},{"cell_type":"code","source":"train_cells = ['NK cells', 'T cells CD4+', 'T regulatory cells', 'T cells CD8+']\n\ntrain = de_train_feats[de_train_feats['cell_type'].isin(train_cells)]\nX_pre = train[train.columns[5:]].values\n\nprint(f\"Training data shape: {X_pre.shape}\")","metadata":{"execution":{"iopub.status.busy":"2023-09-30T12:40:33.267954Z","iopub.execute_input":"2023-09-30T12:40:33.268258Z","iopub.status.idle":"2023-09-30T12:40:33.989789Z","shell.execute_reply.started":"2023-09-30T12:40:33.268232Z","shell.execute_reply":"2023-09-30T12:40:33.988994Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train.head(10)","metadata":{"execution":{"iopub.status.busy":"2023-09-30T12:40:33.990783Z","iopub.execute_input":"2023-09-30T12:40:33.991062Z","iopub.status.idle":"2023-09-30T12:40:34.013907Z","shell.execute_reply.started":"2023-09-30T12:40:33.991025Z","shell.execute_reply":"2023-09-30T12:40:34.013086Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.preprocessing import StandardScaler\n\nss = StandardScaler()\nss.fit(X_pre)\nX = ss.transform(X_pre)","metadata":{"execution":{"iopub.status.busy":"2023-09-30T12:40:34.014954Z","iopub.execute_input":"2023-09-30T12:40:34.015303Z","iopub.status.idle":"2023-09-30T12:40:34.745575Z","shell.execute_reply.started":"2023-09-30T12:40:34.015257Z","shell.execute_reply":"2023-09-30T12:40:34.744636Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Creating our Autoencoder\n\nHere we define the autoencoder using `flax.linen` layers, a functional style JAX wrapping library that provides more high-level routines (some which overlapy with `pytorch`, some more similar to wrapping libs in that ecosystem like `pytorch-lightning`).\n\nWe define the Encoder and Decoder separately, and for the quickstart example, we use a very basic Autoencoder class to wrap these, which just calls the two in sequence, returning the reconstructed output.","metadata":{}},{"cell_type":"code","source":"class Encoder(nn.Module):\n    c_hid : int\n    latent_dim : int\n    training: bool\n\n    @nn.compact\n    def __call__(self, x):\n        x = nn.Dropout(rate=0.10, deterministic=not self.training)(x)\n        x = nn.Dense(features=2*self.c_hid)(x)\n        x = nn.gelu(x)\n        x = nn.Dropout(rate=0.25, deterministic=not self.training)(x)\n        x = nn.Dense(features=self.c_hid)(x)\n        x = nn.gelu(x)\n        x = nn.Dropout(rate=0.25, deterministic=not self.training)(x)\n        x = nn.Dense(features=self.c_hid)(x)\n        x = nn.gelu(x)\n        x = nn.Dropout(rate=0.25, deterministic=not self.training)(x)\n        x = nn.Dense(features=self.latent_dim)(x)\n        return x\n    \n    \nclass Decoder(nn.Module):\n    c_out : int\n    c_hid : int\n    latent_dim : int\n    training: bool\n\n    @nn.compact\n    def __call__(self, x):\n        x = nn.Dense(features=self.c_hid)(x)\n        x = nn.gelu(x)\n        x = nn.Dropout(rate=0.25, deterministic=not self.training)(x)\n        x = nn.Dense(features=2*self.c_hid)(x)\n        x = nn.gelu(x)\n        x = nn.Dropout(rate=0.25, deterministic=not self.training)(x)\n        x = nn.Dense(features=2*self.c_hid)(x)\n        x = nn.gelu(x)\n        x = nn.Dropout(rate=0.25, deterministic=not self.training)(x)\n        x = nn.Dense(features=self.c_out)(x)\n        x = nn.tanh(x)\n        return x\n\n    \nclass AutoEncoder(nn.Module):\n    c_hid: int\n    latent_dim : int\n    input_dim: int\n    training: bool\n\n    def setup(self):\n        self.encoder = Encoder(c_hid=self.c_hid,\n                               latent_dim=self.latent_dim,\n                               training=self.training)\n        self.decoder = Decoder(c_hid=self.c_hid,\n                               latent_dim=self.latent_dim,\n                               c_out=self.input_dim, training=self.training)\n\n    def __call__(self, x):\n        z = self.encoder(x)\n        x_hat = self.decoder(z)\n        return x_hat","metadata":{"execution":{"iopub.status.busy":"2023-09-30T12:40:34.749145Z","iopub.execute_input":"2023-09-30T12:40:34.749522Z","iopub.status.idle":"2023-09-30T12:40:34.766203Z","shell.execute_reply.started":"2023-09-30T12:40:34.749497Z","shell.execute_reply":"2023-09-30T12:40:34.765424Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training Regime\n\nOne way our JAX approach differs from other NN training packages, is that it decouples our model objects from the state of training. Here we use the flax `TrainState` helper to group the state of our training run as properties on a single object.\n\nIt also makes random keys required arguments, and provides a lot of functions for creating additional random keys when need be, by using `split` on our existing keys. This way we can link up the random components of our model so that it will be fully reproducible in a deterministic way, while also specifying that some of our random state needs to be independent from the rest.\n\nApart from that, we use an all caps variable convention as is common in other notebooks to define the model & training regime's hyperparameters (which you're most likely to target with model searches and/or manual adjustments). As a note: a lot of JAX models you'll find elsewhere use a wrapping object or dict and a  `config.property` or `config[\"property\"]` convention for settings these.\n\nIn this case, we use a very simple mean squared error metric as our training loss. There are other ways to define reconstruction loss, and we'll use some of these in the _Bonus: VAE_ section of the notebook.","metadata":{}},{"cell_type":"code","source":"# ------ model & training hyperparams ------\nLATENT_DIM = 512\nHIDDEN_BASE_DIM = 1024\nINPUT_DIM = X.shape[1]  # 19319\nBATCH_SIZE = 8\nLEARNING_RATE = 1e-4\nEPOCHS=200\n\n# ------------------------------------------\nrng = random.PRNGKey(0)\nmain_key, params_key, dropout_key = jax.random.split(key=rng, num=3)\n# ------------------------------------------\n\n\n# ----------- initialize model -------------\nae = AutoEncoder(\n    input_dim=INPUT_DIM,\n    c_hid=HIDDEN_BASE_DIM,\n    latent_dim=LATENT_DIM,\n    training=False\n)\nvariables = ae.init(params_key, jnp.ones([BATCH_SIZE, INPUT_DIM]))\nstate = train_state.TrainState.create(\n        apply_fn = ae.apply,\n        tx=optax.adam(LEARNING_RATE),\n        params=variables['params']\n)\n\n\n# ------------ fns to drive training -------\n@jax.jit\ndef mse(params, x_batched, y_batched):\n    def squared_error(x, y):\n        pred = ae.apply({'params': params}, x, rngs={'dropout': dropout_key})\n        return jnp.inner(y - pred, y - pred) / 2.0\n    return jnp.mean(jax.vmap(squared_error)(x_batched, y_batched), axis=0)\n\n@jax.jit\ndef train_step(\n    state: train_state.TrainState, batch: jnp.ndarray\n):\n\n    def loss_fn(params):\n        # logits = state.apply_fn({'params': params}, batch)\n        loss = mse(params, batch, batch)\n        return loss\n\n    gradient_fn = jax.value_and_grad(loss_fn)\n    loss, grads = gradient_fn(state.params)\n    state = state.apply_gradients(grads=grads)\n    return state, loss","metadata":{"execution":{"iopub.status.busy":"2023-09-30T12:40:34.767259Z","iopub.execute_input":"2023-09-30T12:40:34.767528Z","iopub.status.idle":"2023-09-30T12:40:46.349495Z","shell.execute_reply.started":"2023-09-30T12:40:34.767504Z","shell.execute_reply":"2023-09-30T12:40:46.348371Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The actual training loop is below. We might the `EPOCHS` value here while developing to see if the model looks like it's getting anywhere.\n\nThis is a cartoonishly simple training regime. We're not tracking additional metrics, taking snapshots and using a validation reconstruction loss to trigger early stopping, etc. There are many ways to improve this, deliberately left out to keep this:\n\n- simple for beginners\n- easy to improve","metadata":{}},{"cell_type":"code","source":"from itertools import cycle\n\nEPOCHS = 200\n\n# the jitter term is a crude way to see \ntotal = X.shape[0]\niters = total // BATCH_SIZE\njitter = cycle(range(total % BATCH_SIZE))\n\nX_jax = jnp.array(X)\nloss = 1e6\n\nprint(f\"Training for {EPOCHS} epochs.\")\n\nfor e in range(EPOCHS):\n    epoch_offset = next(jitter)\n    print(f\"Epoch: {e}, most recent training loss: {loss}\")\n\n    # -- you can uncomment this to permute rows randomly across epochs --\n    # random permutation of rows seems to hurt, probably because we lose\n    # being distributed evenly across cell types, why we use jitter strategy\n    # X_jax = random.permutation(main_key, X_jax, axis=0, independent=False)\n    for i in range(iters):\n        x0 = i*4 + epoch_offset\n        x1 = (i+1)*4 + epoch_offset\n        state, loss = train_step(state, X_jax[x0:x1, :])","metadata":{"execution":{"iopub.status.busy":"2023-09-30T12:40:46.350655Z","iopub.execute_input":"2023-09-30T12:40:46.350938Z","iopub.status.idle":"2023-09-30T12:42:40.209245Z","shell.execute_reply.started":"2023-09-30T12:40:46.350912Z","shell.execute_reply":"2023-09-30T12:42:40.208091Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Using our Trained Encoder\n\nThe encoder wouldn't do us much good if we couldn't use us to encode features! We create a function that will let us call the encoder model on data we pass in that's of the same shape as training data. We then use this to encode features for our cell type + chemical + gene perturbation training space.\n\n_Note_: we have to use the `state.params` here, as again the model weights, as a stateful part of our training, are not a part of our model definition. Instead, this is more like partial function changing, i.e. `model(state.params, arr) => embeddings` as `model.bind(state.params).encoder(arr) => embeddings`.","metadata":{}},{"cell_type":"code","source":"@jax.jit\ndef encode(arr):\n    return ae.bind({'params': state.params}).encoder(arr)\n\nembedded = encode(X[:,:])\nprint(f\"Original dims: {X.shape} reduced to {embedded.shape} dimensionality.\")","metadata":{"execution":{"iopub.status.busy":"2023-09-30T12:42:40.210372Z","iopub.execute_input":"2023-09-30T12:42:40.210637Z","iopub.status.idle":"2023-09-30T12:42:47.661380Z","shell.execute_reply.started":"2023-09-30T12:42:40.210613Z","shell.execute_reply":"2023-09-30T12:42:47.660416Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Cell Type Cross-Prediction\n\nWe have a big challenge in this dataset, in that we have very limited training samples in `de_train` for our target cell types, B cells and Myeloid cells. We can also see that these cell phenotypes are fairly different from our training data, three of which are subtypes of T cells!\n\nThis notebook uses an oversimplified model selection technique. We use cross-prediction to predict cell types against the other cell types we do have, using this as a proxy task to predict the cells we don't have. We then use that to set a regularization term for basic ridge regression, then use that regularization term in the model we fit for predicting B cells and Myeloid cells. We don't use any sophisticated cross-validation (you could, for instance do k-fold where the prediction task is similarly limited, from the same number of instances of B cells we have, predict all the rest). Left as an exercise for the reader...","metadata":{}},{"cell_type":"code","source":"import numpy as np\nfrom sklearn.linear_model import Ridge\nfrom sklearn.multioutput import MultiOutputRegressor\nfrom sklearn.ensemble import RandomForestRegressor\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.decomposition import TruncatedSVD\nfrom sklearn.metrics import mean_squared_error","metadata":{"execution":{"iopub.status.busy":"2023-09-30T12:42:47.662429Z","iopub.execute_input":"2023-09-30T12:42:47.662702Z","iopub.status.idle":"2023-09-30T12:42:48.084271Z","shell.execute_reply.started":"2023-09-30T12:42:47.662678Z","shell.execute_reply":"2023-09-30T12:42:48.083278Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_cells","metadata":{"execution":{"iopub.status.busy":"2023-09-30T12:42:48.085249Z","iopub.execute_input":"2023-09-30T12:42:48.085506Z","iopub.status.idle":"2023-09-30T12:42:48.090695Z","shell.execute_reply.started":"2023-09-30T12:42:48.085483Z","shell.execute_reply":"2023-09-30T12:42:48.089907Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Cross-prediction training details\n\nWe use our autoencoder to train the source/training cell type's data. To ensure we're not overly coupled to the dimensionality reduction learned by the autoencoder, we make our training target a more straight forward dimensionality reduction of the other cell type's gene logfold change space, using `TruncatedSVD`. Note that the `n_components` param was set with naive manual inspection, and could be more rigorously tuned.\n\nWe also have to re-use the `StandardScaler` we used as an input for model training, to match the autoencoder's expected range and distributino of data, then invert our TruncatedSVD transform to get our predictions back into gene expression space, then evaluate our resulting model error there.","metadata":{}},{"cell_type":"code","source":"genes = de_train.columns[5:]\n\n\ndef naive_cross_predict(df, cell1='T cells CD8+', cell2='T cells CD4+'):\n    cell1_df = df[df['cell_type'] == cell1]\n    cell2_df = df[df['cell_type'] == cell2]\n    cell1_sms = list(cell1_df['sm_name'].unique())\n    cell2_sms = list(cell2_df['sm_name'].unique())\n    common = [sm for sm in cell1_sms if sm in cell2_sms]\n    print(f\"Found {len(common)} sm_names shared between cell types.\")\n    cell1_sub = (cell1_df[cell1_df['sm_name'].isin(common)]).sort_values(by='sm_name')\n    cell2_sub = (cell2_df[cell2_df['sm_name'].isin(common)]).sort_values(by='sm_name')\n    Xc1 = cell1_sub[cell1_sub.columns[5:]].values\n    y = cell2_sub[genes].values\n    tsvd = TruncatedSVD(n_components=40)\n    tsvd.fit(y)\n    scaled = ss.transform(Xc1)\n    embed = encode(jnp.array(scaled))\n    y_dr = tsvd.transform(y)\n    lm = Ridge(alpha=55)\n    lm.fit(embed, y_dr)\n    y_dr_hat = lm.predict(embed)\n    y_hat = tsvd.inverse_transform(y_dr_hat)\n    mse = mean_squared_error(y, y_hat)\n    return lm, tsvd, mse","metadata":{"execution":{"iopub.status.busy":"2023-09-30T12:42:48.091530Z","iopub.execute_input":"2023-09-30T12:42:48.091751Z","iopub.status.idle":"2023-09-30T12:42:48.102893Z","shell.execute_reply.started":"2023-09-30T12:42:48.091732Z","shell.execute_reply":"2023-09-30T12:42:48.102181Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Cross-prediction Loop\n\nWith our training defined, we run it across all combinations of cell types we havei n the data.","metadata":{}},{"cell_type":"code","source":"from itertools import combinations\n\nmodels = {}\n\nfor ct1, ct2 in combinations(train_cells, 2):\n    lm, tsvd, mse = naive_cross_predict(train, cell1=ct1, cell2=ct2)\n    models[f'{ct1} => {ct2}'] = {\n        'linear_regression_model': lm,\n        'tsvd_trasnform_obj': tsvd,\n        'mae': mse\n    }\n    print(f\"Naive cross-prediction from {ct1} => {ct2}, MSE: {mse}\")","metadata":{"execution":{"iopub.status.busy":"2023-09-30T12:42:48.103709Z","iopub.execute_input":"2023-09-30T12:42:48.103942Z","iopub.status.idle":"2023-09-30T12:43:07.122420Z","shell.execute_reply.started":"2023-09-30T12:42:48.103919Z","shell.execute_reply":"2023-09-30T12:43:07.121040Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Predicting B and Myeloid Cell Gene Perturbation\n\nNow that we've done some test cross-prediction, we read in the data for which we have no gene perturbation provided, to use the same strategy to make predictions. We re-use our same `naive_cross_predict` function from our prior cross-type prediction to fit a model to the data included in `de_train`.","metadata":{}},{"cell_type":"code","source":"id_map = pd.read_csv(comp_dir / 'id_map.csv')\nid_map","metadata":{"execution":{"iopub.status.busy":"2023-09-30T12:43:07.123598Z","iopub.execute_input":"2023-09-30T12:43:07.123863Z","iopub.status.idle":"2023-09-30T12:43:07.140888Z","shell.execute_reply.started":"2023-09-30T12:43:07.123839Z","shell.execute_reply":"2023-09-30T12:43:07.140019Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for ct1 in train_cells:\n    for ct2 in ['B cells', 'Myeloid cells']:\n        lm, tsvd, mse = naive_cross_predict(de_train_feats, cell1=ct1, cell2=ct2)\n        models[f'{ct1} => {ct2}'] = {\n            'linear_regression_model': lm,\n            'tsvd_trasnform_obj': tsvd,\n            'mae': mse\n        }\n        print(f\"Naive cross-prediction from {ct1} => {ct2}, MSE: {mse}\")","metadata":{"execution":{"iopub.status.busy":"2023-09-30T12:43:07.141952Z","iopub.execute_input":"2023-09-30T12:43:07.142487Z","iopub.status.idle":"2023-09-30T12:43:19.766429Z","shell.execute_reply.started":"2023-09-30T12:43:07.142460Z","shell.execute_reply":"2023-09-30T12:43:19.764333Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"This next section is just to show how we transform the ` => ` key to and back from tuples. It's a little bit of extra string munging, but was helpful for me to see the ` => ` prediction ordering while developing this, so I lefti t in.","metadata":{}},{"cell_type":"code","source":"import pprint\npp = pprint.PrettyPrinter(indent=1, depth=2)\n\ndef to_mapping(s1, s2):\n    return f\"{s1} => {s2}\"\n\nmodel_key_tuples = [k.split(' => ') for k in models.keys()]\npp.pprint(model_key_tuples)","metadata":{"execution":{"iopub.status.busy":"2023-09-30T12:43:19.768917Z","iopub.execute_input":"2023-09-30T12:43:19.770681Z","iopub.status.idle":"2023-09-30T12:43:19.783839Z","shell.execute_reply.started":"2023-09-30T12:43:19.770606Z","shell.execute_reply":"2023-09-30T12:43:19.781971Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"[to_mapping(*k) for k in model_key_tuples]","metadata":{"execution":{"iopub.status.busy":"2023-09-30T12:43:19.786815Z","iopub.execute_input":"2023-09-30T12:43:19.787943Z","iopub.status.idle":"2023-09-30T12:43:19.808115Z","shell.execute_reply.started":"2023-09-30T12:43:19.787875Z","shell.execute_reply":"2023-09-30T12:43:19.806114Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"b_cell_keys = [(src, target) for src, target in model_key_tuples if target == 'B cells']\nmyeloid_cell_keys = [(src, target) for src, target in model_key_tuples if target == 'Myeloid cells']\n[to_mapping(*k) for k in b_cell_keys] + [to_mapping(*k) for k in myeloid_cell_keys]","metadata":{"execution":{"iopub.status.busy":"2023-09-30T12:43:19.810296Z","iopub.execute_input":"2023-09-30T12:43:19.810938Z","iopub.status.idle":"2023-09-30T12:43:19.829075Z","shell.execute_reply.started":"2023-09-30T12:43:19.810875Z","shell.execute_reply":"2023-09-30T12:43:19.827018Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Using Our Simple Models to Make Submission Predictions\n\nUsing the models we fit to our data above, we make now fit our data. There is some shared logic with our naive cross prediction function here, none of it refactored out. (In this case, it's a little easier to read with each in its own place). A lot of this is just collection manipulation in base Python to ensure we're predicting B cells from the other cell types.\n\nNote that the strategy being used here is to fit a model for each of our reference cell types --(CD8+, CD4+, regulatory) T cells, NK cells -- then average the predictions. You could also make one big predictive modeling using the information you have on all cells, and thoughtfully built model of that form is likely to do better than the simpler averaging strategy.","metadata":{}},{"cell_type":"code","source":"bcell_predictions = {}\nbcells = id_map[id_map['cell_type'] == 'B cells']\n\n\nfor mod_key, bcell in b_cell_keys:\n    mk = to_mapping(mod_key, bcell)\n    lrm = models[mk]['linear_regression_model']\n    tsvd = models[mk]['tsvd_trasnform_obj']\n    mod_cells_df = train[train['cell_type'] == mod_key]\n    mod_sms = list(mod_cells_df['sm_name'].unique())\n    bcell_sms = list(bcells['sm_name'].unique())\n    common = [sm for sm in mod_sms if sm in bcell_sms]\n    print(f\"Found {len(common)} sm_names shared between cell types, for {len(bcell_sms)} needed for submission.\")\n    mod_sub = (mod_cells_df[mod_cells_df['sm_name'].isin(common)]).sort_values(by='sm_name')\n    bcell_sub = (bcells[bcells['sm_name'].isin(common)]).sort_values(by='sm_name')\n    Xmc = mod_sub[mod_sub.columns[5:]].values\n    Xmc = ss.transform(Xmc)\n    Xmc = encode(jnp.array(Xmc))\n    bcell_gx_pred_dr = lrm.predict(Xmc)\n    bcell_gx_pred = tsvd.inverse_transform(bcell_gx_pred_dr)\n    bcell_predictions[mod_key] = {\n        'pred': bcell_gx_pred,\n        'sm_names': bcell_sub['sm_name']\n    }","metadata":{"execution":{"iopub.status.busy":"2023-09-30T12:43:19.832050Z","iopub.execute_input":"2023-09-30T12:43:19.832696Z","iopub.status.idle":"2023-09-30T12:43:32.004596Z","shell.execute_reply.started":"2023-09-30T12:43:19.832633Z","shell.execute_reply":"2023-09-30T12:43:32.002835Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"bcells = id_map[id_map['cell_type'] == 'Myeloid cells']\nmyeloid_predictions = {}\n\n# just use bcell name again for myeloid cells, fix later\nfor mod_key, bcell in myeloid_cell_keys:\n    mk = to_mapping(mod_key, bcell)\n    lrm = models[mk]['linear_regression_model']\n    tsvd = models[mk]['tsvd_trasnform_obj']\n    mod_cells_df = train[train['cell_type'] == mod_key]\n    mod_sms = list(mod_cells_df['sm_name'].unique())\n    bcell_sms = list(bcells['sm_name'].unique())\n    common = [sm for sm in mod_sms if sm in bcell_sms]\n    print(f\"Found {len(common)} sm_names shared between cell types, for {len(bcell_sms)} needed for submission.\")\n    mod_sub = (mod_cells_df[mod_cells_df['sm_name'].isin(common)]).sort_values(by='sm_name')\n    bcell_sub = (bcells[bcells['sm_name'].isin(common)]).sort_values(by='sm_name')\n    Xmc = mod_sub[mod_sub.columns[5:]].values\n    Xmc = ss.transform(Xmc)\n    Xmc = encode(jnp.array(Xmc))\n    bcell_gx_pred_dr = lrm.predict(Xmc)\n    bcell_gx_pred = tsvd.inverse_transform(bcell_gx_pred_dr)\n    myeloid_predictions[mod_key] = {\n        'pred': bcell_gx_pred,\n        'sm_names': bcell_sub['sm_name']\n    }","metadata":{"execution":{"iopub.status.busy":"2023-09-30T12:43:32.005897Z","iopub.execute_input":"2023-09-30T12:43:32.006221Z","iopub.status.idle":"2023-09-30T12:43:38.794295Z","shell.execute_reply.started":"2023-09-30T12:43:32.006192Z","shell.execute_reply":"2023-09-30T12:43:38.792774Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"As a sanity check, we assert that our predictions are the right shape w/r/t the space of gene logfold change.","metadata":{}},{"cell_type":"code","source":"assert myeloid_predictions['NK cells']['pred'].shape[1] == len(genes)\nassert bcell_predictions['NK cells']['pred'].shape[1] == len(genes)","metadata":{"execution":{"iopub.status.busy":"2023-09-30T12:43:38.795612Z","iopub.execute_input":"2023-09-30T12:43:38.795902Z","iopub.status.idle":"2023-09-30T12:43:38.802454Z","shell.execute_reply.started":"2023-09-30T12:43:38.795874Z","shell.execute_reply":"2023-09-30T12:43:38.801237Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Then we average our predictions together.","metadata":{}},{"cell_type":"code","source":"from functools import reduce\n\neach_pred = [k for k in bcell_predictions.keys()]\nfor pred in each_pred:\n    bcell_predictions[pred]['df'] = pd.DataFrame(bcell_predictions[pred]['pred'], columns=genes,\n                                                 index=bcell_predictions[pred]['sm_names'].index)\n    myeloid_predictions[pred]['df'] = pd.DataFrame(myeloid_predictions[pred]['pred'], columns=genes,\n                                                   index=myeloid_predictions[pred]['sm_names'].index)\nbcell_all = reduce(lambda x, y: x.add(y, fill_value=0.0), [val['df'] for val in bcell_predictions.values()])\nbcell_all = bcell_all / 4.0\nmyeloid_all = reduce(lambda x, y: x.add(y, fill_value=0.0), [val['df'] for val in myeloid_predictions.values()])\nmyeloid_all = myeloid_all / 4.0","metadata":{"execution":{"iopub.status.busy":"2023-09-30T12:43:38.804016Z","iopub.execute_input":"2023-09-30T12:43:38.804309Z","iopub.status.idle":"2023-09-30T12:43:38.963173Z","shell.execute_reply.started":"2023-09-30T12:43:38.804284Z","shell.execute_reply":"2023-09-30T12:43:38.962027Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Creating Submissions File\n\nWe then concatenate our results into the form expected for submissions.\n\n_NOTE_: on its own, this very simple use of a very basic autoencoder model hits ~0.72 on the public leaderboard. Like with other notebooks, we ensemble to get a better score.","metadata":{}},{"cell_type":"code","source":"result = pd.concat([bcell_all, myeloid_all])\n\n# we use a different copy of the dataframe to munge into\n# submissions column form, so we can use `result` later\n# in ensembling\nresult_id = result.copy(deep=True)\nresult_id['id'] = result_id.index\nresult_id.to_csv('ae_submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2023-09-30T12:44:31.701934Z","iopub.execute_input":"2023-09-30T12:44:31.702322Z","iopub.status.idle":"2023-09-30T12:44:39.972930Z","shell.execute_reply.started":"2023-09-30T12:44:31.702292Z","shell.execute_reply":"2023-09-30T12:44:39.971798Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Ensembling\n\nBefore ensembling this with the result of simple models and average predictions (what most of the leaderboard and public notebooks do at the moment), we check how our predictions correlate. Remember that we want some evidence that it's only weak correlation! (If it's perfectly anti-correlated the predictions just cancel out. If it's too correlated, we'll increase our combined models' bias, rather than correct for it when averaging.)","metadata":{}},{"cell_type":"code","source":"should_be_06_02 = pd.read_csv('/kaggle/input/sep28-0-602-submission-reference/submission.csv')\nshould_be_06_02","metadata":{"execution":{"iopub.status.busy":"2023-09-30T12:44:45.278570Z","iopub.execute_input":"2023-09-30T12:44:45.279236Z","iopub.status.idle":"2023-09-30T12:44:49.205818Z","shell.execute_reply.started":"2023-09-30T12:44:45.279204Z","shell.execute_reply":"2023-09-30T12:44:49.204822Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"gene_corr = should_be_06_02.corrwith(result_id)\nprint(gene_corr.mean(), gene_corr.std())\ngene_corr.hist(bins=50)","metadata":{"execution":{"iopub.status.busy":"2023-09-30T12:44:52.380221Z","iopub.execute_input":"2023-09-30T12:44:52.381188Z","iopub.status.idle":"2023-09-30T12:44:54.817111Z","shell.execute_reply.started":"2023-09-30T12:44:52.381150Z","shell.execute_reply":"2023-09-30T12:44:54.815816Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"new_sub = 0.88*should_be_06_02 + 0.12*result\nnew_sub","metadata":{"execution":{"iopub.status.busy":"2023-09-30T12:44:59.661485Z","iopub.execute_input":"2023-09-30T12:44:59.662103Z","iopub.status.idle":"2023-09-30T12:44:59.838232Z","shell.execute_reply.started":"2023-09-30T12:44:59.662062Z","shell.execute_reply":"2023-09-30T12:44:59.836906Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"new_sub['id'] = new_sub.index\nnew_sub.to_csv('submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2023-09-30T12:45:30.006946Z","iopub.execute_input":"2023-09-30T12:45:30.007461Z","iopub.status.idle":"2023-09-30T12:45:38.917603Z","shell.execute_reply.started":"2023-09-30T12:45:30.007424Z","shell.execute_reply":"2023-09-30T12:45:38.916200Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Bonus: VAE\n\nA variational auto-encoder, in simple terms, expands on the basic idea of an autoencoder by learning a probability distribution in the latent space/embeddings, rather than just where to place single training points. This is done using the 'reparameterization trick' -- there's a \n[good video lecture here](https://www.youtube.com/watch?v=iL1c1KmYPM0) that goes into details on latent space models and distributions in general, and uses that as a starting poitn for explaining VAEs.\n\nI've included a VAE starting point here, with none of the follow-on steps as above, for those who want to explore that model. Note that this model has not been tuned much, though I have provided a cosine learning rate scheduler as a tool for tuning the training regime more carefully.\n\nMany of the below examples are adapted from the JAX and Flax official documentation.","metadata":{}},{"cell_type":"code","source":"@jax.jit\ndef reparameterize(rng, mean, logvar):\n    std = jnp.exp(0.5 * logvar)\n    eps = random.normal(rng, logvar.shape)\n    return mean + eps * std\n\n\nclass VAEncoder(nn.Module):\n    c_hid : int\n    latent_dim : int\n    training: bool\n\n    @nn.compact\n    def __call__(self, x):\n        x = nn.Dropout(rate=0.1, deterministic=not self.training)(x)\n        x = nn.Dense(features=2*self.c_hid)(x)\n        x = nn.gelu(x)\n        x = nn.Dropout(rate=0.25, deterministic=not self.training)(x)\n        x = nn.Dense(features=self.c_hid)(x)\n        x = nn.gelu(x)\n        x = nn.Dropout(rate=0.25, deterministic=not self.training)(x)\n        mean_x = nn.Dense(self.latent_dim, name='fc2_mean')(x)\n        logvar_x = nn.Dense(self.latent_dim, name='fc2_logvar',\n                            kernel_init=jax.nn.initializers.variance_scaling(\n                                scale=0.25, distribution='truncated_normal', mode='fan_in'\n                            ))(x)\n        return mean_x, logvar_x\n    \n    \nclass VADecoder(nn.Module):\n    c_out : int\n    c_hid : int\n    latent_dim : int\n    training: bool\n\n    @nn.compact\n    def __call__(self, x):\n        x = nn.Dense(features=self.c_hid)(x)\n        x = nn.gelu(x)\n        x = nn.Dropout(rate=0.25, deterministic=not self.training)(x)\n        x = nn.Dense(features=2*self.c_hid)(x)\n        x = nn.gelu(x)\n        x = nn.Dropout(rate=0.25, deterministic=not self.training)(x)\n        x = nn.Dense(features=self.c_out)(x)\n        x = nn.tanh(x)\n        return x\n\n    \nclass VAE(nn.Module):\n    latent_dim: int\n    c_hid: int\n    input_dim: int\n    training: bool\n\n    def setup(self):\n        self.encoder = VAEncoder(c_hid=self.c_hid,\n                                 latent_dim=self.latent_dim,\n                                 training=self.training)\n        self.decoder = VADecoder(c_hid=self.c_hid,\n                                 latent_dim=self.latent_dim,\n                                 c_out=self.input_dim, training=self.training)\n\n    def __call__(self, x, z_rng):\n        mean, logvar = self.encoder(x)\n        z = reparameterize(z_rng, mean, logvar)\n        recon_x = self.decoder(z)\n        return recon_x, mean, logvar","metadata":{"execution":{"iopub.status.busy":"2023-09-30T12:45:47.701298Z","iopub.execute_input":"2023-09-30T12:45:47.702209Z","iopub.status.idle":"2023-09-30T12:45:47.718932Z","shell.execute_reply.started":"2023-09-30T12:45:47.702169Z","shell.execute_reply":"2023-09-30T12:45:47.717882Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def create_learning_rate_fn(n_epochs, base_learning_rate, steps_per_epoch, warmup_epochs=5):\n    warmup_fn = optax.linear_schedule(\n          init_value=0., end_value=base_learning_rate,\n          transition_steps=1 * steps_per_epoch)\n    cosine_epochs = max(n_epochs - warmup_epochs, 1)\n    cosine_fn = optax.cosine_decay_schedule(\n          init_value=base_learning_rate,\n          decay_steps=cosine_epochs * steps_per_epoch)\n    schedule_fn = optax.join_schedules(\n          schedules=[warmup_fn, cosine_fn],\n          boundaries=[warmup_epochs * steps_per_epoch])\n    return schedule_fn","metadata":{"execution":{"iopub.status.busy":"2023-09-30T12:45:48.488413Z","iopub.execute_input":"2023-09-30T12:45:48.488707Z","iopub.status.idle":"2023-09-30T12:45:48.494998Z","shell.execute_reply.started":"2023-09-30T12:45:48.488681Z","shell.execute_reply":"2023-09-30T12:45:48.493989Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import functools\n\n# ------ model & training hyperparams ------\nLATENT_DIM = 1024\nHIDDEN_BASE_DIM = 1024\nINPUT_DIM = X.shape[1]  # 19319\nBATCH_SIZE = 16\nINIT_LEARNING_RATE = 5e-5\n# ------------------------------------------\nrng = random.PRNGKey(0)\nmain_key, params_key, dropout_key, z_key = jax.random.split(key=rng, num=4)\n# ------------------------------------------\n\n\n# ----------- initialize model -------------\nvae = VAE(\n    input_dim=INPUT_DIM,\n    c_hid=HIDDEN_BASE_DIM,\n    latent_dim=LATENT_DIM,\n    training=False\n)\n\n# ------------ fns to drive training -------\n# note that we need a more complex loss function, including\n# both KL divergence and binary cross-entropy.\n# ------------------------------------------\n@jax.vmap\ndef kl_divergence(mean, logvar):\n    return -0.5 * jnp.sum(1 + logvar - jnp.square(mean) - jnp.exp(logvar))\n\n@jax.vmap\ndef binary_cross_entropy_with_logits(logits, labels):\n    logits = nn.log_sigmoid(logits)\n    return -jnp.sum(labels * logits + (1. - labels) * jnp.log(-jnp.expm1(logits)))\n\n@functools.partial(jax.jit, static_argnums=3)\ndef train_step(state, batch, z_rng, learning_rain_fn):\n    def loss_fn(params):\n        recon_x, mean, logvar = vae.apply({'params': params}, batch, z_rng)\n        bce_loss = binary_cross_entropy_with_logits(recon_x, batch).mean()\n        kld_loss = kl_divergence(mean, logvar).mean()\n        loss = bce_loss + kld_loss\n        return loss, recon_x\n\n    grad_fn = jax.value_and_grad(loss_fn, has_aux=True)\n    (loss, _), grads = grad_fn(state.params)\n    state = state.apply_gradients(grads=grads)\n    lr = learning_rate_fn(state.step)\n    metrics = {'learning_rate': lr,\n               'loss': loss}\n    return state, metrics","metadata":{"execution":{"iopub.status.busy":"2023-09-30T12:46:14.216915Z","iopub.execute_input":"2023-09-30T12:46:14.217357Z","iopub.status.idle":"2023-09-30T12:46:14.252195Z","shell.execute_reply.started":"2023-09-30T12:46:14.217324Z","shell.execute_reply":"2023-09-30T12:46:14.251032Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from itertools import cycle\n\nEPOCHS = 20\n\ntotal = X.shape[0]\niters = total // BATCH_SIZE\njitter = cycle(range(total % BATCH_SIZE))\n\nlearning_rate_fn = create_learning_rate_fn(EPOCHS, INIT_LEARNING_RATE, iters)\n\ninit_data = jnp.ones([BATCH_SIZE, INPUT_DIM], jnp.float32)\nvariables = vae.init(params_key, init_data, z_key)\nstate = train_state.TrainState.create(\n        apply_fn = vae.apply,\n        tx=optax.adam(learning_rate_fn),\n        params=variables['params']\n)\n\nX_jax = jnp.array(X)\n\nprint(f\"Training for {EPOCHS} epochs.\")\n\nfor e in range(EPOCHS + 1):\n    epoch_offset = next(jitter)\n    if e != 0:\n        print(f\"Epoch: {e}, most recent training loss (bce): {metrics['loss']}\")\n    # random permutation of rows seems to hurt, probably because we lose\n    # being distributed evenly across cell types, why we use jitter strategy\n    # instead\n    # X_jax = random.permutation(main_key, X_jax, axis=0, independent=False)\n    for i in range(iters):\n        z_rng, _ = random.split(z_key)\n        x0 = i*4 + epoch_offset\n        x1 = (i+1)*4 + epoch_offset\n        state, metrics = train_step(state, X_jax[x0:x1, :], z_rng, learning_rate_fn)","metadata":{"execution":{"iopub.status.busy":"2023-09-30T12:46:15.613905Z","iopub.execute_input":"2023-09-30T12:46:15.614327Z","iopub.status.idle":"2023-09-30T12:46:24.742329Z","shell.execute_reply.started":"2023-09-30T12:46:15.614294Z","shell.execute_reply":"2023-09-30T12:46:24.741332Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Where To Go From Here\n\nNow that we've got some interesting models, maybe we can make use of the single cell data? If that's an angle you want to pursue, you might check out the [Vendekagon Labs Single Cell EDA notebook](https://www.kaggle.com/code/vendekagonlabs/op2-single-cell-eda-10x-multiome/data) as a starting point. How could you use some of those features to represent cells differently? Or to account for the expression pertubation space as a distribution of values, rather than the aggregated values we have in `de_train`?\n\nPerhaps there's a better way to explore the information we have on our small molecules? Maybe we can get more interesting data from the supplemental compound information provided with this contest? We explore that [here](https://www.kaggle.com/code/vendekagonlabs/supplemental-compound-info/data). Maybe we can incorporate prior expectations on gene perturbation with the compound lincs-ids using [other data sources](https://lincs.hms.harvard.edu/db/sm/10018-101/).\n\nMaybe we could also come up with other ways to encode these features? The features uses here came from [Come on Chemicals!](https://www.kaggle.com/code/vendekagonlabs/come-on-chemicals-r-version/data) which was influenced by other public work in this competition. But you'll see there it only looks weakly associated with our gene perturbation results. Perhaps some better deep model embeddings could be useful?\n\nMaybe we can enrich our gene information with techniques from [Pathway/Gene Set Enrichment](https://www.kaggle.com/code/vendekagonlabs/op2-pathway-enrichment)? Will our knowledge of how the proteins encoded by genes interact help or restrict our models?\n\nWe've all got a lot more to explore. Good luck!","metadata":{}}]}