{"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":"# Multiome Quickstart\n\nThis notebook shows how to cross-validate a baseline model and create a submission for the Multiome part of the *Multimodal Single-Cell Integration* competition without running out of memory.\n\nIt does not show the EDA - see the separate notebook [MSCI EDA which makes sense ⭐️⭐️⭐️⭐️⭐️](https://www.kaggle.com/ambrosm/msci-eda-which-makes-sense).\n\nThe baseline model for the other part of the competition (CITEseq) is [here](https://www.kaggle.com/ambrosm/msci-citeseq-quickstart).","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19"}},{"cell_type":"code","source":"import os, gc, pickle\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport numpy as np\nfrom colorama import Fore, Back, Style\nfrom matplotlib.ticker import MaxNLocator\n\nfrom sklearn.base import BaseEstimator, TransformerMixin\nfrom sklearn.model_selection import KFold\nfrom sklearn.preprocessing import StandardScaler, scale\nfrom sklearn.decomposition import PCA\nfrom sklearn.dummy import DummyRegressor\nfrom sklearn.pipeline import make_pipeline, Pipeline\nfrom sklearn.linear_model import Ridge, LinearRegression\nfrom sklearn.metrics import mean_squared_error\n\nDATA_DIR = \"/kaggle/input/open-problems-multimodal/\"\nFP_CELL_METADATA = os.path.join(DATA_DIR,\"metadata.csv\")\n\nFP_CITE_TRAIN_INPUTS = os.path.join(DATA_DIR,\"train_cite_inputs.h5\")\nFP_CITE_TRAIN_TARGETS = os.path.join(DATA_DIR,\"train_cite_targets.h5\")\nFP_CITE_TEST_INPUTS = os.path.join(DATA_DIR,\"test_cite_inputs.h5\")\n\nFP_MULTIOME_TRAIN_INPUTS = os.path.join(DATA_DIR,\"train_multi_inputs.h5\")\nFP_MULTIOME_TRAIN_TARGETS = os.path.join(DATA_DIR,\"train_multi_targets.h5\")\nFP_MULTIOME_TEST_INPUTS = os.path.join(DATA_DIR,\"test_multi_inputs.h5\")\n\nFP_SUBMISSION = os.path.join(DATA_DIR,\"sample_submission.csv\")\nFP_EVALUATION_IDS = os.path.join(DATA_DIR,\"evaluation_ids.csv\")","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-11-09T10:06:10.218358Z","iopub.execute_input":"2022-11-09T10:06:10.218791Z","iopub.status.idle":"2022-11-09T10:06:10.231361Z","shell.execute_reply.started":"2022-11-09T10:06:10.218757Z","shell.execute_reply":"2022-11-09T10:06:10.229978Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# If you see a warning \"Failed to establish a new connection\" running this cell,\n# go to \"Settings\" on the right hand side, \n# and turn on internet. Note, you need to be phone verified.\n# We need this library to read HDF files.\n!pip install --quiet tables\n","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-11-09T08:10:25.070863Z","iopub.execute_input":"2022-11-09T08:10:25.071270Z","iopub.status.idle":"2022-11-09T08:10:36.477038Z","shell.execute_reply.started":"2022-11-09T08:10:25.071239Z","shell.execute_reply":"2022-11-09T08:10:36.475580Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Loading the common metadata table\n\nThe current version of the model is so primitive that it doesn't use the metadata, but we load it anyway.","metadata":{}},{"cell_type":"code","source":"df_cell = pd.read_csv(FP_CELL_METADATA)\ndf_cell_cite = df_cell[df_cell.technology==\"citeseq\"]\ndf_cell_multi = df_cell[df_cell.technology==\"multiome\"]\ndf_cell_cite.shape, df_cell_multi.shape","metadata":{"execution":{"iopub.status.busy":"2022-11-09T08:10:36.480480Z","iopub.execute_input":"2022-11-09T08:10:36.481061Z","iopub.status.idle":"2022-11-09T08:10:36.843340Z","shell.execute_reply.started":"2022-11-09T08:10:36.481012Z","shell.execute_reply":"2022-11-09T08:10:36.841798Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# The scoring function\n\nThis competition has a special metric: For every row, it computes the Pearson correlation between y_true and y_pred, and then all these correlation coefficients are averaged.","metadata":{}},{"cell_type":"code","source":"def correlation_score(y_true, y_pred):\n    \"\"\"Scores the predictions according to the competition rules. \n    \n    It is assumed that the predictions are not constant.\n    \n    Returns the average of each sample's Pearson correlation coefficient\"\"\"\n    if type(y_true) == pd.DataFrame: y_true = y_true.values\n    if type(y_pred) == pd.DataFrame: y_pred = y_pred.values\n    if y_true.shape != y_pred.shape: raise ValueError(\"Shapes are different.\")\n    corrsum = 0\n    for i in range(len(y_true)):\n        corrsum += np.corrcoef(y_true[i], y_pred[i])[1, 0]\n    return corrsum / len(y_true)\n","metadata":{"execution":{"iopub.status.busy":"2022-11-09T08:10:36.845205Z","iopub.execute_input":"2022-11-09T08:10:36.845717Z","iopub.status.idle":"2022-11-09T08:10:36.854879Z","shell.execute_reply.started":"2022-11-09T08:10:36.845666Z","shell.execute_reply":"2022-11-09T08:10:36.853253Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Preprocessing and cross-validation\n\nThe Multiome dataset is way too large to fit into 16 GByte RAM:\n- train inputs:  105942 * 228942 float32 values (97 GByte)\n- train targets: 105942 *  23418 float32 values (10 GByte)\n- test inputs:    55935 * 228942 float32 values (13 GByte)\n\nTo get a result with only 16 GByte RAM, we simplify the problem as follows:\n- We ignore the complete metadata (donors, days, cell types).\n- We read only 6000 rows of the training data.\n- We drop all feature columns which are constant.\n- Of the remaining columns, we keep only 4000.\n- We do a PCA and keep only the 4 most important components.\n- We fit a ridge regression model with 6000\\*4 inputs and 6000\\*23418 targets.","metadata":{}},{"cell_type":"code","source":"#%%time\n# Preprocessing\n\nclass PreprocessMultiome(BaseEstimator, TransformerMixin):\n    columns_to_use = slice(10000, 14000)\n    \n    @staticmethod\n    def take_column_subset(X):\n        return X[:,PreprocessMultiome.columns_to_use]\n    \n    def transform(self, X):\n        print(X.shape)\n        X = X[:,~self.all_zero_columns]\n        print(X.shape)\n        X = PreprocessMultiome.take_column_subset(X) # use only a part of the columns\n        print(X.shape)\n        gc.collect()\n\n        X = self.pca.transform(X)\n        print(X.shape)\n        return X\n\n    def fit_transform(self, X):\n        print(X.shape)\n        self.all_zero_columns = (X == 0).all(axis=0)\n        X = X[:,~self.all_zero_columns]\n        print(X.shape)\n        X = PreprocessMultiome.take_column_subset(X) # use only a part of the columns\n        print(X.shape)\n        gc.collect()\n\n        self.pca = PCA(n_components=100, copy=False, random_state=1)\n        X = self.pca.fit_transform(X)\n        plt.plot(self.pca.explained_variance_ratio_.cumsum())\n        plt.title(\"Cumulative explained variance ratio\")\n        plt.gca().xaxis.set_major_locator(MaxNLocator(integer=True))\n        plt.xlabel('PCA component')\n        plt.ylabel('Cumulative explained variance ratio')\n        plt.show()\n        print(X.shape)\n        return X\n\npreprocessor = PreprocessMultiome()\n\nmulti_train_x = None\nstart, stop = 0, 6000\nmulti_train_x = preprocessor.fit_transform(pd.read_hdf(FP_MULTIOME_TRAIN_INPUTS, start=start, stop=stop).values)\n\nmulti_train_y = pd.read_hdf(FP_MULTIOME_TRAIN_TARGETS, start=start, stop=stop)\ny_columns = multi_train_y.columns\nmulti_train_y = multi_train_y.values\nmulti_train_y_pca = preprocessor.fit_transform(pd.read_hdf(FP_MULTIOME_TRAIN_TARGETS, start=start, stop=stop).values)\nprint(multi_train_y.shape)","metadata":{"execution":{"iopub.status.busy":"2022-11-09T10:04:10.482054Z","iopub.execute_input":"2022-11-09T10:04:10.482578Z","iopub.status.idle":"2022-11-09T10:05:03.585207Z","shell.execute_reply.started":"2022-11-09T10:04:10.482536Z","shell.execute_reply":"2022-11-09T10:05:03.583792Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"multi_train_x_norm = pd.read_hdf(FP_MULTIOME_TRAIN_INPUTS, start=start, stop=stop)\nprint(\"MULTI TRAIN X NON-PCA\")\nprint(multi_train_x_norm.shape)","metadata":{"execution":{"iopub.status.busy":"2022-11-09T08:11:22.081531Z","iopub.execute_input":"2022-11-09T08:11:22.081988Z","iopub.status.idle":"2022-11-09T08:11:48.031051Z","shell.execute_reply.started":"2022-11-09T08:11:22.081947Z","shell.execute_reply":"2022-11-09T08:11:48.029805Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"multi_train_y_norm = pd.read_hdf(FP_MULTIOME_TRAIN_TARGETS, start=start, stop=stop)\n","metadata":{"execution":{"iopub.status.busy":"2022-11-09T10:06:51.817560Z","iopub.execute_input":"2022-11-09T10:06:51.818035Z","iopub.status.idle":"2022-11-09T10:06:57.062789Z","shell.execute_reply.started":"2022-11-09T10:06:51.817996Z","shell.execute_reply":"2022-11-09T10:06:57.061643Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\n# Cross-validation\n\nkf = KFold(n_splits=5, shuffle=True, random_state=1)\nscore_list = []\nfor fold, (idx_tr, idx_va) in enumerate(kf.split(multi_train_x)):\n    model = None\n    gc.collect()\n    X_tr = multi_train_x[idx_tr] # creates a copy, https://numpy.org/doc/stable/user/basics.copies.html\n    y_tr = multi_train_y[idx_tr]\n    del idx_tr\n\n    model = Ridge(copy_X=False)\n    model.fit(X_tr, y_tr)\n    del X_tr, y_tr\n    gc.collect()\n\n    # We validate the model\n    X_va = multi_train_x[idx_va]\n    y_va = multi_train_y[idx_va]\n    del idx_va\n    y_va_pred = model.predict(X_va)\n    mse = mean_squared_error(y_va, y_va_pred)\n    corrscore = correlation_score(y_va, y_va_pred)\n    del X_va, y_va\n\n    print(f\"Fold {fold}: mse = {mse:.5f}, corr =  {corrscore:.3f}\")\n    score_list.append((mse, corrscore))\n\n# Show overall score\nresult_df = pd.DataFrame(score_list, columns=['mse', 'corrscore'])\nprint(f\"{Fore.GREEN}{Style.BRIGHT}{multi_train_x.shape} Average  mse = {result_df.mse.mean():.5f}; corr = {result_df.corrscore.mean():.3f}{Style.RESET_ALL}\")\n","metadata":{"execution":{"iopub.status.busy":"2022-11-09T08:11:48.032660Z","iopub.execute_input":"2022-11-09T08:11:48.033031Z","iopub.status.idle":"2022-11-09T08:12:02.539137Z","shell.execute_reply.started":"2022-11-09T08:11:48.032997Z","shell.execute_reply":"2022-11-09T08:12:02.537702Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"By the way, this ridge regression is not much better than DummyRegressor, which scores `mse = 2.01718; corr = 0.679`.","metadata":{}},{"cell_type":"markdown","source":"# Retraining\n","metadata":{}},{"cell_type":"code","source":"# We retrain the model and then delete the training data, which is no longer needed\nmodel, score_list, result_df = None, None, None # free the RAM occupied by the old model\ngc.collect()\nmodel = Ridge(copy_X=False) # we overwrite the training data\nmodel.fit(multi_train_x, multi_train_y)\n# del multi_train_x, multi_train_y # free the RAM\n# _ = gc.collect()","metadata":{"execution":{"iopub.status.busy":"2022-11-09T08:12:02.540936Z","iopub.execute_input":"2022-11-09T08:12:02.541599Z","iopub.status.idle":"2022-11-09T08:12:03.703946Z","shell.execute_reply.started":"2022-11-09T08:12:02.541551Z","shell.execute_reply":"2022-11-09T08:12:03.702245Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install activations\nmulti_train_x","metadata":{"execution":{"iopub.status.busy":"2022-11-09T08:12:03.712601Z","iopub.execute_input":"2022-11-09T08:12:03.713588Z","iopub.status.idle":"2022-11-09T08:12:16.128324Z","shell.execute_reply.started":"2022-11-09T08:12:03.713520Z","shell.execute_reply":"2022-11-09T08:12:16.126355Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nfrom torch import nn","metadata":{"execution":{"iopub.status.busy":"2022-11-09T08:12:16.131145Z","iopub.execute_input":"2022-11-09T08:12:16.131563Z","iopub.status.idle":"2022-11-09T08:12:16.137444Z","shell.execute_reply.started":"2022-11-09T08:12:16.131523Z","shell.execute_reply":"2022-11-09T08:12:16.136383Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Encoder(nn.Module):\n    def __init__(self, num_inputs: int, num_units=32, activation=nn.PReLU):\n        super().__init__()\n        self.num_inputs = num_inputs\n        self.num_units = num_units\n        self.encode1 = nn.Linear(self.num_inputs, 16)\n        nn.init.xavier_uniform_(self.encode1.weight)\n        self.bn1 = nn.BatchNorm1d(num_units)\n        self.act1 = activation()\n\n        self.encode2 = nn.Linear(16, self.num_units)\n        nn.init.xavier_uniform_(self.encode2.weight)\n        self.bn2 = nn.BatchNorm1d(num_units)\n        self.act2 = activation()\n\n    def forward(self, x):       \n#  mat1 and mat2 shapes cannot be multiplied (228942x1 and 4x64)\n        print(x.shape)\n        n = self.encode1(x)\n        print(n.shape) \n        b = self.bn1(n)\n        print(b.shape)\n        x = self.act1(b)\n        \n        x2 = self.encode2(x)\n        b2 = self.bn2(x2)\n        x = self.act2(b2)\n        return x\n\nclass ChromEncoder(nn.Module):\n    \"\"\"\n    Consumes multiple inputs (i.e. one feature vector for each chromosome)\n    After processing everything to be the same dimensionality, concatenate\n    to form a single latent dimension\n    \"\"\"\n\n    def __init__(\n        self, num_inputs, latent_dim: int = 32, activation=nn.PReLU\n    ):\n        super(ChromEncoder, self).__init__()\n        self.num_inputs = num_inputs\n        self.act = activation\n\n        self.initial_modules = nn.ModuleList()\n        for n in self.num_inputs:\n            assert isinstance(n, int)\n            layer1 = nn.Linear(n, 32)\n            nn.init.xavier_uniform_(layer1.weight)\n            bn1 = nn.BatchNorm1d(32)\n            act1 = self.act()\n            layer2 = nn.Linear(32, 16)\n            nn.init.xavier_uniform_(layer2.weight)\n            bn2 = nn.BatchNorm1d(16)\n            act2 = self.act()\n            self.initial_modules.append(\n                nn.ModuleList([layer1, bn1, act1, layer2, bn2, act2])\n            )\n\n        self.encode2 = nn.Linear(16 * len(self.num_inputs), latent_dim)\n        nn.init.xavier_uniform_(self.encode2.weight)\n        self.bn2 = nn.BatchNorm1d(latent_dim)\n        self.act2 = self.act()\n\n    def forward(self, x):\n        assert len(x) == len(\n            self.num_inputs\n        ), f\"Expected {len(self.num_inputs)} inputs but got {len(x)}\"\n        enc_chroms = []\n        for init_mod, chrom_input in zip(self.initial_modules, x):\n            for f in init_mod:\n                chrom_input = f(chrom_input)\n            enc_chroms.append(chrom_input)\n        enc1 = torch.cat(\n            enc_chroms, dim=1\n        )  # Concatenate along the feature dim not batch dim\n        enc2 = self.act2(self.bn2(self.encode2(enc1)))\n        return enc2\n\nclass Decoder(nn.Module):\n    def __init__(\n        self,\n        num_outputs: int,\n        num_units: int = 32,\n        intermediate_dim: int = 64,\n        activation=nn.PReLU,\n        final_activation=None,\n    ):\n        super().__init__()\n        self.num_outputs = num_outputs\n        self.num_units = num_units\n\n        self.decode1 = nn.Linear(self.num_units, intermediate_dim)\n        nn.init.xavier_uniform_(self.decode1.weight)\n        self.bn1 = nn.BatchNorm1d(intermediate_dim)\n        self.act1 = activation()\n\n        self.decode21 = nn.Linear(intermediate_dim, self.num_outputs)\n        nn.init.xavier_uniform_(self.decode21.weight)\n        self.decode22 = nn.Linear(intermediate_dim, self.num_outputs)\n        nn.init.xavier_uniform_(self.decode22.weight)\n        self.decode23 = nn.Linear(intermediate_dim, self.num_outputs)\n        nn.init.xavier_uniform_(self.decode23.weight)\n\n        self.final_activations = nn.ModuleDict()\n        if final_activation is not None:\n            if isinstance(final_activation, list) or isinstance(\n                final_activation, tuple\n            ):\n                assert len(final_activation) <= 3\n                for i, act in enumerate(final_activation):\n                    if act is None:\n                        continue\n                    self.final_activations[f\"act{i+1}\"] = act\n            elif isinstance(final_activation, nn.Module):\n                self.final_activations[\"act1\"] = final_activation\n            else:\n                raise ValueError(\n                    f\"Unrecognized type for final_activation: {type(final_activation)}\"\n                )\n\n    def forward(self, x, size_factors=None):\n        \"\"\"include size factor here because we may want to scale the output by that\"\"\"\n        x = self.act1(self.bn1(self.decode1(x)))\n\n        retval1 = self.decode21(x)  # This is invariably the counts\n        if \"act1\" in self.final_activations.keys():\n            retval1 = self.final_activations[\"act1\"](retval1)\n        if size_factors is not None:\n            sf_scaled = size_factors.view(-1, 1).repeat(1, retval1.shape[1])\n            retval1 = retval1 * sf_scaled  # Elementwise multiplication\n\n        retval2 = self.decode22(x)\n        if \"act2\" in self.final_activations.keys():\n            retval2 = self.final_activations[\"act2\"](retval2)\n\n        retval3 = self.decode23(x)\n        if \"act3\" in self.final_activations.keys():\n            retval3 = self.final_activations[\"act3\"](retval3)\n\n        return retval1, retval2, retval3\n\nclass ChromDecoder(nn.Module):\n    \"\"\"\n    Network that is per-chromosome aware, but does not does not output\n    per-chromsome values, instead concatenating them into a single vector\n    \"\"\"\n\n    def __init__(\n        self,\n        num_outputs,  # Per-chromosome list of output sizes\n        latent_dim: int = 32,\n        activation=nn.PReLU,\n        final_activations=[nn.ELU(), nn.Softplus()],\n    ):\n        super().__init__()\n        self.num_outputs = num_outputs\n        self.latent_dim = latent_dim\n\n        self.decode1 = nn.Linear(self.latent_dim, len(self.num_outputs) * 16)\n        nn.init.xavier_uniform_(self.decode1.weight)\n        self.bn1 = nn.BatchNorm1d(len(self.num_outputs) * 16)\n        self.act1 = activation()\n\n        self.final_activations = nn.ModuleDict()\n        if final_activations is not None:\n            if isinstance(final_activations, list) or isinstance(\n                final_activations, tuple\n            ):\n                assert len(final_activations) <= 3\n                for i, act in enumerate(final_activations):\n                    if act is None:\n                        continue\n                    self.final_activations[f\"act{i+1}\"] = act\n            elif isinstance(final_activations, nn.Module):\n                self.final_activations[\"act1\"] = final_activations\n            else:\n                raise ValueError(\n                    f\"Unrecognized type for final_activation: {type(final_activation)}\"\n                )\n        logging.info(\n            f\"ChromDecoder with {len(self.final_activations)} output activations\"\n        )\n\n        self.final_decoders = nn.ModuleList()  # List[List[Module]]\n        for n in self.num_outputs:\n            layer0 = nn.Linear(16, 32)\n            nn.init.xavier_uniform_(layer0.weight)\n            bn0 = nn.BatchNorm1d(32)\n            act0 = activation()\n            # l = [layer0, bn0, act0]\n            # for _i in range(len(self.final_activations)):\n            #     fc_layer = nn.Linear(32, n)\n            #     nn.init.xavier_uniform_(fc_layer.weight)\n            #     l.append(fc_layer)\n            # self.final_decoders.append(nn.ModuleList(l))\n            layer1 = nn.Linear(32, n)\n            nn.init.xavier_uniform_(layer1.weight)\n            layer2 = nn.Linear(32, n)\n            nn.init.xavier_uniform_(layer2.weight)\n            layer3 = nn.Linear(32, n)\n            nn.init.xavier_uniform_(layer3.weight)\n            self.final_decoders.append(\n                nn.ModuleList([layer0, bn0, act0, layer1, layer2, layer3])\n            )\n\n    def forward(self, x):\n        x = self.act1(self.bn1(self.decode1(x)))\n        # This is the reverse operation of cat\n        x_chunked = torch.chunk(x, chunks=len(self.num_outputs), dim=1)\n\n        retval1, retval2, retval3 = [], [], []\n        for chunk, processors in zip(x_chunked, self.final_decoders):\n            # Each processor is a list of 3 different decoders\n            # decode1, bn1, act1, *output_decoders = processors\n            decode1, bn1, act1, decode21, decode22, decode23 = processors\n            chunk = act1(bn1(decode1(chunk)))\n            temp1 = decode21(chunk)\n            temp2 = decode22(chunk)\n            temp3 = decode23(chunk)\n\n            if \"act1\" in self.final_activations.keys():\n                # temp1 = output_decoders[0](chunk)\n                temp1 = self.final_activations[\"act1\"](temp1)\n                # retval1.append(temp1)\n            if \"act2\" in self.final_activations.keys():\n                # temp2 = output_decoders[1](chunk)\n                temp2 = self.final_activations[\"act2\"](temp2)\n                # retval2.append(temp2)\n            if \"act3\" in self.final_activations.keys():\n                # temp3 = output_decoders[2](chunk)\n                temp3 = self.final_activations[\"act3\"](temp3)\n                # retval3.append(temp3)\n            retval1.append(temp1)\n            retval2.append(temp2)\n            retval3.append(temp3)\n        retval1 = torch.cat(retval1, dim=1)\n        retval2 = torch.cat(retval2, dim=1)\n        retval3 = torch.cat(retval3, dim=1)\n        return retval1, retval2, retval3\n    \n","metadata":{"execution":{"iopub.status.busy":"2022-11-09T08:12:16.139728Z","iopub.execute_input":"2022-11-09T08:12:16.140141Z","iopub.status.idle":"2022-11-09T08:12:16.188535Z","shell.execute_reply.started":"2022-11-09T08:12:16.140104Z","shell.execute_reply":"2022-11-09T08:12:16.187315Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The final submission will contain 65744180 predictions, of which the first 6812820 are CITEseq predictions and the remaining 58931360 are Multiome. \n\nThe Multiome test predictions have 55935 rows and 23418 columns. 55935 \\* 23418 = 1’309’885’830 predictions. We'll only submit 4.5 % of these predictions. According to the data description, this subset was created by sampling 30 % of the Multiome rows, and for each row, 15 % of the columns (i.e., 16780 rows and 3512 columns per row). Consequently, when reading the test data, we can immediately drop 70 % of the rows and keep only the remaining 16780.\n\nThe eval_ids table specifies which predictions are required for the submission file.","metadata":{}},{"cell_type":"code","source":"print(multi_train_x.shape)\nmulti_train_y.shape","metadata":{"execution":{"iopub.status.busy":"2022-11-09T08:12:16.192851Z","iopub.execute_input":"2022-11-09T08:12:16.193443Z","iopub.status.idle":"2022-11-09T08:12:16.214890Z","shell.execute_reply.started":"2022-11-09T08:12:16.193362Z","shell.execute_reply":"2022-11-09T08:12:16.213785Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nclass Autoencoder(nn.Module):\n    def __init__(\n        self,\n        input_dim1: int,\n        input_dim2,\n        hidden_dim: int = 16,\n        final_activations1: list = [nn.ELU(), nn.Softplus()],\n        final_activations2=nn.Sigmoid(),\n        flat_mode: bool = True,  # Controls if we have to re-split inputs\n        seed: int = 182822,\n    ):\n        # https://stackoverflow.com/questions/9575409/calling-parent-class-init-with-multiple-inheritance-whats-the-right-way\n        nn.Module.__init__(self)\n        torch.manual_seed(seed)\n\n        self.flat_mode = flat_mode\n        self.input_dim1 = input_dim1\n        self.input_dim2 = input_dim2\n        self.num_outputs1 = (\n            len(final_activations1)\n            if isinstance(final_activations1, (list, set, tuple))\n            else 1\n        )\n        self.num_outputs2 = (\n            len(final_activations2)\n            if isinstance(final_activations2, (list, set, tuple))\n            else 1\n        )\n\n        self.encoder1 = Encoder(num_inputs=input_dim1, num_units=hidden_dim)\n        self.encoder2 = Encoder(num_inputs=input_dim2[1], num_units=hidden_dim)\n        \n# (in1 * in2) * (num_inputs * 64) \n# batchnorm(num_units)\n\n#         self.encoder2 = ChromEncoder(num_inputs=input_dim2, latent_dim=hidden_dim)\n\n        self.decoder1 = Decoder(\n            num_outputs=input_dim1,\n            num_units=hidden_dim,\n            final_activation=final_activations1,\n        )\n        self.decoder2 = Decoder(\n            num_outputs=input_dim2[1],\n            num_units=hidden_dim,\n            final_activation=final_activations1\n        )\n#         self.decoder2 = ChromDecoder(\n#             num_outputs=input_dim2,\n#             latent_dim=hidden_dim,\n#             final_activations=final_activations2,\n#         )\n\n\n    def split_catted_input(self, x):\n        \"\"\"Split the input into chunks that goes to each input to model\"\"\"\n        a, b = torch.split(x, [self.input_dim1, sum(self.input_dim2)], dim=-1)\n        return (a, torch.split(b, self.input_dim2, dim=-1))\n\n    def _combine_output_and_encoded(self, decoded, encoded, num_outputs: int):\n        \"\"\"\n        Combines the output and encoded in a single output\n        \"\"\"\n        if num_outputs > 1:\n            retval = *decoded, encoded\n        else:\n            if isinstance(decoded, tuple):\n                decoded = decoded[0]\n            retval = decoded, encoded\n        assert isinstance(retval, (list, tuple))\n        assert isinstance(\n            retval[0], (torch.TensorType, torch.Tensor)\n        ), f\"Expected tensor but got {type(retval[0])}\"\n        return retval\n    \n    def forward(self, x, size_factors=None, mode= None):\n#         if self.flat_mode:\n#             x = self.split_catted_input(x)\n        assert isinstance(x, (tuple, list))\n        assert len(x) == 2, \"There should be two inputs to spliced autoencoder\"\n        print(x)\n        encoded1 = self.encoder1(x[0])\n        encoded2 = self.encoder2(x[1])\n\n        decoded11 = self.decoder1(encoded1)\n        retval11 = self._combine_output_and_encoded(\n            decoded11, encoded1, self.num_outputs1\n        )\n        decoded12 = self.decoder2(encoded1)\n        retval12 = self._combine_output_and_encoded(\n            decoded12, encoded1, self.num_outputs2\n        )\n        decoded22 = self.decoder2(encoded2)\n        retval22 = self._combine_output_and_encoded(\n            decoded22, encoded2, self.num_outputs2\n        )\n        decoded21 = self.decoder1(encoded2)\n        retval21 = self._combine_output_and_encoded(\n            decoded21, encoded2, self.num_outputs1\n        )\n\n        if mode is None:\n            return retval11, retval12, retval21, retval22\n        retval_dict = {\n            (1, 1): retval11,\n            (1, 2): retval12,\n            (2, 1): retval21,\n            (2, 2): retval22,\n        }\n        if mode not in retval_dict:\n            raise ValueError(f\"Invalid mode code: {mode}\")\n        return retval_dict[mode]","metadata":{"execution":{"iopub.status.busy":"2022-11-09T08:12:16.216554Z","iopub.execute_input":"2022-11-09T08:12:16.216950Z","iopub.status.idle":"2022-11-09T08:12:16.240392Z","shell.execute_reply.started":"2022-11-09T08:12:16.216885Z","shell.execute_reply":"2022-11-09T08:12:16.239265Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"features0 = multi_train_x.shape[1]\nfeatures1 = multi_train_y_pca.shape[1]\nprint(features0)\nprint(features1)\nimport logging\n\n\nmodel = Autoencoder(input_dim1=features0, input_dim2=(features0, features1))\n","metadata":{"execution":{"iopub.status.busy":"2022-11-09T08:12:16.242409Z","iopub.execute_input":"2022-11-09T08:12:16.242874Z","iopub.status.idle":"2022-11-09T08:12:16.264211Z","shell.execute_reply.started":"2022-11-09T08:12:16.242829Z","shell.execute_reply":"2022-11-09T08:12:16.262875Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"m_train_x = pd.read_hdf(FP_MULTIOME_TRAIN_INPUTS, start=start, stop=stop).values\nm_train_y = pd.read_hdf(FP_MULTIOME_TRAIN_TARGETS, start=start, stop=stop).values","metadata":{"execution":{"iopub.status.busy":"2022-11-09T08:12:16.270199Z","iopub.execute_input":"2022-11-09T08:12:16.270624Z","iopub.status.idle":"2022-11-09T08:12:47.169988Z","shell.execute_reply.started":"2022-11-09T08:12:16.270587Z","shell.execute_reply":"2022-11-09T08:12:47.168955Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"optimizer = torch.optim.SGD(model.parameters(), lr=0.001, momentum=0.9)\nloss_fn = torch.nn.CrossEntropyLoss()","metadata":{"execution":{"iopub.status.busy":"2022-11-09T08:12:47.171279Z","iopub.execute_input":"2022-11-09T08:12:47.172641Z","iopub.status.idle":"2022-11-09T08:12:47.179751Z","shell.execute_reply.started":"2022-11-09T08:12:47.172541Z","shell.execute_reply":"2022-11-09T08:12:47.178299Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"training_loader = torch.utils.data.DataLoader([multi_train_x, multi_train_y_pca], batch_size=4, shuffle=True, num_workers=2)","metadata":{"execution":{"iopub.status.busy":"2022-11-09T08:29:50.348977Z","iopub.execute_input":"2022-11-09T08:29:50.349639Z","iopub.status.idle":"2022-11-09T08:29:50.356800Z","shell.execute_reply.started":"2022-11-09T08:29:50.349584Z","shell.execute_reply":"2022-11-09T08:29:50.355439Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"training_loader_2 = torch.utils.data.DataLoader(multi_train_y_pca, batch_size=4, shuffle=True, num_workers=2)","metadata":{"execution":{"iopub.status.busy":"2022-11-09T08:28:47.878784Z","iopub.execute_input":"2022-11-09T08:28:47.879273Z","iopub.status.idle":"2022-11-09T08:28:47.885690Z","shell.execute_reply.started":"2022-11-09T08:28:47.879239Z","shell.execute_reply":"2022-11-09T08:28:47.884662Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pip install skorch","metadata":{"execution":{"iopub.status.busy":"2022-11-09T09:14:07.893853Z","iopub.execute_input":"2022-11-09T09:14:07.894370Z","iopub.status.idle":"2022-11-09T09:14:20.197746Z","shell.execute_reply.started":"2022-11-09T09:14:07.894328Z","shell.execute_reply":"2022-11-09T09:14:20.195886Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import skorch\nimport skorch.utils \n\nclass PairedAutoEncoderSkorchNet(skorch.NeuralNet):\n    def forward_iter(self, X, training=False, device=\"cpu\"):\n        \"\"\"Subclassed to work with tuples\"\"\"\n        dataset = self.get_dataset(X)\n        iterator = self.get_iterator(dataset, training=training)\n        for data in iterator:\n            Xi = skorch.dataset.unpack_data(data)[0]\n            yp = self.evaluation_step(Xi, training=training)\n            if isinstance(yp, tuple):\n                yield model_utils.recursive_to_device(yp)  # <- modification here\n            else:\n                yield yp.to(device)\n\n    def predict_proba(self, x):\n        \"\"\"Subclassed so calling predict produces a tuple of outputs\"\"\"\n        y_probas1, y_probas2 = [], []\n        for yp in self.forward_iter(x, training=False):\n            assert isinstance(yp, tuple)\n            yp1 = yp[0][0]\n            yp2 = yp[1][0]\n            y_probas1.append(skorch.utils.to_numpy(yp1))\n            y_probas2.append(skorch.utils.to_numpy(yp2))\n        y_proba1 = np.concatenate(y_probas1, 0)\n        y_proba2 = np.concatenate(y_probas2, 0)\n        return y_proba1, y_proba2\n\n    def get_encoded_layer(self, x):\n        \"\"\"Get the encoded representation as a TUPLE of two elements\"\"\"\n        encoded1, encoded2 = [], []\n        for out1, out2, *_other in self.forward_iter(x, training=False):\n            encoded1.append(out1[-1])\n            encoded2.append(out2[-1])\n        return np.concatenate(encoded1, axis=0), np.concatenate(encoded2, axis=0)\n\n    def translate_1_to_2(self, x):\n        enc1, enc2 = self.get_encoded_layer(x)\n        device = next(self.module_.parameters()).device\n        enc1_torch = torch.from_numpy(enc1).to(device)\n        return self.module_.translate_1_to_2(enc1_torch)[0].detach().cpu().numpy()\n\n    def translate_2_to_1(self, x):\n        enc1, enc2 = self.get_encoded_layer(x)\n        device = next(self.module_.parameters()).device\n        enc2_torch = torch.from_numpy(enc2).to(device)\n        return self.module_.translate_2_to_1(enc2_torch)[0].detach().cpu().numpy()\n\n","metadata":{"execution":{"iopub.status.busy":"2022-11-09T09:14:38.248151Z","iopub.execute_input":"2022-11-09T09:14:38.248817Z","iopub.status.idle":"2022-11-09T09:14:38.267707Z","shell.execute_reply.started":"2022-11-09T09:14:38.248772Z","shell.execute_reply":"2022-11-09T09:14:38.266157Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from scipy import sparse\n","metadata":{"execution":{"iopub.status.busy":"2022-11-09T09:14:58.737147Z","iopub.execute_input":"2022-11-09T09:14:58.737603Z","iopub.status.idle":"2022-11-09T09:14:58.743800Z","shell.execute_reply.started":"2022-11-09T09:14:58.737565Z","shell.execute_reply":"2022-11-09T09:14:58.742418Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nclass SplicedAutoEncoderSkorchNet(PairedAutoEncoderSkorchNet):\n    \"\"\"\n    Skorch wrapper for the SplicedAutoEncoder above.\n    Mostly here to take care of how we calculate loss\n    \"\"\"\n\n    def predict_proba(self, x):\n        \"\"\"\n        Subclassed so that calling predict produces a tuple of 4 outputs\n        \"\"\"\n        y_probas1, y_probas2, y_probas3, y_probas4 = [], [], [], []\n        for yp in self.forward_iter(x, training=False):\n            assert isinstance(yp, tuple)\n            yp1 = yp[0][0]\n            yp2 = yp[1][0]\n            yp3 = yp[2][0]\n            yp4 = yp[3][0]\n            y_probas1.append(skorch.utils.to_numpy(yp1))\n            y_probas2.append(skorch.utils.to_numpy(yp2))\n            y_probas3.append(skorch.utils.to_numpy(yp3))\n            y_probas4.append(skorch.utils.to_numpy(yp4))\n        y_proba1 = np.concatenate(y_probas1)\n        y_proba2 = np.concatenate(y_probas2)\n        y_proba3 = np.concatenate(y_probas3)\n        y_proba4 = np.concatenate(y_probas4)\n        # Order: 1to1, 1to2, 2to1, 2to2\n        return y_proba1, y_proba2, y_proba3, y_proba4\n\n    def get_encoded_layer(self, x):\n        \"\"\"Get the encoded representation as a TUPLE of two elements\"\"\"\n        encoded1, encoded2 = [], []\n        for out11, out12, out21, out22 in self.forward_iter(x, training=False):\n            encoded1.append(out11[-1])\n            encoded2.append(out22[-1])\n        return np.concatenate(encoded1, axis=0), np.concatenate(encoded2, axis=0)\n\n    def translate_1_to_1(self, x) -> sparse.csr_matrix:\n        retval = [\n            sparse.csr_matrix(skorch.utils.to_numpy(yp[0][0]))\n            for yp in self.forward_iter(x, training=False)\n        ]\n        return sparse.vstack(retval)\n\n    def translate_1_to_2(self, x) -> sparse.csr_matrix:\n        retval = [\n            sparse.csr_matrix(skorch.utils.to_numpy(yp[1][0]))\n            for yp in self.forward_iter(x, training=False)\n        ]\n        return sparse.vstack(retval)\n\n    def translate_2_to_1(self, x) -> sparse.csr_matrix:\n        retval = [\n            sparse.csr_matrix(skorch.utils.to_numpy(yp[2][0]))\n            for yp in self.forward_iter(x, training=False)\n        ]\n        return sparse.vstack(retval)\n\n    def translate_2_to_2(self, x) -> sparse.csr_matrix:\n        retval = [\n            sparse.csr_matrix(skorch.utils.to_numpy(yp[3][0]))\n            for yp in self.forward_iter(x, training=False)\n        ]\n        return sparse.vstack(retval)\n\n    def score(self, true, pred):\n        \"\"\"\n        Required for sklearn gridsearch\n        Since sklearn uses convention of (true, pred) in its score functions\n        We use the same\n        https://scikit-learn.org/stable/modules/classes.html#module-sklearn.metrics\n        \"\"\"\n        return self.get_loss(pred, true)","metadata":{"execution":{"iopub.status.busy":"2022-11-09T09:15:00.387328Z","iopub.execute_input":"2022-11-09T09:15:00.387783Z","iopub.status.idle":"2022-11-09T09:15:00.405804Z","shell.execute_reply.started":"2022-11-09T09:15:00.387744Z","shell.execute_reply":"2022-11-09T09:15:00.404646Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"features0 = multi_train_x.shape[1]\nfeatures1 = multi_train_y_pca.shape[1]\nprint(features0)\nprint(features1)\n\n# model = Autoencoder(input_dim1=features0, input_dim2=(features0, features1))","metadata":{"execution":{"iopub.status.busy":"2022-11-09T09:20:55.879519Z","iopub.execute_input":"2022-11-09T09:20:55.880174Z","iopub.status.idle":"2022-11-09T09:20:55.887971Z","shell.execute_reply.started":"2022-11-09T09:20:55.880125Z","shell.execute_reply":"2022-11-09T09:20:55.886598Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\"\"\"\nLoss functions\n\"\"\"\n\nimport os\nimport sys\nimport functools\nfrom typing import *\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.tensorboard import SummaryWriter\n\n\nclass BCELoss(nn.BCELoss):\n    \"\"\"Custom BCE loss that can correctly ignore the encoded latent space output\"\"\"\n\n    def forward(self, x, target):\n        input = x[0]\n        return F.binary_cross_entropy(\n            input, target, weight=self.weight, reduction=self.reduction\n        )\n\n\nclass L1Loss(nn.L1Loss):\n    \"\"\"Custom L1 loss that ignores all but first input\"\"\"\n\n    def forward(self, x, target) -> torch.Tensor:\n        return F.l1_loss(x[0], target, reduction=self.reduction)\n\n\nclass ClassWeightedBCELoss(nn.BCELoss):\n    \"\"\"BCE that has different weight for 1/0 class\"\"\"\n\n    def __init__(self, class0_weight: float, class1_weight: float, **kwargs):\n        super(ClassWeightedBCELoss, self).__init__(**kwargs)\n        self.w0 = torch.tensor(class0_weight)\n        self.w1 = torch.tensor(class1_weight)\n\n    def forward(self, preds, target):\n        # This is batch size x num_features\n        bce = F.binary_cross_entropy(\n            preds, target, weight=self.weight, reduction=\"none\"\n        )\n        weights = torch.where(\n            target > 0, self.w1.to(preds.device), self.w0.to(preds.device)\n        )\n        retval = weights * bce\n        assert retval.shape == preds.shape\n        return torch.mean(retval)\n\n\nclass LogProbLoss(nn.Module):\n    \"\"\"\n    Log probability loss (originally written for RealNVP). Negates output (because log)\n\n    The prior needs to support a .log_prob(x) method\n    \"\"\"\n\n    def __init__(self, prior):\n        super(LogProbLoss, self).__init__()\n        self.prior = prior\n\n    def forward(self, x, _target=None):\n        z, logp = x[:2]\n        p = self.prior.log_prob(z)\n        if len(p.shape) == 2:\n            p = torch.mean(p, dim=1)\n        per_ex = p + logp\n        # assert len(per_ex) == z.shape[0]\n        retval = -torch.mean(per_ex)\n        if retval != retval:  # Detect NaN\n            raise ValueError(f\"Got NaN for loss with input z and logp: {z} {logp}\")\n        # if retval < 0:\n        #     raise ValueError(f\"Got negative loss with input z and logp: {z} {logp}\")\n        return retval\n\n\nclass DistanceProbLoss(nn.Module):\n    \"\"\"\n    Analog of above log prob loss, but using distances\n\n    May be useful for aligning latent spaces\n    \"\"\"\n\n    def __init__(self, weight: float = 5.0, norm: int = 1):\n        super(DistanceProbLoss, self).__init__()\n        assert weight > 0\n        self.weight = weight\n        self.norm = norm\n\n    def forward(self, x, target_z):\n        z, logp = x[:2]\n        d = F.pairwise_distance(\n            z,\n            target_z,\n            p=self.norm,\n            eps=1e-6,\n            keepdim=False,  # Default value\n        )\n        if len(d.shape) == 2:\n            d = torch.mean(d, dim=1)  # Drop 1 dimension\n        per_ex = self.weight * d - logp\n        retval = torch.mean(per_ex)\n        if retval != retval:\n            raise ValueError(\"NaN\")\n        return retval\n        return torch.mean(d)\n\n\nclass MSELoss(nn.MSELoss):\n    \"\"\"MSE loss\"\"\"\n\n    def forward(self, x: torch.Tensor, target: torch.Tensor) -> torch.Tensor:\n        return F.mse_loss(x[0], target, reduction=self.reduction)\n\n\nclass MSELogLoss(nn.modules.loss._Loss):\n    \"\"\"\n    MSE loss after applying log2\n\n    Based on:\n    https://pytorch.org/docs/stable/_modules/torch/nn/modules/loss.html#MSELoss\n    \"\"\"\n\n    __constants__ = [\"reduction\"]\n\n    def __init__(self, size_average=None, reduce=None, reduction=\"mean\"):\n        super(MSELogLoss, self).__init__(size_average, reduce, reduction)\n\n    def forward(self, input, target):\n        input_log = torch.log1p(input)\n        target_log = torch.log1p(target)\n        return F.mse_loss(input_log, target_log, reduction=self.reduction)\n\n\nclass MyNegativeBinomialLoss(nn.Module):\n    \"\"\"\n    Re-derived negative binomial loss.\n    \"\"\"\n\n    def __init__(self):\n        super(MyNegativeBinomialLoss, self).__init__()\n\n    def forward(self, preds, target):\n        preds, theta = preds[:2]\n        # Compare to:\n        # reconst_loss = vae.get_reconstruction_loss(sample_batch, px_rate, px_r, px_dropout)\n        l = -scvi_log_nb_positive(target, preds, theta)\n        return l\n\n\nclass NegativeBinomialLoss(nn.Module):\n    \"\"\"\n    Negative binomial loss. Preds should be a tuple of (mean, dispersion)\n    \"\"\"\n\n    def __init__(\n        self,\n        scale_factor: float = 1.0,\n        eps: float = 1e-10,\n        l1_lambda: float = 0.0,\n        mean: bool = True,\n    ):\n        super(NegativeBinomialLoss, self).__init__()\n        self.loss = negative_binom_loss(\n            scale_factor=scale_factor,\n            eps=eps,\n            mean=mean,\n            debug=True,\n        )\n        self.l1_lambda = l1_lambda\n\n    def forward(self, preds, target):\n        preds, theta = preds[:2]\n        l = self.loss(\n            preds=preds,\n            theta=theta,\n            truth=target,\n        )\n        encoded = preds[:-1]\n        l += self.l1_lambda * torch.abs(encoded).sum()\n        return l\n\n\nclass ZeroInflatedNegativeBinomialLoss(nn.Module):\n    \"\"\"\n    ZINB loss. Preds should be a tuple of (mean, dispersion, dropout)\n\n    General notes:\n    total variation seems to do poorly (at least for atacseq)\n    \"\"\"\n\n    def __init__(\n        self,\n        ridge_lambda: float = 0.0,\n        tv_lambda: float = 0.0,\n        l1_lambda: float = 0.0,\n        eps: float = 1e-10,\n        scale_factor: float = 1.0,\n        debug: bool = True,\n    ):\n        super(ZeroInflatedNegativeBinomialLoss, self).__init__()\n        self.loss = zero_inflated_negative_binom_loss(\n            ridge_lambda=ridge_lambda,\n            tv_lambda=tv_lambda,\n            eps=eps,\n            scale_factor=scale_factor,\n            debug=debug,\n        )\n        self.l1_lambda = l1_lambda\n\n    def forward(self, preds, target):\n        preds, theta, pi = preds[:3]\n        l = self.loss(\n            preds=preds,\n            theta_disp=theta,\n            pi_dropout=pi,\n            truth=target,\n        )\n        encoded = preds[:-1]\n        l += self.l1_lambda * torch.abs(encoded).sum()\n        return l\n\n\nclass MyZeroInflatedNegativeBinomialLoss(nn.Module):\n    \"\"\"\n    ZINB loss, based on scvi\n    \"\"\"\n\n    def forward(self, preds, target):\n        preds, theta, pi = preds[:3]\n        l = -scvi_log_zinb_positive(target, preds, theta, pi)\n        return l\n\n\nclass PairedLoss(nn.Module):\n    \"\"\"\n    Paired loss function. Automatically unpacks and encourages the encoded representation to be similar\n    using a given distance function. link_strength parameter controls how strongly we encourage this\n    loss2_weight controls how strongly we weight the second loss, relative to the first\n    A value of 1.0 indicates that they receive equal weight, and a value larger indicates\n    that the second loss receives greater weight.\n\n    link_func should be a callable that takes in the two encoded representations and outputs a metric\n    where a larger value indicates greater divergence\n    \"\"\"\n\n    def __init__(\n        self,\n        loss1=NegativeBinomialLoss,\n        loss2=ZeroInflatedNegativeBinomialLoss,\n        link_func=lambda x, y: (x - y).abs().mean(),\n        link_strength=1e-3,\n    ):\n        super(PairedLoss, self).__init__()\n        self.loss1 = loss1()\n        self.loss2 = loss2()\n        self.link = link_strength\n        self.link_f = link_func\n\n        self.warmup = layers.SigmoidWarmup(\n            midpoint=1000,\n            maximum=link_strength,\n        )\n\n    def forward(self, preds, target):\n        \"\"\"Unpack and feed to each loss, averaging at end\"\"\"\n        preds1, preds2 = preds\n        target1, target2 = target\n\n        loss1 = self.loss1(preds1, target1)\n        loss2 = self.loss2(preds2, target2)\n        retval = loss1 + loss2\n\n        # Align the encoded representation assuming the last output is encoded representation\n        encoded1 = preds1[-1]\n        encoded2 = preds2[-1]\n        if self.link > 0:\n            l = next(self.warmup)\n            if l > 1e-6:\n                d = self.link_f(encoded1, encoded2).mean()\n                retval += l * d\n\n        return retval\n\n\nclass PairedLossInvertible(nn.Module):\n    \"\"\"\n    Paired loss function with additional invertible (RealNVP) layer loss\n    Loss 1 is for the first autoencoder\n    Loss 2 is for the second autoencoder\n    Loss 3 is for the invertible network at bottleneck\n    \"\"\"\n\n    def __init__(\n        self,\n        loss1=NegativeBinomialLoss,\n        loss2=ZeroInflatedNegativeBinomialLoss,\n        loss3=DistanceProbLoss,\n        link_func=lambda x, y: (x - y).abs().mean(),\n        link_strength=1e-3,\n        inv_strength=1.0,\n    ):\n        super(PairedLossInvertible, self).__init__()\n        self.loss1 = loss1()\n        self.loss2 = loss2()\n        self.loss3 = loss3()\n        self.link = link_strength\n        self.link_f = link_func\n\n        # self.link_warmup = layers.SigmoidWarmup(\n        #     midpoint=1000,\n        #     maximum=link_strength,\n        # )\n        self.link_warmup = layers.DelayedLinearWarmup(\n            delay=1000,\n            inc=5e-3,\n            t_max=link_strength,\n        )\n\n        self.inv_warmup = layers.DelayedLinearWarmup(\n            delay=2000,\n            inc=5e-3,\n            t_max=inv_strength,\n        )\n\n    def forward(self, preds, target):\n        \"\"\"Unpack and feed to each loss\"\"\"\n        # Both enc1_pred and enc2_pred are tuples of 2 values\n        preds1, preds2, (enc1_pred, enc2_pred) = preds\n        target1, target2 = target\n\n        loss1 = self.loss1(preds1, target1)\n        loss2 = self.loss2(preds2, target2)\n        retval = loss1 + loss2\n\n        # Align the encoded representations\n        encoded1 = preds1[-1]\n        encoded2 = preds2[-1]\n        if self.link > 0:\n            l = next(self.link_warmup)\n            if l > 1e-6:\n                d = self.link_f(encoded1, encoded2).mean()\n                retval += l * d\n\n        # Add a term for invertible network\n        inv_loss1 = self.loss3(enc1_pred, enc2_pred[0])\n        inv_loss2 = self.loss3(enc2_pred, enc1_pred[0])\n        retval += next(self.inv_warmup) * (inv_loss1 + inv_loss2)\n\n        return retval\n\n\nclass QuadLoss(PairedLoss):\n    \"\"\"\n    Paired loss, but for the spliced autoencoder with 4 outputs\n    \"\"\"\n\n    def __init__(\n        self,\n        loss1=NegativeBinomialLoss,\n        loss2=BCELoss,\n        loss2_weight: float = 3.0,\n        cross_weight: float = 1.0,\n        cross_warmup_delay: int = 0,\n        link_strength: float = 0.0,\n        link_func: Callable = lambda x, y: (x - y).abs().mean(),\n        link_warmup_delay: int = 0,\n        record_history: bool = False,\n    ):\n        super(QuadLoss, self).__init__()\n        self.loss1 = loss1()\n        self.loss2 = loss2()\n        self.loss2_weight = loss2_weight\n        self.history = []  # Eventually contains list of tuples per call\n        self.record_history = record_history\n\n        if link_warmup_delay:\n            self.warmup = layers.SigmoidWarmup(\n                midpoint=link_warmup_delay,\n                maximum=link_strength,\n            )\n            # self.warmup = layers.DelayedLinearWarmup(\n            #     delay=warmup_delay,\n            #     t_max=link_strength,\n            #     inc=1e-3,\n            # )\n        else:\n            self.warmup = layers.NullWarmup(t_max=link_strength)\n        if cross_warmup_delay:\n            self.cross_warmup = layers.SigmoidWarmup(\n                midpoint=cross_warmup_delay,\n                maximum=cross_weight,\n            )\n        else:\n            self.cross_warmup = layers.NullWarmup(t_max=cross_weight)\n\n        self.link_strength = link_strength\n        self.link_func = link_func\n\n    def get_component_losses(self, preds, target):\n        \"\"\"\n        Return the four losses that go into the overall loss, without scaling\n        \"\"\"\n        preds11, preds12, preds21, preds22 = preds\n        if not isinstance(target, (list, tuple)):\n            # Try to unpack into the correct parts\n            target = torch.split(\n                target, [preds11[0].shape[-1], preds22[0].shape[-1]], dim=-1\n            )\n        target1, target2 = target  # Both are torch tensors\n\n        loss11 = self.loss1(preds11, target1)\n        loss21 = self.loss1(preds21, target1)\n        loss12 = self.loss2(preds12, target2)\n        loss22 = self.loss2(preds22, target2)\n\n        return loss11, loss21, loss12, loss22\n\n    def forward(self, preds, target):\n        loss11, loss21, loss12, loss22 = self.get_component_losses(preds, target)\n        if self.record_history:\n            detensor = lambda x: x.detach().cpu().numpy().item()\n            self.history.append([detensor(l) for l in (loss11, loss21, loss12, loss22)])\n\n        loss = loss11 + self.loss2_weight * loss22\n        loss += next(self.cross_warmup) * (loss21 + self.loss2_weight * loss12)\n\n        if self.link_strength > 0:\n            l = next(self.warmup)\n            if l > 1e-6:  # If too small we disregard\n                preds11, preds12, preds21, preds22 = preds\n                encoded1 = preds11[-1]  # Could be preds12\n                encoded2 = preds22[-1]  # Could be preds21\n                d = self.link_func(encoded1, encoded2)\n                loss += self.link_strength * d\n        return loss\n\n\ndef scvi_log_nb_positive(x, mu, theta, eps=1e-8):\n    \"\"\"\n    Taken from scVI log_likelihood.py - scVI invocation is:\n    reconst_loss = -log_nb_positive(x, px_rate, px_r).sum(dim=-1)\n    scVI decoder outputs px_scale, px_r, px_rate, px_dropout\n    px_scale is subject to Softmax\n    px_r is just a Linear layer\n    px_rate = torch.exp(library) * px_scale\n\n    mu = mean of NB\n    theta = indverse dispersion parameter\n\n    Here, x appears to correspond to y_true in the below negative_binom_loss (aka the observed counts)\n    \"\"\"\n    # if theta.ndimension() == 1:\n    #     theta = theta.view(\n    #         1, theta.size(0)\n    #     )  # In this case, we reshape theta for broadcasting\n\n    log_theta_mu_eps = torch.log(theta + mu + eps)\n    res = (\n        theta * (torch.log(theta + eps) - log_theta_mu_eps)\n        + x * (torch.log(mu + eps) - log_theta_mu_eps)\n        + torch.lgamma(x + theta)\n        - torch.lgamma(theta)  # Present (in negative) for DCA\n        - torch.lgamma(x + 1)\n    )\n\n    return res.mean()\n\n\ndef scvi_log_zinb_positive(x, mu, theta, pi, eps=1e-8):\n    \"\"\"\n    https://github.com/YosefLab/scVI/blob/6c9f43e3332e728831b174c1c1f0c9127b77cba0/scvi/models/log_likelihood.py#L206\n    \"\"\"\n    # theta is the dispersion rate. If .ndimension() == 1, it is shared for all cells (regardless of batch or labels)\n    if theta.ndimension() == 1:\n        theta = theta.view(\n            1, theta.size(0)\n        )  # In this case, we reshape theta for broadcasting\n\n    softplus_pi = F.softplus(-pi)\n    log_theta_eps = torch.log(theta + eps)\n    log_theta_mu_eps = torch.log(theta + mu + eps)\n    pi_theta_log = -pi + theta * (log_theta_eps - log_theta_mu_eps)\n\n    case_zero = F.softplus(pi_theta_log) - softplus_pi\n    mul_case_zero = torch.mul((x < eps).type(torch.float32), case_zero)\n\n    case_non_zero = (\n        -softplus_pi\n        + pi_theta_log\n        + x * (torch.log(mu + eps) - log_theta_mu_eps)  # Found above\n        + torch.lgamma(x + theta)  # Found above\n        - torch.lgamma(theta)  # Found above\n        - torch.lgamma(x + 1)  # Found above\n    )\n    mul_case_non_zero = torch.mul((x > eps).type(torch.float32), case_non_zero)\n\n    res = mul_case_zero + mul_case_non_zero\n\n    return res.mean()\n\n\ndef negative_binom_loss(\n    scale_factor: float = 1.0,\n    eps: float = 1e-10,\n    mean: bool = True,\n    debug: bool = False,\n    tb: SummaryWriter = None,\n) -> Callable:\n    \"\"\"\n    Return a function that calculates the binomial loss\n    https://github.com/theislab/dca/blob/master/dca/loss.py\n\n    combination of the Poisson distribution and a gamma distribution is a negative binomial distribution\n    \"\"\"\n\n    def loss(preds, theta, truth, tb_step: int = None):\n        \"\"\"Calculates negative binomial loss as defined in the NB class in link above\"\"\"\n        y_true = truth\n        y_pred = preds * scale_factor\n\n        if debug:  # Sanity check before loss calculation\n            assert not torch.isnan(y_pred).any(), y_pred\n            assert not torch.isinf(y_pred).any(), y_pred\n            assert not (y_pred < 0).any()  # should be non-negative\n            assert not (theta < 0).any()\n\n        # Clip theta values\n        theta = torch.clamp(theta, max=1e6)\n\n        t1 = (\n            torch.lgamma(theta + eps)\n            + torch.lgamma(y_true + 1.0)\n            - torch.lgamma(y_true + theta + eps)\n        )\n        t2 = (theta + y_true) * torch.log1p(y_pred / (theta + eps)) + (\n            y_true * (torch.log(theta + eps) - torch.log(y_pred + eps))\n        )\n        if debug:  # Sanity check after calculating loss\n            assert not torch.isnan(t1).any(), t1\n            assert not torch.isinf(t1).any(), (t1, torch.sum(torch.isinf(t1)))\n            assert not torch.isnan(t2).any(), t2\n            assert not torch.isinf(t2).any(), t2\n\n        retval = t1 + t2\n        if debug:\n            assert not torch.isnan(retval).any(), retval\n            assert not torch.isinf(retval).any(), retval\n\n        if tb is not None and tb_step is not None:\n            tb.add_histogram(\"nb/t1\", t1, global_step=tb_step)\n            tb.add_histogram(\"nb/t2\", t2, global_step=tb_step)\n\n        return torch.mean(retval) if mean else retval\n\n    return loss\n\n\ndef zero_inflated_negative_binom_loss(\n    ridge_lambda: float = 0.0,\n    tv_lambda: float = 0.0,\n    eps: float = 1e-10,\n    scale_factor: float = 1.0,\n    debug: bool = False,\n    tb: SummaryWriter = None,\n):\n    \"\"\"\n    Return a function that calculates ZINB loss\n    https://github.com/theislab/dca/blob/master/dca/loss.py\n    \"\"\"\n    nb_loss_func = negative_binom_loss(\n        mean=False, eps=eps, scale_factor=scale_factor, debug=debug, tb=tb\n    )\n\n    def loss(preds, theta_disp, pi_dropout, truth, tb_step: int = None):\n        if debug:\n            assert not (pi_dropout > 1.0).any()\n            assert not (pi_dropout < 0.0).any()\n        nb_case = nb_loss_func(preds, theta_disp, truth, tb_step=tb_step) - torch.log(\n            1.0 - pi_dropout + eps\n        )\n\n        y_true = truth\n        y_pred = preds * scale_factor\n        theta = torch.clamp(theta_disp, max=1e6)\n\n        zero_nb = torch.pow(theta / (theta + y_pred + eps), theta)\n        zero_case = -torch.log(pi_dropout + ((1.0 - pi_dropout) * zero_nb) + eps)\n        result = torch.where(y_true < 1e-8, zero_case, nb_case)\n\n        # Ridge regularization on pi dropout term\n        ridge = ridge_lambda * torch.pow(pi_dropout, 2)\n        result += ridge\n\n        # Total variation regularization on pi dropout term\n        tv = tv_lambda * total_variation(pi_dropout)\n        result += tv\n\n        if tb is not None and tb_step is not None:\n            tb.add_histogram(\"zinb/nb_case\", nb_case, global_step=tb_step)\n            tb.add_histogram(\"zinb/zero_nb\", zero_nb, global_step=tb_step)\n            tb.add_histogram(\"zinb/zero_case\", zero_case, global_step=tb_step)\n            tb.add_histogram(\"zinb/ridge\", ridge, global_step=tb_step)\n            tb.add_histogram(\"zinb/zinb_loss\", result, global_step=tb_step)\n\n        retval = torch.mean(result)\n        # if debug:\n        #     assert retval.item() > 0\n        return retval\n\n    return loss\n\n\ndef mmd(x, y):\n    \"\"\"\n    Compute maximum mean discrepancy\n\n    References:\n    https://ermongroup.github.io/blog/a-tutorial-on-mmd-variational-autoencoders/\n    https://github.com/napsternxg/pytorch-practice/blob/master/Pytorch%20-%20MMD%20VAE.ipynb\n    \"\"\"\n\n    def compute_kernel(x, y):\n        x_size = x.size(0)\n        y_size = y.size(0)\n        dim = x.size(1)\n        x = x.unsqueeze(1)  # (x_size, 1, dim)\n        y = y.unsqueeze(0)  # (1, y_size, dim)\n        tiled_x = x.expand(x_size, y_size, dim)\n        tiled_y = y.expand(x_size, y_size, dim)\n        kernel_input = (tiled_x - tiled_y).pow(2).mean(2) / float(dim)\n        return torch.exp(-kernel_input)  # (x_size, y_size)\n\n    x_kernel = compute_kernel(x, x)\n    y_kernel = compute_kernel(y, y)\n    xy_kernel = compute_kernel(x, y)\n    mmd = x_kernel.mean() + y_kernel.mean() - 2 * xy_kernel.mean()\n    return mmd\n\n\ndef total_variation(x):\n    \"\"\"\n    Given a 2D input (where one dimension is a batch dimension, the actual values are\n    one dimensional) compute the total variation (within a 1 position shift)\n    \"\"\"\n    t = torch.sum(torch.abs(x[:, :-1] - x[:, 1:]))\n    return t\n\n\nif __name__ == \"__main__\":\n    MSELogLoss()","metadata":{"execution":{"iopub.status.busy":"2022-11-09T09:24:05.146750Z","iopub.execute_input":"2022-11-09T09:24:05.147235Z","iopub.status.idle":"2022-11-09T09:24:05.243858Z","shell.execute_reply.started":"2022-11-09T09:24:05.147200Z","shell.execute_reply":"2022-11-09T09:24:05.242482Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"optim = torch.optim.Adam(model.parameters(), lr=0.001)","metadata":{"execution":{"iopub.status.busy":"2022-11-09T09:45:26.144891Z","iopub.execute_input":"2022-11-09T09:45:26.145468Z","iopub.status.idle":"2022-11-09T09:45:26.155353Z","shell.execute_reply.started":"2022-11-09T09:45:26.145425Z","shell.execute_reply":"2022-11-09T09:45:26.153834Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"REDUCE_LR_ON_PLATEAU_PARAMS = {\n    \"mode\": \"min\",\n    \"factor\": 0.1,\n    \"patience\": 10,\n    \"min_lr\": 1e-6,\n}","metadata":{"execution":{"iopub.status.busy":"2022-11-09T09:47:53.387132Z","iopub.execute_input":"2022-11-09T09:47:53.387596Z","iopub.status.idle":"2022-11-09T09:47:53.394566Z","shell.execute_reply.started":"2022-11-09T09:47:53.387559Z","shell.execute_reply":"2022-11-09T09:47:53.393395Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.utils.data import Dataset\n\n\nclass SplicedDataset(Dataset):\n    \"\"\"\n    Combines two datasets into one, where the first denotes X and the second denotes Y.\n    A spliced datset indicates that the inputs of x should predict the outputs of y.\n    Tries to match var names when possible\n    Flat mode also assumes that the input datasets are also flattened/catted\n    \"\"\"\n\n    def __init__(self, dataset_x, dataset_y, flat_mode: bool = False):\n        assert isinstance(\n            dataset_x, Dataset\n        ), f\"Bad type for dataset_x: {type(dataset_x)}\"\n        assert isinstance(\n            dataset_y, Dataset\n        ), f\"Bad type for dataset_y: {type(dataset_y)}\"\n        assert len(dataset_x) == len(dataset_y), \"Mismatched length\"\n        self.flat_mode = flat_mode\n\n        self.obs_names = None\n        x_obs_names = obs_names_from_dataset(dataset_x)\n        y_obs_names = obs_names_from_dataset(dataset_y)\n        if x_obs_names is not None and y_obs_names is not None:\n            logging.info(\"Checking obs names for two input datasets\")\n            for i, (x, y) in enumerate(zip(x_obs_names, y_obs_names)):\n                if x != y:\n                    raise ValueError(\n                        f\"Datasets have a different label at the {i}th index: {x} {y}\"\n                    )\n            self.obs_names = list(x_obs_names)\n        elif x_obs_names is not None:\n            self.obs_names = x_obs_names\n        elif y_obs_names is not None:\n            self.obs_names = y_obs_names\n        else:\n            raise ValueError(\"Both components of combined dataset appear to be dummy\")\n\n        self.dataset_x = dataset_x\n        self.dataset_y = dataset_y\n\n    def get_feature_labels(self) -> List[str]:\n        \"\"\"Return the names of the combined features\"\"\"\n        return list(self.dataset_x.data_raw.var_names) + list(\n            self.dataset_y.data_raw.var_names\n        )\n\n    def get_obs_labels(self) -> List[str]:\n        \"\"\"Return the names of each example\"\"\"\n        return self.obs_names\n\n    def __len__(self):\n        return len(self.dataset_x)\n\n    def __getitem__(self, i):\n        \"\"\"Assumes both return a single output\"\"\"\n        pair = (self.dataset_x[i][0], self.dataset_y[i][1])\n        if self.flat_mode:\n            raise NotImplementedError(f\"Flat mode is not defined for spliced dataset\")\n        return pair\n\nclass PairedDataset(SplicedDataset):\n    \"\"\"\n    Combines two datasets into one, where input is now (x1, x2) and\n    output is (y1, y2). A Paired dataset simply combines x and y\n    by returning the x input and y input as a tuple, and the x output\n    and y output as a tuple, and does not \"cross\" between the datasets\n    \"\"\"\n\n    # Inherits the init from SplicedDataset since we're doing the same thing - recording\n    # the two different datasets\n    def __getitem__(self, i):\n        x1 = self.dataset_x[i]\n        x2 = self.dataset_y[i]\n        x_pair = (x1[0], x2[0])\n        y_pair = (x1[1], x2[1])\n        if not self.flat_mode:\n            return x_pair, y_pair\n        else:\n            retval = torch.cat(x_pair), torch.cat(y_pair)\n            return \n        \n","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"    # what is dim1 and dim2 \n    # need to add loss functions\n    # change device\n    \n    spliced_net = SplicedAutoEncoderSkorchNet(\n            module=Autoencoder,\n            module__hidden_dim=16,  # Based on hyperparam tuning\n            module__input_dim1=features0,\n            module__input_dim2=(features0, features1),\n            module__final_activations1=[\n                Exp(),\n                ClippedSoftplus(),\n            ],\n            module__final_activations2=nn.Sigmoid(),\n            module__flat_mode=True,\n            lr=0.01,  # Based on hyperparam tuning\n            criterion=QuadLoss,\n            criterion__loss2=BCELoss,  # handle output of encoded layer\n            criterion__loss2_weight=0.5,  # numerically balance the two losses with different magnitudes\n            criterion__record_history=True,\n            optimizer=torch.optim.Adam,\n            iterator_train__shuffle=True,\n            batch_size=512,  # Based on  hyperparam tuning\n            max_epochs=500,\n            callbacks=[\n                skorch.callbacks.EarlyStopping(patience=25),\n                skorch.callbacks.LRScheduler(\n                    policy=torch.optim.lr_scheduler.ReduceLROnPlateau,\n                    **REDUCE_LR_ON_PLATEAU_PARAMS,\n                ),\n                skorch.callbacks.GradientNormClipping(gradient_clip_value=5),\n#                 skorch.callbacks.Checkpoint(\n#                     dirname=outdir_name, fn_prefix=\"net_\", monitor=\"valid_loss_best\",\n#                 ),\n            ],\n#             train_split=skorch.helper.predefined_split(sc_dual_valid_dataset),\n            iterator_train__num_workers=8,\n            iterator_valid__num_workers=8,\n        )","metadata":{"execution":{"iopub.status.busy":"2022-11-09T09:48:30.464546Z","iopub.execute_input":"2022-11-09T09:48:30.465164Z","iopub.status.idle":"2022-11-09T09:48:30.477023Z","shell.execute_reply.started":"2022-11-09T09:48:30.465093Z","shell.execute_reply":"2022-11-09T09:48:30.475670Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.utils.data import Dataset\n\nclass SplicedDataset(Dataset):\n    \"\"\"\n    Combines two datasets into one, where the first denotes X and the second denotes Y.\n    A spliced datset indicates that the inputs of x should predict the outputs of y.\n    Tries to match var names when possible\n    Flat mode also assumes that the input datasets are also flattened/catted\n    \"\"\"\n\n    def __init__(self, dataset_x, dataset_y, flat_mode: bool = False):\n        assert isinstance(\n            dataset_x, Dataset\n        ), f\"Bad type for dataset_x: {type(dataset_x)}\"\n        assert isinstance(\n            dataset_y, Dataset\n        ), f\"Bad type for dataset_y: {type(dataset_y)}\"\n        assert len(dataset_x) == len(dataset_y), \"Mismatched length\"\n        self.flat_mode = flat_mode\n\n        self.obs_names = None\n        x_obs_names = obs_names_from_dataset(dataset_x)\n        y_obs_names = obs_names_from_dataset(dataset_y)\n        if x_obs_names is not None and y_obs_names is not None:\n            logging.info(\"Checking obs names for two input datasets\")\n            for i, (x, y) in enumerate(zip(x_obs_names, y_obs_names)):\n                if x != y:\n                    raise ValueError(\n                        f\"Datasets have a different label at the {i}th index: {x} {y}\"\n                    )\n            self.obs_names = list(x_obs_names)\n        elif x_obs_names is not None:\n            self.obs_names = x_obs_names\n        elif y_obs_names is not None:\n            self.obs_names = y_obs_names\n        else:\n            raise ValueError(\"Both components of combined dataset appear to be dummy\")\n\n        self.dataset_x = dataset_x\n        self.dataset_y = dataset_y\n\n    def get_feature_labels(self) -> List[str]:\n        \"\"\"Return the names of the combined features\"\"\"\n        return list(self.dataset_x.data_raw.var_names) + list(\n            self.dataset_y.data_raw.var_names\n        )\n\n    def get_obs_labels(self) -> List[str]:\n        \"\"\"Return the names of each example\"\"\"\n        return self.obs_names\n\n    def __len__(self):\n        return len(self.dataset_x)\n\n    def __getitem__(self, i):\n        \"\"\"Assumes both return a single output\"\"\"\n        pair = (self.dataset_x[i][0], self.dataset_y[i][1])\n        if self.flat_mode:\n            raise NotImplementedError(f\"Flat mode is not defined for spliced dataset\")\n        return pair\n\n\nclass PairedDataset(SplicedDataset):\n    \"\"\"\n    Combines two datasets into one, where input is now (x1, x2) and\n    output is (y1, y2). A Paired dataset simply combines x and y\n    by returning the x input and y input as a tuple, and the x output\n    and y output as a tuple, and does not \"cross\" between the datasets\n    \"\"\"\n\n    # Inherits the init from SplicedDataset since we're doing the same thing - recording\n    # the two different datasets\n    def __getitem__(self, i):\n        x1 = self.dataset_x[i]\n        x2 = self.dataset_y[i]\n        x_pair = (x1[0], x2[0])\n        y_pair = (x1[1], x2[1])\n        if not self.flat_mode:\n            return x_pair, y_pair\n        else:\n            retval = torch.cat(x_pair), torch.cat(y_pair)\n            return retval\n","metadata":{"execution":{"iopub.status.busy":"2022-11-09T09:56:11.415432Z","iopub.execute_input":"2022-11-09T09:56:11.416104Z","iopub.status.idle":"2022-11-09T09:56:11.438132Z","shell.execute_reply.started":"2022-11-09T09:56:11.416041Z","shell.execute_reply":"2022-11-09T09:56:11.436790Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sc_dual_train_dataset = PairedDataset(\n        multi_train_y_pca, multi_train_x, flat_mode=True,\n    )\n\nspliced_net.fit(sc_dual_train_dataset, y=None)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pip install anndata","metadata":{"execution":{"iopub.status.busy":"2022-11-09T10:01:04.617241Z","iopub.execute_input":"2022-11-09T10:01:04.617829Z","iopub.status.idle":"2022-11-09T10:01:16.771987Z","shell.execute_reply.started":"2022-11-09T10:01:04.617787Z","shell.execute_reply":"2022-11-09T10:01:16.770085Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nfrom anndata import AnnData\nimport scanpy as sc\n\n@functools.lru_cache(4)\ndef sc_read_mtx(fname: str, dtype: str = \"float32\"):\n    \"\"\"Helper function for reading mtx files so we can cache the result\"\"\"\n    return sc.read_mtx(filename=fname, dtype=dtype)\n\n\nclass SingleCellDataset(Dataset):\n    \"\"\"\n    Given a sparse matrix file, load in dataset\n    If transforms is given, it is applied after all the pre-baked transformations. These\n    can be things like sklearn MaxAbsScaler().fit_transform\n    \"\"\"\n\n    def __init__(\n        self,\n        fname: Union[str, List[str]],\n        reader: Callable = sc_read_mtx,\n        raw_adata: Union[AnnData, None] = None,  # Should be raw data\n        transpose: bool = True,\n        mode: str = \"all\",\n        data_split_by_cluster: str = \"leiden\",  # Specify as leiden\n        valid_cluster_id: int = 0,  # Only used if data_split_by_cluster is on\n        test_cluster_id: int = 1,\n        data_split_by_cluster_log: bool = True,\n        predefined_split=None,  # of type SingleCellDataset\n        cell_info: pd.DataFrame = None,\n        gene_info: pd.DataFrame = None,\n        selfsupervise: bool = True,\n        binarize: bool = False,\n        filt_cell_min_counts=None,  # All of these are off by default\n        filt_cell_max_counts=None,\n        filt_cell_min_genes=None,\n        filt_cell_max_genes=None,\n        filt_gene_min_counts=None,\n        filt_gene_max_counts=None,\n        filt_gene_min_cells=None,\n        filt_gene_max_cells=None,\n        pool_genomic_interval: Union[int, List[str]] = 0,\n        calc_size_factors: bool = True,\n        normalize: bool = True,\n        log_trans: bool = True,\n        clip: float = 0,\n        sort_by_pos: bool = False,\n        split_by_chrom: bool = False,\n        concat_outputs: bool = False,  # Instead of outputting a list of tensors, concat\n        autosomes_only: bool = False,\n        # high_confidence_clustering_genes: List[str] = [],  # used to build clustering\n        x_dropout: bool = False,\n        y_mode: str = \"size_norm\",\n        sample_y: bool = False,\n        return_sf: bool = True,\n        return_pbulk: bool = False,\n        filter_features: dict = {},\n        filter_samples: dict = {},\n        transforms: List[Callable] = [],\n#         gtf_file: str = MM10_GTF,  # GTF file mapping genes to chromosomes, unused for ATAC\n        cluster_res: float = 2.0,\n        cache_prefix: str = \"\",\n    ):\n        \"\"\"\n        Clipping is performed AFTER normalization\n        Binarize will turn all counts into binary 0/1 indicators before running normalization code\n        If pool_genomic_interval is -1, then we pool based on proximity to gene\n        \"\"\"\n        assert mode in [\n            \"all\",\n            \"skip\",\n        ], \"SingleCellDataset now operates as a full dataset only. Use SingleCellDatasetSplit to define data splits\"\n        assert y_mode in [\n            \"size_norm\",\n            \"log_size_norm\",\n            \"raw_count\",\n            \"log_raw_count\",\n            \"x\",\n        ], f\"Unrecognized mode for y output: {y_mode}\"\n        if y_mode == \"size_norm\":\n            assert calc_size_factors\n        self.mode = mode\n        self.selfsupervise = selfsupervise\n        self.x_dropout = x_dropout\n        self.y_mode = y_mode\n        self.sample_y = sample_y\n        self.binarize = binarize\n        self.calc_size_factors = calc_size_factors\n        self.return_sf = return_sf\n        self.return_pbulk = return_pbulk\n        self.transforms = transforms\n        self.cache_prefix = cache_prefix\n        self.sort_by_pos = sort_by_pos\n        self.split_by_chrom = split_by_chrom\n        self.concat_outputs = concat_outputs\n        self.autosomes_only = autosomes_only\n        self.cluster_res = cluster_res\n        self.data_split_by_cluster = data_split_by_cluster\n        self.valid_cluster_id = valid_cluster_id\n        self.test_cluster_id = test_cluster_id\n        self.data_split_by_cluster_log = data_split_by_cluster_log\n\n        if raw_adata is not None:\n            logging.info(\n                f\"Got AnnData object {str(raw_adata)}, ignoring reader/fname args\"\n            )\n            self.data_raw = raw_adata\n        else:\n            self.data_raw = reader(fname)\n        assert isinstance(\n            self.data_raw, AnnData\n        ), f\"Expected AnnData but got {type(self.data_raw)}\"\n        if not isinstance(self.data_raw.X, scipy.sparse.csr_matrix):\n            self.data_raw.X = scipy.sparse.csr_matrix(\n                self.data_raw.X\n            )  # Convert to sparse matrix\n\n        if transpose:\n            self.data_raw = self.data_raw.T\n\n        # Filter out undesirable var/obs\n        # self.__filter_obs_metadata(filter_samples=filter_samples)\n        # self.__filter_var_metadata(filter_features=filter_features)\n        self.data_raw = adata_utils.filter_adata(\n            self.data_raw, filt_cells=filter_samples, filt_var=filter_features\n        )\n\n        # Attach obs/var annotations\n        if cell_info is not None:\n            assert isinstance(cell_info, pd.DataFrame)\n            if self.data_raw.obs is not None and not self.data_raw.obs.empty:\n                self.data_raw.obs = self.data_raw.obs.join(\n                    cell_info, how=\"left\", sort=False\n                )\n            else:\n                self.data_raw.obs = cell_info\n            assert (\n                self.data_raw.shape[0] == self.data_raw.obs.shape[0]\n            ), f\"Got discordant shapes for data and obs: {self.data_raw.shape} {self.data_raw.obs.shape}\"\n\n        if gene_info is not None:\n            assert isinstance(gene_info, pd.DataFrame)\n            if (\n                self.data_raw.var is not None and not self.data_raw.var.empty\n            ):  # Is not None and is not empty\n                self.data_raw.var = self.data_raw.var.join(\n                    gene_info, how=\"left\", sort=False\n                )\n            else:\n                self.data_raw.var = gene_info\n            assert (\n                self.data_raw.shape[1] == self.data_raw.var.shape[0]\n            ), f\"Got discordant shapes for data and var: {self.data_raw.shape} {self.data_raw.var.shape}\"\n\n        if sort_by_pos:\n            genes_reordered, chroms_reordered = reorder_genes_by_pos(\n                self.data_raw.var_names, gtf_file=gtf_file, return_chrom=True\n            )\n            self.data_raw = self.data_raw[:, genes_reordered]\n\n        self.__annotate_chroms(gtf_file)\n        if self.autosomes_only:\n            autosomal_idx = [\n                i\n                for i, chrom in enumerate(self.data_raw.var[\"chrom\"])\n                if utils.is_numeric(chrom.strip(\"chr\"))\n            ]\n            self.data_raw = self.data_raw[:, autosomal_idx]\n\n        # Sort by the observation names so we can combine datasets\n        sort_order_idx = np.argsort(self.data_raw.obs_names)\n        self.data_raw = self.data_raw[sort_order_idx, :]\n        # NOTE pooling occurs AFTER feature/observation filtering\n        if pool_genomic_interval:\n            self.__pool_features(pool_genomic_interval=pool_genomic_interval)\n            # Re-annotate because we have lost this information\n            self.__annotate_chroms(gtf_file)\n\n        # Preprocess the data now that we're done filtering\n        if self.binarize:\n            # If we are binarizing data we probably don't care about raw counts\n            # self.data_raw.raw = self.data_raw.copy()  # Store original counts\n            self.data_raw.X[self.data_raw.X.nonzero()] = 1  # .X here is a csr matrix\n\n        adata_utils.annotate_basic_adata_metrics(self.data_raw)\n        adata_utils.filter_adata_cells_and_genes(\n            self.data_raw,\n            filter_cell_min_counts=filt_cell_min_counts,\n            filter_cell_max_counts=filt_cell_max_counts,\n            filter_cell_min_genes=filt_cell_min_genes,\n            filter_cell_max_genes=filt_cell_max_genes,\n            filter_gene_min_counts=filt_gene_min_counts,\n            filter_gene_max_counts=filt_gene_max_counts,\n            filter_gene_min_cells=filt_gene_min_cells,\n            filter_gene_max_cells=filt_gene_max_cells,\n        )\n        self.data_raw = adata_utils.normalize_count_table(  # Normalizes in place\n            self.data_raw,\n            size_factors=calc_size_factors,\n            normalize=normalize,\n            log_trans=log_trans,\n        )\n\n        if clip > 0:\n            assert isinstance(clip, float) and 0.0 < clip < 50.0\n            logging.info(f\"Clipping to {clip} percentile\")\n            clip_low, clip_high = np.percentile(\n                self.data_raw.X.flatten(), [clip, 100.0 - clip]\n            )\n            if clip_low == clip_high == 0:\n                logging.warning(\"Skipping clipping, as clipping intervals are 0\")\n            else:\n                assert (\n                    clip_low < clip_high\n                ), f\"Got discordant values for clipping ends: {clip_low} {clip_high}\"\n                self.data_raw.X = np.clip(self.data_raw.X, clip_low, clip_high)\n\n        # Apply any final transformations\n        if self.transforms:\n            for trans in self.transforms:\n                self.data_raw.X = trans(self.data_raw.X)\n\n        # Make sure the data is a sparse matrix\n        if not isinstance(self.data_raw.X, scipy.sparse.csr_matrix):\n            self.data_raw.X = scipy.sparse.csr_matrix(self.data_raw.X)\n\n        # Do all normalization before we split to make sure all folds get the same normalization\n        self.data_split_to_idx = {}\n        if predefined_split is not None:\n            logging.info(\"Got predefined split, ignoring mode\")\n            # Subset items\n            self.data_raw = self.data_raw[\n                [\n                    i\n                    for i in predefined_split.data_raw.obs.index\n                    if i in self.data_raw.obs.index\n                ],\n            ]\n            assert (\n                self.data_raw.n_obs > 0\n            ), \"No intersected obs names from predefined split\"\n            # Carry over cluster indexing\n            self.data_split_to_idx = copy.copy(predefined_split.data_split_to_idx)\n        elif mode != \"skip\":\n            # Create dicts mapping string to list of indices\n            if self.data_split_by_cluster:\n                self.data_split_to_idx = self.__split_train_valid_test_cluster(\n                    clustering_key=self.data_split_by_cluster\n                    if isinstance(self.data_split_by_cluster, str)\n                    else \"leiden\",\n                    valid_cluster={str(self.valid_cluster_id)},\n                    test_cluster={str(self.test_cluster_id)},\n                )\n            else:\n                self.data_split_to_idx = self.__split_train_valid_test()\n        else:\n            logging.info(\"Got data split skip, skipping data split\")\n        self.data_split_to_idx[\"all\"] = np.arange(len(self.data_raw))\n\n        self.size_factors = (\n            torch.from_numpy(self.data_raw.obs.size_factors.values).type(\n                torch.FloatTensor\n            )\n            if self.return_sf\n            else None\n        )\n        self.cell_sim_mat = (\n            euclidean_sim_matrix(self.size_norm_counts) if self.sample_y else None\n        )  # Skip calculation if we don't need\n\n        # Perform file backing if necessary\n        self.data_raw_cache_fname = \"\"\n        if self.cache_prefix:\n            self.data_raw_cache_fname = self.cache_prefix + \".data_raw.h5ad\"\n            logging.info(f\"Setting cache at {self.data_raw_cache_fname}\")\n            self.data_raw.filename = self.data_raw_cache_fname\n            if hasattr(self, \"_size_norm_counts\"):\n                size_norm_cache_name = self.cache_prefix + \".size_norm_counts.h5ad\"\n                logging.info(\n                    f\"Setting size norm counts cache at {size_norm_cache_name}\"\n                )\n                self._size_norm_counts.filename = size_norm_cache_name\n            if hasattr(self, \"_size_norm_log_counts\"):\n                size_norm_log_cache_name = (\n                    self.cache_prefix + \".size_norm_log_counts.h5ad\"\n                )\n                logging.info(\n                    f\"Setting size log norm counts cache at {size_norm_log_cache_name}\"\n                )\n                self._size_norm_log_counts.filename = size_norm_log_cache_name\n\n    def __annotate_chroms(self, gtf_file: str = \"\") -> None:\n        \"\"\"Annotates chromosome information on the var field, without the chr prefix\"\"\"\n        # gtf_file can be empty if we're using atac intervals\n        feature_chroms = (\n            get_chrom_from_intervals(self.data_raw.var_names)\n            if list(self.data_raw.var_names)[0].startswith(\"chr\")\n            else get_chrom_from_genes(self.data_raw.var_names, gtf_file)\n        )\n        self.data_raw.var[\"chrom\"] = feature_chroms\n\n    def __pool_features(self, pool_genomic_interval: Union[int, List[str]]):\n        n_obs = self.data_raw.n_obs\n        if isinstance(pool_genomic_interval, int):\n            if pool_genomic_interval > 0:\n                # WARNING This will wipe out any existing var information\n                idx, names = get_indices_to_combine(\n                    list(self.data_raw.var.index), interval=pool_genomic_interval\n                )\n                data_raw_aggregated = combine_array_cols_by_idx(  # Returns np ndarray\n                    self.data_raw.X,\n                    idx,\n                )\n                data_raw_aggregated = scipy.sparse.csr_matrix(data_raw_aggregated)\n                self.data_raw = AnnData(\n                    data_raw_aggregated,\n                    obs=self.data_raw.obs,\n                    var=pd.DataFrame(index=names),\n                )\n            elif pool_genomic_interval < 0:\n                assert (\n                    pool_genomic_interval == -1\n                ), f\"Invalid value: {pool_genomic_interval}\"\n                # Pool based on proximity to genes\n                data_raw_aggregated, names = combine_by_proximity(self.data_raw)\n                self.data_raw = AnnData(\n                    data_raw_aggregated,\n                    obs=self.data_raw.obs,\n                    var=pd.DataFrame(index=names),\n                )\n            else:\n                raise ValueError(f\"Invalid integer value: {pool_genomic_interval}\")\n        elif isinstance(pool_genomic_interval, (list, set, tuple)):\n            idx = get_indices_to_form_target_intervals(\n                self.data_raw.var.index, target_intervals=pool_genomic_interval\n            )\n            data_raw_aggregated = scipy.sparse.csr_matrix(\n                combine_array_cols_by_idx(\n                    self.data_raw.X,\n                    idx,\n                )\n            )\n            self.data_raw = AnnData(\n                data_raw_aggregated,\n                obs=self.data_raw.obs,\n                var=pd.DataFrame(index=pool_genomic_interval),\n            )\n        else:\n            raise TypeError(\n                f\"Unrecognized type for pooling features: {type(pool_genomic_interval)}\"\n            )\n        assert self.data_raw.n_obs == n_obs\n\n    def __split_train_valid_test(self) -> Dict[str, List[int]]:\n        \"\"\"\n        Split the dataset into the appropriate split, returning the indices of split\n        \"\"\"\n        logging.warning(\n            f\"Constructing {self.mode} random data split - not recommended due to potential leakage between data split\"\n        )\n        indices = np.arange(self.data_raw.n_obs)\n        (train_idx, valid_idx, test_idx,) = shuffle_indices_train_valid_test(\n            indices, shuffle=True, seed=1234, valid=0.15, test=0.15\n        )\n        assert train_idx, \"Got empty training split\"\n        assert valid_idx, \"Got empty validation split\"\n        assert test_idx, \"Got empty test split\"\n        data_split_idx = {}\n        data_split_idx[\"train\"] = train_idx\n        data_split_idx[\"valid\"] = valid_idx\n        data_split_idx[\"test\"] = test_idx\n        return data_split_idx\n\n    def __split_train_valid_test_cluster(\n        self, clustering_key: str = \"leiden\", valid_cluster={\"0\"}, test_cluster={\"1\"}\n    ) -> Dict[str, List[int]]:\n        \"\"\"\n        Splits the dataset into appropriate split based on clustering\n        Retains similarly sized splits as train/valid/test random from above\n        \"\"\"\n        assert not valid_cluster.intersection(\n            test_cluster\n        ), f\"Overlap between valid and test clusters: {valid_cluster} {test_cluster}\"\n        if clustering_key not in [\"leiden\", \"louvain\"]:\n            raise ValueError(\n                f\"Invalid clustering key for data splits: {clustering_key}\"\n            )\n        logging.info(\n            f\"Constructing {clustering_key} {'log' if self.data_split_by_cluster_log else 'linear'} clustered data split with valid test cluster {valid_cluster} {test_cluster}\"\n        )\n        cluster_labels = (\n            self.size_norm_log_counts.obs[clustering_key]\n            if self.data_split_by_cluster_log\n            else self.size_norm_counts.obs[clustering_key]\n        )\n        cluster_labels_counter = collections.Counter(cluster_labels.to_list())\n        assert not valid_cluster.intersection(\n            test_cluster\n        ), \"Valid and test clusters overlap\"\n\n        train_idx, valid_idx, test_idx = [], [], []\n        for i, label in enumerate(cluster_labels):\n            if label in valid_cluster:\n                valid_idx.append(i)\n            elif label in test_cluster:\n                test_idx.append(i)\n            else:\n                train_idx.append(i)\n\n        assert train_idx, \"Got empty training split\"\n        assert valid_idx, \"Got empty validation split\"\n        assert test_idx, \"Got empty test split\"\n        data_split_idx = {}\n        data_split_idx[\"train\"] = train_idx\n        data_split_idx[\"valid\"] = valid_idx\n        data_split_idx[\"test\"] = test_idx\n        return data_split_idx\n\n    def __sample_similar_cell(self, i, threshold=5, leakage=0.1) -> None:\n        \"\"\"\n        Sample a similar cell for the ith cell\n        Uses a very naive approach where we separately sample things above\n        and below the threshold. 0.6 samples about 13.89 neighbors, 0.5 samples about 60.14\n        Returns index of that similar cell\n        \"\"\"\n\n        def exp_sample(sims):\n            \"\"\"Sample from the vector of similarities\"\"\"\n            w = np.exp(sims)\n            assert not np.any(np.isnan(w)), \"Got NaN in exp(s)\"\n            assert np.sum(w) > 0, \"Got a zero-vector of weights!\"\n            w_norm = w / np.sum(w)\n            idx = np.random.choice(np.arange(len(w_norm)), p=w_norm)\n            return idx\n\n        assert self.cell_sim_mat is not None\n        sim_scores = self.cell_sim_mat[i]\n        high_scores = sim_scores[np.where(sim_scores > threshold)]\n        low_scores = sim_scores[np.where(sim_scores <= threshold)]\n        if np.random.random() < leakage:\n            idx = exp_sample(low_scores)\n        else:\n            idx = exp_sample(high_scores)\n        return idx\n\n    @functools.lru_cache(32)\n    def __get_chrom_idx(self) -> Dict[str, np.ndarray]:\n        \"\"\"Helper func for figuring out which feature indexes are on each chromosome\"\"\"\n        chromosomes = sorted(\n            list(set(self.data_raw.var[\"chrom\"]))\n        )  # Sort to guarantee consistent ordering\n        chrom_to_idx = collections.OrderedDict()\n        for chrom in chromosomes:\n            chrom_to_idx[chrom] = np.where(self.data_raw.var[\"chrom\"] == chrom)\n        return chrom_to_idx\n\n    def __get_chrom_split_features(self, i):\n        \"\"\"Given an index, get the features split by chromsome, returning in chromosome-sorted order\"\"\"\n        if self.x_dropout:\n            raise NotImplementedError\n        features = torch.from_numpy(\n            utils.ensure_arr(self.data_raw.X[i]).flatten()\n        ).type(torch.FloatTensor)\n        assert len(features.shape) == 1  # Assumes one dimensional vector of features\n\n        chrom_to_idx = self.__get_chrom_idx()\n        retval = tuple([features[indices] for _chrom, indices in chrom_to_idx.items()])\n        if self.concat_outputs:\n            retval = torch.cat(retval)\n        return retval\n\n    def __len__(self):\n        \"\"\"Number of examples\"\"\"\n        return self.data_raw.n_obs\n\n    def get_item_data_split(self, idx: int, split: str):\n        \"\"\"Get the i-th item in the split (e.g. train)\"\"\"\n        assert split in [\"train\", \"valid\", \"test\", \"all\"]\n        if split == \"all\":\n            return self.__getitem__(idx)\n        else:\n            return self.__getitem__(self.data_split_to_idx[split][idx])\n\n    def __getitem__(self, i):\n        # TODO compatibility with slices\n        expression_data = (\n            torch.from_numpy(utils.ensure_arr(self.data_raw.X[i]).flatten()).type(\n                torch.FloatTensor\n            )\n            if not self.split_by_chrom\n            else self.__get_chrom_split_features(i)\n        )\n        if self.x_dropout and not self.split_by_chrom:\n            # Apply dropout to the X input\n            raise NotImplementedError\n\n        # Handle case where we are shuffling y a la noise2noise\n        # Only use shuffled indices if it is specifically enabled and if we are doing TRAINING\n        # I.e. validation/test should never be shuffled\n        y_idx = (\n            self.__sample_similar_cell(i)\n            if (self.sample_y and self.mode == \"train\")\n            else i  # If not sampling y and training, return the same idx\n        )\n        if self.y_mode.endswith(\"raw_count\"):\n            target = torch.from_numpy(\n                utils.ensure_arr(self.data_raw.raw.var_vector(y_idx))\n            ).type(torch.FloatTensor)\n        elif self.y_mode.endswith(\"size_norm\"):\n            target = torch.from_numpy(self.size_norm_counts.var_vector(y_idx)).type(\n                torch.FloatTensor\n            )\n        elif self.y_mode == \"x\":\n            target = torch.from_numpy(\n                utils.ensure_arr(self.data_raw.X[y_idx]).flatten()\n            ).type(torch.FloatTensor)\n        else:\n            raise NotImplementedError(f\"Unrecognized y_mode: {self.y_mode}\")\n        if self.y_mode.startswith(\"log\"):\n            target = torch.log1p(target)  # scapy is also natural logaeritm of 1+x\n\n        # Structure here is a series of inputs, followed by a fixed tuple of expected output\n        retval = [expression_data]\n        if self.return_sf:\n            sf = self.size_factors[i]\n            retval.append(sf)\n        # Build expected truth\n        if self.selfsupervise:\n            if not self.return_pbulk:\n                retval.append(target)\n            else:  # Return both target and psuedobulk\n                ith_cluster = self.data_raw.obs.iloc[i][\"leiden\"]\n                pbulk = torch.from_numpy(\n                    self.get_cluster_psuedobulk().var_vector(ith_cluster)\n                ).type(torch.FloatTensor)\n                retval.append((target, pbulk))\n        elif self.return_pbulk:\n            ith_cluster = self.data_raw.obs.iloc[i][\"leiden\"]\n            pbulk = torch.from_numpy(\n                self.get_cluster_psuedobulk().var_vector(ith_cluster)\n            ).type(torch.FloatTensor)\n            retval.append(pbulk)\n        else:\n            raise ValueError(\"Neither selfsupervise or retur_pbulk is specified\")\n\n        return tuple(retval)\n\n    def get_per_chrom_feature_count(self) -> List[int]:\n        \"\"\"\n        Return the number of features from each chromosome\n        If we were to split a catted feature vector, we would split\n        into these sizes\n        \"\"\"\n        chrom_to_idx = self.__get_chrom_idx()\n        return [len(indices[0]) for _chrom, indices in chrom_to_idx.items()]\n\n    @property\n    def size_norm_counts(self):\n        \"\"\"Computes and stores table of normalized counts w/ size factor adjustment and no other normalization\"\"\"\n        if not hasattr(self, \"_size_norm_counts\"):\n            self._size_norm_counts = self._set_size_norm_counts()\n        assert self._size_norm_counts.shape == self.data_raw.shape\n        return self._size_norm_counts\n\n    def _set_size_norm_counts(self) -> AnnData:\n        logging.info(f\"Setting size normalized counts\")\n        raw_counts_anndata = AnnData(\n            scipy.sparse.csr_matrix(self.data_raw.raw.X),\n            obs=pd.DataFrame(index=self.data_raw.obs_names),\n            var=pd.DataFrame(index=self.data_raw.var_names),\n        )\n        sc.pp.normalize_total(raw_counts_anndata, inplace=True)\n        # After normalizing, do clustering\n        plot_utils.preprocess_anndata(\n            raw_counts_anndata,\n            louvain_resolution=self.cluster_res,\n            leiden_resolution=self.cluster_res,\n        )\n        return raw_counts_anndata\n\n    @property\n    def size_norm_log_counts(self):\n        \"\"\"Compute and store adata of counts with size factor adjustment and log normalization\"\"\"\n        if not hasattr(self, \"_size_norm_log_counts\"):\n            self._size_norm_log_counts = self._set_size_norm_log_counts()\n        assert self._size_norm_log_counts.shape == self.data_raw.shape\n        return self._size_norm_log_counts\n\n    def _set_size_norm_log_counts(self) -> AnnData:\n        retval = self.size_norm_counts.copy()  # Generates a new copy\n        logging.info(f\"Setting log-normalized size-normalized counts\")\n        # Apply log to it\n        sc.pp.log1p(retval, chunked=True, copy=False, chunk_size=10000)\n        plot_utils.preprocess_anndata(\n            retval,\n            louvain_resolution=self.cluster_res,\n            leiden_resolution=self.cluster_res,\n        )\n        return retval\n\n    @functools.lru_cache(4)\n    def get_cluster_psuedobulk(self, mode=\"leiden\", normalize=True):\n        \"\"\"\n        Return a dictionary mapping each cluster label to the normalized psuedobulk\n        estimate for that cluster\n        If normalize is set to true, then we normalize such that each cluster's row\n        sums to the median count from each cell\n        \"\"\"\n        assert mode in self.data_raw.obs.columns\n        cluster_labels = sorted(list(set(self.data_raw.obs[mode])))\n        norm_counts = self.get_normalized_counts()\n        aggs = []\n        for cluster in cluster_labels:\n            cluster_cells = np.where(self.data_raw.obs[mode] == cluster)\n            pbulk = norm_counts.X[cluster_cells]\n            pbulk_aggregate = np.sum(pbulk, axis=0, keepdims=True)\n            if normalize:\n                pbulk_aggregate = (\n                    pbulk_aggregate\n                    / np.sum(pbulk_aggregate)\n                    * self.data_raw.uns[\"median_counts\"]\n                )\n                assert np.isclose(\n                    np.sum(pbulk_aggregate), self.data_raw.uns[\"median_counts\"]\n                )\n            aggs.append(pbulk_aggregate)\n        retval = AnnData(\n            np.vstack(aggs),\n            obs={mode: cluster_labels},\n            var=self.data_raw.var,\n        )\n        return retval\n\n\n\nclass SingleCellDatasetSplit(Dataset):\n    \"\"\"\n    Wraps SingleCellDataset to provide train/valid/test splits\n    \"\"\"\n\n    def __init__(self, sc_dataset: SingleCellDataset, split: str) -> None:\n        assert isinstance(sc_dataset, SingleCellDataset)\n        self.dset = sc_dataset  # Full dataset\n        self.split = split\n        assert self.split in self.dset.data_split_to_idx\n        logging.info(\n            f\"Created {split} data split with {len(self.dset.data_split_to_idx[self.split])} examples\"\n        )\n\n    def __len__(self) -> int:\n        return len(self.dset.data_split_to_idx[self.split])\n\n    def __getitem__(self, index: int):\n        return self.dset.get_item_data_split(index, self.split)\n\n    # These properties facilitate compatibility with old code by forwarding some properties\n    # Note that these are NOT meant to be modified\n#     @cached_property\n    def size_norm_counts(self) -> AnnData:\n        indices = self.dset.data_split_to_idx[self.split]\n        return self.dset.size_norm_counts[indices].copy()\n\n#     @cached_property\n    def data_raw(self) -> AnnData:\n        indices = self.dset.data_split_to_idx[self.split]\n        return self.dset.data_raw[indices].copy()\n\n#     @cached_property\n    def obs_names(self):\n        indices = self.dset.data_split_to_idx[self.split]\n        return self.dset.data_raw.obs_names[indices]\n","metadata":{"execution":{"iopub.status.busy":"2022-11-09T10:09:47.401502Z","iopub.execute_input":"2022-11-09T10:09:47.402012Z","iopub.status.idle":"2022-11-09T10:09:48.656337Z","shell.execute_reply.started":"2022-11-09T10:09:47.401973Z","shell.execute_reply":"2022-11-09T10:09:48.654821Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pip install scanpy","metadata":{"execution":{"iopub.status.busy":"2022-11-09T10:08:41.801496Z","iopub.execute_input":"2022-11-09T10:08:41.801928Z","iopub.status.idle":"2022-11-09T10:08:56.587825Z","shell.execute_reply.started":"2022-11-09T10:08:41.801879Z","shell.execute_reply":"2022-11-09T10:08:56.586011Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sc_rna_test_dataset = SingleCellDatasetSplit(\n    SingleCellDataset(torch.Tensor(multi_train_y_norm.values)), split=\"test\",\n)","metadata":{"execution":{"iopub.status.busy":"2022-11-09T10:09:52.429922Z","iopub.execute_input":"2022-11-09T10:09:52.430363Z","iopub.status.idle":"2022-11-09T10:09:52.582148Z","shell.execute_reply.started":"2022-11-09T10:09:52.430330Z","shell.execute_reply":"2022-11-09T10:09:52.580051Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\"\"\"\nCustom activation functions\n\"\"\"\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\n\nclass Exp(nn.Module):\n    \"\"\"Applies torch.exp, clamped to improve stability during training\"\"\"\n\n    def __init__(self, minimum=1e-5, maximum=1e6):\n        \"\"\"Values taken from DCA\"\"\"\n        super(Exp, self).__init__()\n        self.min_value = minimum\n        self.max_value = maximum\n\n    def forward(self, input):\n        return torch.clamp(\n            torch.exp(input),\n            min=self.min_value,\n            max=self.max_value,\n        )\n\n\nclass ClippedSoftplus(nn.Module):\n    def __init__(self, beta=1, threshold=20, minimum=1e-4, maximum=1e3):\n        super(ClippedSoftplus, self).__init__()\n        self.beta = beta\n        self.threshold = threshold\n        self.min_value = minimum\n        self.max_value = maximum\n\n    def forward(self, input):\n        return torch.clamp(\n            F.softplus(input, self.beta, self.threshold),\n            min=self.min_value,\n            max=self.max_value,\n        )\n\n    def extra_repr(self):\n        return \"beta={}, threshold={}, min={}, max={}\".format(\n            self.beta,\n            self.threshold,\n            self.min_value,\n            self.max_value,\n        )","metadata":{"execution":{"iopub.status.busy":"2022-11-09T09:27:04.284120Z","iopub.execute_input":"2022-11-09T09:27:04.284581Z","iopub.status.idle":"2022-11-09T09:27:04.298244Z","shell.execute_reply.started":"2022-11-09T09:27:04.284532Z","shell.execute_reply":"2022-11-09T09:27:04.296686Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n        if args.pretrain:\n            # Load in the warm start parameters\n            spliced_net.load_params(f_params=args.pretrain)\n            spliced_net.partial_fit(sc_dual_train_dataset, y=None)\n        else:\n            spliced_net.fit(sc_dual_train_dataset, y=None)\n\n        fig = plot_loss_history(\n            spliced_net.history, os.path.join(outdir_name, f\"loss.{args.ext}\")\n        )\n        plt.close(fig)\n\n        logging.info(\"Evaluating on test set\")\n        logging.info(\"Evaluating RNA > RNA\")\n        sc_rna_test_preds = spliced_net.translate_1_to_1(sc_dual_test_dataset)\n        sc_rna_test_preds_anndata = sc.AnnData(\n            sc_rna_test_preds,\n            var=sc_rna_test_dataset.data_raw.var,\n            obs=sc_rna_test_dataset.data_raw.obs,\n        )\n        sc_rna_test_preds_anndata.write_h5ad(\n            os.path.join(outdir_name, \"rna_rna_test_preds.h5ad\")\n        )","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_one_epoch(epoch_index, tb_writer):\n    running_loss = 0.\n    last_loss = 0.\n\n    # Here, we use enumerate(training_loader) instead of\n    # iter(training_loader) so that we can track the batch\n    # index and do some intra-epoch reporting\n    \n    for i in range(0, multi_train_x.shape[0], 4):\n        # Every data instance is an input + label pair\n        data = (torch.Tensor([multi_train_x[j] for j in range(i, i+4)]), torch.Tensor([multi_train_y_pca[j] for j in range(i, i+4)]))\n        \n        # Zero your gradients for every batch!\n        optimizer.zero_grad()\n\n        # Make predictions for this batch\n        outputs = model(data)\n\n        # Compute the loss and its gradients\n        loss = loss_fn(outputs, labels) # ??? \n        loss.backward()\n\n        # Adjust learning weights\n        optimizer.step()\n\n        # Gather data and report\n        running_loss += loss.item()\n        if i % 1000 == 999:\n            last_loss = running_loss / 1000 # loss per batch\n            print('  batch {} loss: {}'.format(i + 1, last_loss))\n            tb_x = epoch_index * len(training_loader) + i + 1\n            tb_writer.add_scalar('Loss/train', last_loss, tb_x)\n            running_loss = 0.\n\n    return last_loss","metadata":{"execution":{"iopub.status.busy":"2022-11-09T09:01:00.677550Z","iopub.execute_input":"2022-11-09T09:01:00.678134Z","iopub.status.idle":"2022-11-09T09:01:00.690029Z","shell.execute_reply.started":"2022-11-09T09:01:00.678094Z","shell.execute_reply":"2022-11-09T09:01:00.688595Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"itr = iter(training_loader)\nm = next(itr)","metadata":{"execution":{"iopub.status.busy":"2022-11-09T08:33:08.010719Z","iopub.execute_input":"2022-11-09T08:33:08.011185Z","iopub.status.idle":"2022-11-09T08:33:08.140976Z","shell.execute_reply.started":"2022-11-09T08:33:08.011151Z","shell.execute_reply":"2022-11-09T08:33:08.138513Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"multi_train_x[0]","metadata":{"execution":{"iopub.status.busy":"2022-11-09T08:39:25.686453Z","iopub.execute_input":"2022-11-09T08:39:25.687016Z","iopub.status.idle":"2022-11-09T08:39:25.697559Z","shell.execute_reply.started":"2022-11-09T08:39:25.686975Z","shell.execute_reply":"2022-11-09T08:39:25.696135Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"multi_train_x[1]","metadata":{"execution":{"iopub.status.busy":"2022-11-09T08:34:21.792505Z","iopub.execute_input":"2022-11-09T08:34:21.793252Z","iopub.status.idle":"2022-11-09T08:34:21.807000Z","shell.execute_reply.started":"2022-11-09T08:34:21.793189Z","shell.execute_reply":"2022-11-09T08:34:21.806002Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"multi_train_y_pca[2].shape","metadata":{"execution":{"iopub.status.busy":"2022-11-09T08:35:44.567420Z","iopub.execute_input":"2022-11-09T08:35:44.567829Z","iopub.status.idle":"2022-11-09T08:35:44.576554Z","shell.execute_reply.started":"2022-11-09T08:35:44.567798Z","shell.execute_reply":"2022-11-09T08:35:44.575211Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from datetime import datetime\nfrom torch.utils.tensorboard import SummaryWriter","metadata":{"execution":{"iopub.status.busy":"2022-11-09T08:05:09.443994Z","iopub.execute_input":"2022-11-09T08:05:09.444605Z","iopub.status.idle":"2022-11-09T08:05:09.450296Z","shell.execute_reply.started":"2022-11-09T08:05:09.444557Z","shell.execute_reply":"2022-11-09T08:05:09.449195Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"timestamp = datetime.now().strftime('%Y%m%d_%H%M%S')\nwriter = SummaryWriter('runs/fashion_trainer_{}'.format(timestamp))\nepoch_number = 0\n\nEPOCHS = 5\n\nbest_vloss = 1_000_000.\n\nfor epoch in range(EPOCHS):\n    print('EPOCH {}:'.format(epoch_number + 1))\n\n    # Make sure gradient tracking is on, and do a pass over the data\n    model.train(True)\n    avg_loss = train_one_epoch(epoch_number, writer)\n\n    # We don't need gradients on to do reporting\n    model.train(False)\n\n    running_vloss = 0.0\n    for i, vdata in enumerate(validation_loader):\n        vinputs, vlabels = vdata\n        voutputs = model(vinputs)\n        vloss = loss_fn(voutputs, vlabels)\n        running_vloss += vloss\n\n    avg_vloss = running_vloss / (i + 1)\n    print('LOSS train {} valid {}'.format(avg_loss, avg_vloss))\n\n    # Log the running loss averaged per batch\n    # for both training and validation\n    writer.add_scalars('Training vs. Validation Loss',\n                    { 'Training' : avg_loss, 'Validation' : avg_vloss },\n                    epoch_number + 1)\n    writer.flush()\n\n    # Track best performance, and save the model's state\n    if avg_vloss < best_vloss:\n        best_vloss = avg_vloss\n        model_path = 'model_{}_{}'.format(timestamp, epoch_number)\n        torch.save(model.state_dict(), model_path)\n\n    epoch_number += 1","metadata":{"execution":{"iopub.status.busy":"2022-11-09T08:45:04.383425Z","iopub.execute_input":"2022-11-09T08:45:04.383926Z","iopub.status.idle":"2022-11-09T08:45:04.479595Z","shell.execute_reply.started":"2022-11-09T08:45:04.383874Z","shell.execute_reply":"2022-11-09T08:45:04.477846Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"code","source":"print(multi_train_y_pca.shape)\nmulti_train_x.shape","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"t0 = torch.Tensor(multi_train_x)\nt1 = torch.Tensor(multi_train_y_pca)\nprint(t1.shape)\nmodel.eval()\nm = model([t0, t1])","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nclass AutoEncoder(nn.Module):\n    \"\"\"Vanilla autoencoder\"\"\"\n\n    def __init__(\n        self,\n        num_inputs: int,\n        num_units: int = 32,\n        num_outputs: int = None,\n        activation=nn.PReLU,\n        final_activation=None,\n        seed=8947,\n        output_encoded: bool = True,\n    ):\n        super().__init__()\n        torch.manual_seed(seed)\n        self.output_encoded = output_encoded\n        self.num_inputs = num_inputs\n        self.num_outputs = self.num_inputs if num_outputs is None else num_outputs\n        self.num_units = num_units\n\n        self.encoder = Encoder(\n            self.num_inputs, num_units=self.num_units, activation=activation\n        )\n        self.decoder = Decoder(\n            self.num_outputs,\n            num_units=self.num_units,\n            activation=activation,\n            final_activation=final_activation,\n        )\n\n    def forward(self, X):\n        encoded = self.encoder(X)\n        decoded = self.decoder(encoded)[\n            0\n        ]  # Drop the second output that corresponds to addtl parameters\n        if self.output_encoded:\n            return decoded, encoded\n        return decoded\n","metadata":{"execution":{"iopub.status.busy":"2022-11-09T01:14:42.833217Z","iopub.execute_input":"2022-11-09T01:14:42.833570Z","iopub.status.idle":"2022-11-09T01:14:42.851633Z","shell.execute_reply.started":"2022-11-09T01:14:42.833539Z","shell.execute_reply":"2022-11-09T01:14:42.850299Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\n# Read the table of rows and columns required for submission\neval_ids = pd.read_csv(FP_EVALUATION_IDS, index_col='row_id')\n\n# Convert the string columns to more efficient categorical types\n#eval_ids.cell_id = eval_ids.cell_id.apply(lambda s: int(s, base=16))\neval_ids.cell_id = eval_ids.cell_id.astype(pd.CategoricalDtype())\neval_ids.gene_id = eval_ids.gene_id.astype(pd.CategoricalDtype())\ndisplay(eval_ids)\n\n# Create the set of needed cell_ids\ncell_id_set = set(eval_ids.cell_id)\n\n# Convert the string gene_ids to a more efficient categorical dtype\ny_columns = pd.CategoricalIndex(y_columns, dtype=eval_ids.gene_id.dtype, name='gene_id')\n","metadata":{"execution":{"iopub.status.busy":"2022-11-08T03:44:41.974583Z","iopub.execute_input":"2022-11-08T03:44:41.975008Z","iopub.status.idle":"2022-11-08T03:44:43.879359Z","shell.execute_reply.started":"2022-11-08T03:44:41.974975Z","shell.execute_reply":"2022-11-08T03:44:43.878009Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Prepare an empty series which will be filled with predictions\nsubmission = pd.Series(name='target',\n                       index=pd.MultiIndex.from_frame(eval_ids), \n                       dtype=np.float32)\nsubmission","metadata":{"execution":{"iopub.status.busy":"2022-11-04T21:13:40.77053Z","iopub.execute_input":"2022-11-04T21:13:40.770944Z","iopub.status.idle":"2022-11-04T21:13:40.79617Z","shell.execute_reply.started":"2022-11-04T21:13:40.770908Z","shell.execute_reply":"2022-11-04T21:13:40.794173Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"We now compute the predictions in chunks of 5000 rows and match them with the eval_ids table row by row. The matching is very slow, but space-efficient.","metadata":{}},{"cell_type":"code","source":"%%time\n# Process the test data in chunks of 5000 rows\n\nstart = 0\nchunksize = 5000\ntotal_rows = 0\nwhile True:\n    multi_test_x = None # Free the memory if necessary\n    gc.collect()\n    # Read the 5000 rows and select the 30 % subset which is needed for the submission\n    multi_test_x = pd.read_hdf(FP_MULTIOME_TEST_INPUTS, start=start, stop=start+chunksize)\n    rows_read = len(multi_test_x)\n    needed_row_mask = multi_test_x.index.isin(cell_id_set)\n    multi_test_x = multi_test_x.loc[needed_row_mask]\n    \n    # Keep the index (the cell_ids) for later\n    multi_test_index = multi_test_x.index\n    \n    # Predict\n    multi_test_x = multi_test_x.values\n    multi_test_x = preprocessor.transform(multi_test_x)\n    test_pred = model.predict(multi_test_x)\n    \n    # Convert the predictions to a dataframe so that they can be matched with eval_ids\n    test_pred = pd.DataFrame(test_pred,\n                             index=pd.CategoricalIndex(multi_test_index,\n                                                       dtype=eval_ids.cell_id.dtype,\n                                                       name='cell_id'),\n                             columns=y_columns)\n    gc.collect()\n    \n    # Fill the predictions into the submission series row by row\n    for i, (index, row) in enumerate(test_pred.iterrows()):\n        row = row.reindex(eval_ids.gene_id[eval_ids.cell_id == index])\n        submission.loc[index] = row.values\n    print('na:', submission.isna().sum())\n\n    #test_pred_list.append(test_pred)\n    total_rows += len(multi_test_x)\n    print(total_rows)\n    if rows_read < chunksize: break # this was the last chunk\n    start += chunksize\n    \ndel multi_test_x, multi_test_index, needed_row_mask\n","metadata":{"execution":{"iopub.status.busy":"2022-11-04T20:08:41.137754Z","iopub.execute_input":"2022-11-04T20:08:41.139236Z","iopub.status.idle":"2022-11-04T20:33:55.567671Z","shell.execute_reply.started":"2022-11-04T20:08:41.139196Z","shell.execute_reply":"2022-11-04T20:33:55.566028Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Submission\n\nAs we don't yet have the CITEseq predictions, we save the partial predictions so that they can be used in the [CITEseq notebook](https://www.kaggle.com/ambrosm/msci-citeseq-quickstart).","metadata":{}},{"cell_type":"code","source":"submission.reset_index(drop=True, inplace=True)\nsubmission.index.name = 'row_id'\nwith open(\"partial_submission_multi.pickle\", 'wb') as f: pickle.dump(submission, f)\nsubmission","metadata":{"execution":{"iopub.status.busy":"2022-11-04T20:33:55.570345Z","iopub.execute_input":"2022-11-04T20:33:55.571707Z","iopub.status.idle":"2022-11-04T20:33:56.552592Z","shell.execute_reply.started":"2022-11-04T20:33:55.57165Z","shell.execute_reply":"2022-11-04T20:33:56.551457Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}