{"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":"<div class=\"alert alert-block alert-success\" style=\"font-size:30px\">\n🔬 PyTorch Swiss Army Knife for MSCI Competition 🔬\n</div>\n\n<div class=\"alert alert-block alert-danger\" style=\"text-align:center; font-size:20px;\">\n    ❤️ Dont forget to ▲upvote▲ if you find this notebook usefull!  ❤️\n</div>\n\nThis notebook includes a set of tools to build your deep learning solution for MSCI-2022 competition.\nThis is your one-stop shop to build a winning deep learning model.\nHope you'll find it useful!\n\nHere is the list of features implemented in this notebook:\n\n### Training both Multiome and CITEseq Regressors\nThere is no much difference between Multiome and CITEseq problems apart from the scale. With minibatch training both problmes can be soled within a single framework.\n* set `CFG=CFG_MULTIOME_SVD` or `CFG=CFG_MULTIOME_SPARSE` to train your Multiome regressor\n* set `CFG=CFG_CITESEQ_SVD` or `CFG=CFG_CITESEQ_SPARSE` to train your CITEseq regressor\n\n\n### Both SVD-compressed and Raw Features are Supported\n* SVD-compressed data is loaded from [this notebook](https://www.kaggle.com/code/vslaykovsky/multiome-citeseq-svd-transforms), where TruncatedSVD is used to project raw features to 512 dimensional space. SVD features are concatenated with cell type features in `MSCIDatasetSVD` class\n* Raw data is loaded to memory as sparse matrices and is lazily uncomressed and concatenated with cell_id features in the `MSCIDatasetSparse` class.\n\n### Both Kaggle and Custom Training is Supported\nThis notebook can be easily customised for local training or training on more powerful machines like Colab/google cloud/AWS.\njust fill out your constants in `# local run` section to get started.\n\n### K-fold Training\nTraining on multiple folds enables cross-validation score.\nEnable `TRAIN=True` for training\n\n### Accurate Correlation Metric/Loss Function\n`CorrError` metric is implemented to match competition requirements. This can be used for both training and evaluation of your solution.\n\n### Optuna Hyperparameter Optimization\nOptimize your hyperparameters with Optuna. Multiple hyperparameters are already supported by `MSCIModel`. You are free to add more parameters of course!\nEnable Optuna with `OPTUNA=True`\n\n### Ensembling of k-fold models\nSubmission is generated from K models that come from k-fold training. Implementation is memory-efficient. Only a single copy of predictions (the largest matrix) is loaded in memory at any time.\n\n### Wandb Logging\nTrain with Wandb and track your scores even in background execution! Here is the list of implemented metrics:\n* `eval_score` - correlation score on evaluation set\n* `eval_mse` - MSE score on evaluation set\n* `lr` - learning rate.\n* `epoch` - epoch\n* `train_score` - correlation score on training set\n* `train_loss` - training loss (most of the time == `train_score`)\n* `train_epoch_mse` - epoch-average MSE on training set\n\n### Patching you Predictions Into the Best Public Solution\nYou can always patch your model's outputs into the best public solution if you only generate predictions for a single technology (CITEseq or Multiome).","metadata":{"pycharm":{"name":"#%% md\n"}}},{"cell_type":"markdown","source":"## Model diagram\n\nThis is a simplified diagram of the model generated by pytorchviz. \nYou can see sequences of Linear->LayerNorm->SiLU(ReLU) layers here\n\n<img src=\"https://images2.imgbox.com/be/27/9vy3PmRH_o.png\" alt=\"image host\"/>","metadata":{"execution":{"iopub.status.busy":"2022-09-06T19:01:35.033156Z","iopub.execute_input":"2022-09-06T19:01:35.033798Z","iopub.status.idle":"2022-09-06T19:01:35.040995Z","shell.execute_reply.started":"2022-09-06T19:01:35.033750Z","shell.execute_reply":"2022-09-06T19:01:35.039509Z"}}},{"cell_type":"markdown","source":"# Configuration\n\nSet `CFG` to the right configuration to train a model","metadata":{"pycharm":{"name":"#%% md\n"}}},{"cell_type":"code","source":"!pip install -q torchviz","metadata":{"execution":{"iopub.status.busy":"2022-09-07T07:32:37.377034Z","iopub.execute_input":"2022-09-07T07:32:37.377546Z","iopub.status.idle":"2022-09-07T07:32:48.824005Z","shell.execute_reply.started":"2022-09-07T07:32:37.377439Z","shell.execute_reply":"2022-09-07T07:32:48.822695Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import gc\nimport os.path\n\nimport numpy as np\nimport optuna\nimport pandas as pd\nimport torch\nimport wandb\nfrom optuna.study import StudyDirection\nfrom scipy import sparse\nfrom tqdm.notebook import tqdm\nimport copy\nfrom torchviz import make_dot\n\n\nDEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\n\ntry:\n    import kaggle_secrets\n    IS_KAGGLE = True\nexcept:\n    IS_KAGGLE = False\n\nif IS_KAGGLE:\n    # kaggle run\n    MSCI_ROOT = '../input/open-problems-multimodal'\n    SPARSE_ROOT = '../input/multimodal-single-cell-as-sparse-matrix'\n    SVD_ROOT = '../input/multiome-citeseq-sv-transforms-ds'\n    MODELS_ROOT = '../input/multiomekfoldmodels'\n    PREDICTIONS_ROOT = '.'\n    SUBMISSIONS_ROOT = '../input/lb-0-858-normalized-ensembles-for-pearson-s-r/'\nelse:\n    # local run\n    MSCI_ROOT = '/mnt/msci'\n    SPARSE_ROOT = f'data/sparse'\n    SVD_ROOT = f'data/svd'\n    MODELS_ROOT = 'models'\n    PREDICTIONS_ROOT = '.'\n    SUBMISSIONS_ROOT = 'data/sub'\n\n\nMETA_FILE = f'{SPARSE_ROOT}/metadata.parquet'\n\nN_FOLDS = 5\n\n# Enable/disable parts of the notebook\nTRAIN = False  # training\nCROSS_VALIDATE = False\nOPTUNA = False # hyperparameters search with Optuna\nOPTUNA_N_TRIALS = 30\nPREDICT = True\nSUBMISSION = False\n\nSUBMISSION_FOR_PATCHING = f'{SUBMISSIONS_ROOT}/submission.csv'\nPATCH_CITESEQ = False\nPATCH_MULTI = True\n","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2022-09-07T07:32:48.827880Z","iopub.execute_input":"2022-09-07T07:32:48.828241Z","iopub.status.idle":"2022-09-07T07:32:53.068123Z","shell.execute_reply.started":"2022-09-07T07:32:48.828205Z","shell.execute_reply":"2022-09-07T07:32:53.066840Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"CFG_MULTIOME_SVD = {\n    'TECHNOLOGY': 'Multiome',\n    'MSE_LOSS': False,\n    'SCHEDULER': 'onecycle',\n    'SKIP_CONNECTION': False,\n\n    'TRAIN_INPUTS_VALUES_NPZ': f'{SVD_ROOT}/train_multi_inputs.npz',\n\n    'TRAIN_TARGETS_VALUES_NPZ': f'{SPARSE_ROOT}/train_multi_targets_values.sparse.npz',\n    'TRAIN_TARGETS_IDXCOL_NPZ': f'{SPARSE_ROOT}/train_multi_targets_idxcol.npz',\n    'TRAIN_INPUTS_IDXCOL_NPZ': f'{SPARSE_ROOT}/train_multi_inputs_idxcol.npz',\n\n    'TEST_INPUTS_VALUES_NPZ': f'{SVD_ROOT}/test_multi_inputs.npz',\n    'TEST_INPUTS_IDXCOL_NPZ': f'{SPARSE_ROOT}/test_multi_inputs_idxcol.npz',\n\n    'WANDB_PROJECT': 'MSCI-MULTI-SVD',\n    'NUM_WORKERS': 1,\n    'BATCH_SIZE': 1024,\n    'EPOCHS': 20,\n    'MAX_LR': 0.001,\n    'N_LAYERS': 5,\n    'DROPOUT': False,\n    'HIDDEN_SIZE': 1024,\n    'ADAMW': False,\n    'WEIGHT_DECAY': 0.05,\n    'ACTIVATION': torch.nn.SiLU,\n    'N_FEATURES': 520,\n    'N_TARGETS': 23418,\n}\n\nCFG_MULTIOME_SPARSE = {\n    'TECHNOLOGY': 'Multiome',\n    'MSE_LOSS': False,\n    'SCHEDULER': 'onecycle',\n    'SKIP_CONNECTION': False,\n\n    'TRAIN_INPUTS_VALUES_NPZ': f'{SPARSE_ROOT}/train_multi_inputs_values.sparse.npz',\n    'TRAIN_TARGETS_VALUES_NPZ': f'{SPARSE_ROOT}/train_multi_targets_values.sparse.npz',\n    'TRAIN_TARGETS_IDXCOL_NPZ': f'{SPARSE_ROOT}/train_multi_targets_idxcol.npz',\n    'TRAIN_INPUTS_IDXCOL_NPZ': f'{SPARSE_ROOT}/train_multi_inputs_idxcol.npz',\n\n    'TEST_INPUTS_VALUES_NPZ': f'{SPARSE_ROOT}/test_multi_inputs_values.sparse.npz',\n    'TEST_INPUTS_IDXCOL_NPZ': f'{SPARSE_ROOT}/test_multi_inputs_idxcol.npz',\n\n    'WANDB_PROJECT': 'MSCI-MULTI-Sparse',\n    'NUM_WORKERS': 1,\n    'BATCH_SIZE': 64,\n    'EPOCHS': 7,\n    'WEIGHT_DECAY': 0.05,\n    'MAX_LR': 0.0001,\n    'ADAMW': True,\n    'N_LAYERS': 7,\n    'DROPOUT': False,\n    'HIDDEN_SIZE': 1024,\n    'ACTIVATION': torch.nn.SiLU,\n    'N_FEATURES': 228950,\n    'N_TARGETS': 23418,\n}\n\nCFG_CITESEQ_SVD = {\n    'TECHNOLOGY': 'CITEseq',\n    'MSE_LOSS': False,\n    'SCHEDULER': 'onecycle',\n\n    'TRAIN_INPUTS_VALUES_NPZ': f'{SVD_ROOT}/train_cite_inputs.npz',\n    \n    'TRAIN_TARGETS_VALUES_NPZ': f'{SPARSE_ROOT}/train_cite_targets_values.sparse.npz',\n    'TRAIN_TARGETS_IDXCOL_NPZ': f'{SPARSE_ROOT}/train_cite_targets_idxcol.npz',\n    'TRAIN_INPUTS_IDXCOL_NPZ': f'{SPARSE_ROOT}/train_cite_inputs_idxcol.npz',\n    \n    'TEST_INPUTS_VALUES_NPZ': f'{SVD_ROOT}/test_cite_inputs.npz',\n    'TEST_INPUTS_IDXCOL_NPZ': f'{SPARSE_ROOT}/test_cite_inputs_idxcol.npz',\n    \n    'SKIP_CONNECTION': False,\n\n    'WANDB_PROJECT': 'MSCI-CITE-SVD',\n    'NUM_WORKERS': 2,\n    'BATCH_SIZE': 1024,\n    'EPOCHS': 10,\n    'ADAMW': False,\n    'WEIGHT_DECAY': 0.05,\n    'MAX_LR': 0.003,\n    'N_LAYERS': 7,\n    'DROPOUT': False,\n    'HIDDEN_SIZE': 1024,\n    'ACTIVATION': torch.nn.SiLU,\n    'N_FEATURES': 520,\n    'N_TARGETS': 140,\n}\n\nCFG_CITESEQ_SPARSE = {\n    'TECHNOLOGY': 'CITEseq',\n    'MSE_LOSS': False,\n    'SCHEDULER': 'onecycle',\n\n    'TRAIN_INPUTS_VALUES_NPZ': f'{SPARSE_ROOT}/train_cite_inputs_values.sparse.npz',\n    \n    'TRAIN_TARGETS_VALUES_NPZ': f'{SPARSE_ROOT}/train_cite_targets_values.sparse.npz',\n    'TRAIN_TARGETS_IDXCOL_NPZ': f'{SPARSE_ROOT}/train_cite_targets_idxcol.npz',\n    'TRAIN_INPUTS_IDXCOL_NPZ': f'{SPARSE_ROOT}/train_cite_inputs_idxcol.npz',\n    \n    'TEST_INPUTS_VALUES_NPZ': f'{SPARSE_ROOT}/test_cite_inputs_values.sparse.npz',\n    'TEST_INPUTS_IDXCOL_NPZ': f'{SPARSE_ROOT}/test_cite_inputs_idxcol.npz',\n\n    'SKIP_CONNECTION': False,\n\n    'WANDB_PROJECT': 'MSCI-CITE-Sparse',\n    'NUM_WORKERS': 2,\n    'BATCH_SIZE': 256,\n    'EPOCHS': 10,\n    'WEIGHT_DECAY': 0.05,\n    'MAX_LR': 0.0001,\n    'ADAMW': True,\n    'N_LAYERS': 7,\n    'DROPOUT': False,\n    'HIDDEN_SIZE': 1024,\n    'ACTIVATION': torch.nn.SiLU,\n    'N_FEATURES': 22058,\n    'N_TARGETS': 140,\n}\n\n\n# Choose any of the above configurations to train your model!\nCFG = CFG_MULTIOME_SVD","metadata":{"collapsed":false,"pycharm":{"name":"#%%\n"},"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-07T07:32:53.071242Z","iopub.execute_input":"2022-09-07T07:32:53.073110Z","iopub.status.idle":"2022-09-07T07:32:53.087166Z","shell.execute_reply.started":"2022-09-07T07:32:53.073064Z","shell.execute_reply":"2022-09-07T07:32:53.086054Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Datasets\n\nBoth SVD and raw sparse featurers are supported in MSCIDatasetSVD and MSCIDatasetSparse","metadata":{"pycharm":{"name":"#%% md\n"}}},{"cell_type":"code","source":"def load_meta():\n    df_meta = pd.read_parquet(META_FILE).set_index('cell_id')\n    df_meta = pd.get_dummies(df_meta['cell_type'])\n    return df_meta\n\n# quick test\nload_meta().shape","metadata":{"collapsed":false,"pycharm":{"name":"#%%\n"},"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-07T07:32:53.088694Z","iopub.execute_input":"2022-09-07T07:32:53.089602Z","iopub.status.idle":"2022-09-07T07:32:53.568152Z","shell.execute_reply.started":"2022-09-07T07:32:53.089562Z","shell.execute_reply":"2022-09-07T07:32:53.566979Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class MSCIDatasetSVD(torch.utils.data.Dataset):\n\n    def __init__(self, df_meta, input_index, input_svd, targets=None):\n        cell_type = df_meta.loc[input_index].values\n        self.data = np.concatenate([cell_type, input_svd[:len(cell_type)]], axis=1)\n        self.targets = targets\n\n    def __getitem__(self, item):\n        if self.targets is not None:\n            return self.data[item], np.asarray(self.targets[item].todense())[0]\n        else:\n            return self.data[item]\n\n    def __len__(self):\n        return len(self.data)\n\n\nclass MSCIDatasetSparse(torch.utils.data.Dataset):\n\n    def __init__(self, df_meta, input_index, input_sparse, targets=None):\n        self.cell_type = df_meta.loc[input_index].values.astype('float32')\n        self.data = input_sparse\n        self.targets = targets\n\n    def __getitem__(self, item):\n        if self.targets is not None:\n            return np.concatenate([\n                self.cell_type[item],\n                np.asarray(self.data[item].todense())[0]\n            ]), np.asarray(self.targets[item].todense())[0]\n        else:\n            return np.concatenate([\n                self.cell_type[item],\n                np.asarray(self.data[item].todense())[0]\n            ])\n\n    def __len__(self):\n        return len(self.cell_type)\n\n\ndef load_dataset():\n    df_meta = load_meta()\n    if 'sparse.npz' in CFG['TRAIN_INPUTS_VALUES_NPZ']:\n        ds_data = MSCIDatasetSparse(\n            df_meta,\n            np.load(CFG['TRAIN_INPUTS_IDXCOL_NPZ'], allow_pickle=True)['index'],\n            sparse.load_npz(CFG['TRAIN_INPUTS_VALUES_NPZ']),\n            sparse.load_npz(CFG['TRAIN_TARGETS_VALUES_NPZ'])\n        )\n    else:\n        ds_data = MSCIDatasetSVD(\n            df_meta,\n            np.load(CFG['TRAIN_INPUTS_IDXCOL_NPZ'], allow_pickle=True)['index'],\n            np.load(CFG['TRAIN_INPUTS_VALUES_NPZ'])['values'],\n            sparse.load_npz(CFG['TRAIN_TARGETS_VALUES_NPZ'])\n        )\n    return ds_data","metadata":{"collapsed":false,"pycharm":{"name":"#%%\n"},"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-07T07:32:53.573220Z","iopub.execute_input":"2022-09-07T07:32:53.574038Z","iopub.status.idle":"2022-09-07T07:32:53.587475Z","shell.execute_reply.started":"2022-09-07T07:32:53.573980Z","shell.execute_reply":"2022-09-07T07:32:53.586220Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model\n\nA set of dense layers with LayerNorm. LayerNorm and SiLU result in better convergence","metadata":{"pycharm":{"name":"#%% md\n"}}},{"cell_type":"code","source":"!apt install graphviz","metadata":{"execution":{"iopub.status.busy":"2022-09-07T07:32:53.589683Z","iopub.execute_input":"2022-09-07T07:32:53.590188Z","iopub.status.idle":"2022-09-07T07:32:56.421311Z","shell.execute_reply.started":"2022-09-07T07:32:53.590064Z","shell.execute_reply":"2022-09-07T07:32:56.419540Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class MSCIModel(torch.nn.Module):\n\n    def __init__(self, cfg):\n        super().__init__()\n        input_size, hidden_size, n_layers, output_size, activation, dropout, skip_connection = cfg['N_FEATURES'], cfg['HIDDEN_SIZE'], cfg['N_LAYERS'], cfg['N_TARGETS'], cfg['ACTIVATION'], cfg['DROPOUT'], cfg['SKIP_CONNECTION']\n\n        self.skip_connection = skip_connection\n        self.encoder = torch.nn.Sequential(\n            torch.nn.Linear(input_size, hidden_size),\n            torch.nn.LayerNorm(hidden_size),\n            activation(),\n        )\n        self.blocks = torch.nn.ModuleList([\n            torch.nn.Sequential(\n                torch.nn.Linear(hidden_size, hidden_size),\n                torch.nn.LayerNorm(hidden_size),\n                activation(),\n            )\n            for _ in range(n_layers)]\n        )\n\n        self.output = torch.nn.Sequential(\n            *(\n                    [torch.nn.Dropout(0.1)] if dropout else [] +\n                    [\n                        torch.nn.Linear(hidden_size, output_size),\n                        torch.nn.LayerNorm(output_size),\n                        torch.nn.ReLU(),\n                    ]\n            )\n        )\n\n    def forward(self, x):\n        x = self.encoder(x)\n        for block in self.blocks:\n            if self.skip_connection:\n                x = block(x) + x\n            else:\n                x = block(x)\n        x = self.output(x)\n        return x\n\n\n# quick test\n\ncfg = copy.copy(CFG)\ncfg['N_LAYERS'] = 1\nm = MSCIModel(cfg)\nprint(m)\nprint(m.forward(torch.randn(2, cfg['N_FEATURES'])))\n\n\nx = torch.randn(2, cfg['N_FEATURES']).requires_grad_(True)\ny = m(x)   \nmake_dot(y, params=dict(list(m.named_parameters()) + [('x', x)]))\n","metadata":{"collapsed":false,"pycharm":{"name":"#%%\n"},"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-07T07:32:56.424917Z","iopub.execute_input":"2022-09-07T07:32:56.425942Z","iopub.status.idle":"2022-09-07T07:32:56.966790Z","shell.execute_reply.started":"2022-09-07T07:32:56.425892Z","shell.execute_reply":"2022-09-07T07:32:56.965497Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training\n\nDifferentiable correlation error function.\nWe test that it's correct by comparing to `torch.corrcoef`","metadata":{"pycharm":{"name":"#%% md\n"}}},{"cell_type":"code","source":"class CorrError():\n    def __init__(self, reduction='mean', normalize=True):\n        self.reduction, self.normalize = reduction, normalize\n\n    def __call__(self, y, y_target):\n        y = y - torch.mean(y, dim=1).unsqueeze(1)\n        y_target = y_target - torch.mean(y_target, dim=1).unsqueeze(1)\n        loss = -torch.sum(y * y_target, dim=1) / (y_target.shape[-1] - 1)  # minus because we want gradient ascend\n        if self.normalize:\n            s1 = torch.sqrt(torch.sum(y * y, dim=1) / (y.shape[-1] - 1))\n            s2 = torch.sqrt(torch.sum(y_target * y_target, dim=1) / (y_target.shape[-1] - 1))\n            loss = loss / s1 / s2\n        if self.reduction == 'mean':\n            return torch.mean(loss)\n        return loss\n\n\n# quick test\na = torch.tensor([[0, 1, 1., 0.1, 0.3, 0.4]])\nb = torch.tensor([[0, 0, 1., 10, 10, -13]])\n\ncorr1 = CorrError()(a, b).item()\ncorr2 = torch.corrcoef(torch.stack([a[0], b[0]]))[0, 1].item()\nassert abs(-corr1 - corr2) < 1e-5\n\nCorrError(reduction='none')(torch.randn(2, 3), torch.randn(2, 3))","metadata":{"collapsed":false,"pycharm":{"name":"#%%\n"},"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-07T07:32:56.968762Z","iopub.execute_input":"2022-09-07T07:32:56.969572Z","iopub.status.idle":"2022-09-07T07:32:56.999814Z","shell.execute_reply.started":"2022-09-07T07:32:56.969524Z","shell.execute_reply":"2022-09-07T07:32:56.998519Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_parameter_names(model, forbidden_layer_types):\n    \"\"\"\n    Returns the names of the model parameters that are not inside a forbidden layer.\n    \"\"\"\n    result = []\n    for name, child in model.named_children():\n        result += [\n            f\"{name}.{n}\"\n            for n in get_parameter_names(child, forbidden_layer_types)\n            if not isinstance(child, tuple(forbidden_layer_types))\n        ]\n    # Add model specific parameters (defined with nn.Parameter) since they are not in any child.\n    result += list(model._parameters.keys())\n    return result\n\ndef adamw_optimizer(model, weight_decay):\n    decay_parameters = get_parameter_names(model, [torch.nn.LayerNorm])\n    decay_parameters = [name for name in decay_parameters if \"bias\" not in name]\n    optimizer_grouped_parameters = [\n        {\n            \"params\": [p for n, p in model.named_parameters() if n in decay_parameters],\n            \"weight_decay\": weight_decay,\n        },\n        {\n            \"params\": [p for n, p in model.named_parameters() if n not in decay_parameters],\n            \"weight_decay\": 0.0,\n        },\n    ]\n    return torch.optim.AdamW(optimizer_grouped_parameters)\n\n\n# quick test\nadamw_optimizer(torch.nn.Linear(3, 4), 0.001).step()","metadata":{"collapsed":false,"pycharm":{"name":"#%%\n"},"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-07T07:32:57.001334Z","iopub.execute_input":"2022-09-07T07:32:57.001953Z","iopub.status.idle":"2022-09-07T07:32:57.013513Z","shell.execute_reply.started":"2022-09-07T07:32:57.001912Z","shell.execute_reply":"2022-09-07T07:32:57.012491Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def evaluate(model, ds_eval, cfg, batch_progress):\n    dl_eval = torch.utils.data.DataLoader(ds_eval, batch_size=cfg['BATCH_SIZE'], num_workers=cfg['NUM_WORKERS'])\n\n    with torch.no_grad():\n        model.eval()\n        with tqdm(dl_eval, miniters=100, desc='Batch', disable=not batch_progress) as progress:\n            scores = []\n            mses = []\n            for batch_idx, (X, y) in enumerate(progress):\n                y_pred = model.forward(X.to(DEVICE))\n                score = CorrError()(y_pred.detach(), y.to(DEVICE)).item()\n                progress.set_description(f'Eval Loss: {score:02f}', refresh=False)\n                scores.append(score)\n                mses.append(torch.nn.MSELoss()(y_pred.detach(), y.to(DEVICE)).item())\n\n            score = np.mean(scores)\n            mses = np.mean(mses)\n            model.train()\n            return score, mses\n","metadata":{"collapsed":false,"pycharm":{"name":"#%%\n"},"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-07T07:32:57.015200Z","iopub.execute_input":"2022-09-07T07:32:57.016523Z","iopub.status.idle":"2022-09-07T07:32:57.026405Z","shell.execute_reply.started":"2022-09-07T07:32:57.016474Z","shell.execute_reply":"2022-09-07T07:32:57.025374Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train(cfg, ds_train, ds_eval, wandb_run=None, batch_progress=True, store_best=True, name=''):\n    dl_train = torch.utils.data.DataLoader(ds_train, batch_size=cfg['BATCH_SIZE'], num_workers=cfg['NUM_WORKERS'])\n\n    model = MSCIModel(cfg)\n    model.to(DEVICE)\n    model.train()\n    best_score = 1.\n\n    if cfg['ADAMW']:\n        optim = adamw_optimizer(model, cfg['WEIGHT_DECAY'])\n    else:\n        optim = torch.optim.Adam(model.parameters(), lr=(cfg['MAX_LR']))\n\n    if cfg['MSE_LOSS']:\n        criterion = torch.nn.MSELoss()\n    else:\n        criterion = CorrError()\n\n    if cfg['SCHEDULER'] == 'onecycle':\n        scheduler = torch.optim.lr_scheduler.OneCycleLR(optim, max_lr=(cfg['MAX_LR']), epochs=cfg['EPOCHS'],\n                                                        steps_per_epoch=len(dl_train))\n    else:\n        scheduler = torch.optim.lr_scheduler.ExponentialLR(optim, gamma=0.3)\n\n    with tqdm(range(cfg['EPOCHS']), desc='Epoch') as epoch_progress:\n        for epoch in epoch_progress:\n            # ************** train cycle **************\n            mses = []\n            scores = []\n            with tqdm(dl_train, miniters=100, desc='Batch', disable=not batch_progress) as progress:\n                for batch_idx, (X, y) in enumerate(progress):\n                    y_pred = model.forward(X.to(DEVICE))\n                    loss = criterion(y_pred, y.to(DEVICE))\n                    optim.zero_grad()\n                    loss.backward()\n                    optim.step()\n                    if cfg['SCHEDULER'] == 'onecycle':\n                        scheduler.step()\n\n                    score = CorrError()(y_pred.detach(), y.to(DEVICE)).item()\n                    scores.append(score)\n                    mses.append(torch.nn.MSELoss()(y_pred.detach(), y.to(DEVICE)).item())\n                    progress.set_description(f'Loss: {score:02f}', refresh=False)\n                    if wandb_run is not None:\n                        wandb_run.log({'lr': float(scheduler.get_last_lr()[0]),\n                                       'train_score': score,\n                                       'train_loss': loss.item(),\n                                       'epoch': epoch})\n            if wandb_run is not None:\n                wandb_run.log({'train_epoch_score': np.mean(scores), 'train_epoch_mse': np.mean(mses)})\n            if cfg['SCHEDULER'] != 'onecycle':\n                scheduler.step()\n\n            # ************** eval cycle **************\n            score, mses = evaluate(model, ds_eval, cfg, batch_progress)\n            if wandb_run is not None:\n                wandb_run.log({'eval_score': score, 'eval_mse': mses})\n            if score < best_score:\n                best_score = score\n                if store_best:\n                    !mkdir -p {MODELS_ROOT}\n                    torch.save(model.state_dict(), f'{cfg[\"WANDB_PROJECT\"]}-{name}.pth')\n            epoch_progress.set_description(f'Epochs, eval loss:{score:.03f}')\n    return best_score","metadata":{"collapsed":false,"pycharm":{"name":"#%%\n"},"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-07T07:32:57.027905Z","iopub.execute_input":"2022-09-07T07:32:57.028927Z","iopub.status.idle":"2022-09-07T07:32:57.053812Z","shell.execute_reply.started":"2022-09-07T07:32:57.028888Z","shell.execute_reply":"2022-09-07T07:32:57.052598Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def kfold_split(ds):\n    fold_sizes = [len(ds) // N_FOLDS] * (N_FOLDS - 1) + [len(ds) // N_FOLDS + len(ds) % N_FOLDS]\n    ds_folds = torch.utils.data.random_split(ds, fold_sizes, generator=torch.Generator().manual_seed(42))\n    for fold in range(N_FOLDS):\n        yield torch.utils.data.ConcatDataset(ds_folds[:fold] + ds_folds[fold + 1:]), ds_folds[fold]\n\n# quick test\nfor ds_train, ds_test in kfold_split(list(range(10))):\n    print([i for i in ds_train], [i for i in ds_test])","metadata":{"collapsed":false,"pycharm":{"name":"#%%\n"},"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-07T07:32:57.056811Z","iopub.execute_input":"2022-09-07T07:32:57.057827Z","iopub.status.idle":"2022-09-07T07:32:57.070443Z","shell.execute_reply.started":"2022-09-07T07:32:57.057788Z","shell.execute_reply":"2022-09-07T07:32:57.069476Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if TRAIN:\n    ds_data = load_dataset()\n    for fold, (ds_train, ds_eval) in enumerate(tqdm(kfold_split(ds_data), desc='Train fold', total=N_FOLDS)):\n        with wandb.init(project=CFG['WANDB_PROJECT'], name=f'pytorch-{fold}') as run:\n            train(CFG, ds_train, ds_eval, wandb_run=run, batch_progress=False, name=str(fold))","metadata":{"collapsed":false,"pycharm":{"name":"#%%\n"},"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-07T07:32:57.072306Z","iopub.execute_input":"2022-09-07T07:32:57.072921Z","iopub.status.idle":"2022-09-07T07:32:57.083099Z","shell.execute_reply.started":"2022-09-07T07:32:57.072881Z","shell.execute_reply":"2022-09-07T07:32:57.082142Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Wandb output for Multiome SVD training:\n\n<img src=\"https://images2.imgbox.com/d2/98/Dm7OTXD3_o.png\" alt=\"image host\"/>","metadata":{"pycharm":{"name":"#%% md\n"}}},{"cell_type":"markdown","source":"# Cross-validation","metadata":{"pycharm":{"name":"#%% md\n"}}},{"cell_type":"code","source":"def load_model(fname):\n    model = MSCIModel(CFG)\n    model.load_state_dict(torch.load(fname))\n    return model","metadata":{"collapsed":false,"pycharm":{"name":"#%%\n"},"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-07T07:32:57.087851Z","iopub.execute_input":"2022-09-07T07:32:57.088891Z","iopub.status.idle":"2022-09-07T07:32:57.094354Z","shell.execute_reply.started":"2022-09-07T07:32:57.088852Z","shell.execute_reply":"2022-09-07T07:32:57.093373Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CROSS_VALIDATE:\n    ds = load_dataset()\n    scores = []\n    for fold, (_, ds_eval) in enumerate(tqdm(kfold_split(ds), desc='Evaluating Folds', total=N_FOLDS)):\n        model = load_model(f'{MODELS_ROOT}/{CFG[\"WANDB_PROJECT\"]}-{fold}.pth').to(DEVICE)\n        scores, mses = evaluate(model, ds_eval, CFG, batch_progress=False)\n        del model, ds_eval\n        gc.collect()\n    print('CV score:', -np.mean(scores))","metadata":{"collapsed":false,"pycharm":{"name":"#%%\n"},"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-07T07:32:57.096006Z","iopub.execute_input":"2022-09-07T07:32:57.096676Z","iopub.status.idle":"2022-09-07T07:32:57.104803Z","shell.execute_reply.started":"2022-09-07T07:32:57.096636Z","shell.execute_reply":"2022-09-07T07:32:57.103772Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Optuna Hyperparameters Tuning","metadata":{"pycharm":{"name":"#%% md\n"}}},{"cell_type":"code","source":"import copy\n\n\ndef objective(trial: optuna.trial.Trial, ds_data):\n    is_adamw = trial.suggest_categorical('adamw', [True, False])\n    weight_decay = 0.\n    if is_adamw:\n        weight_decay = trial.suggest_float('weight_decay', 0.0001, 0.1, log=True)\n    max_lr = trial.suggest_float('max_lr', 0.0001, 0.01, log=True)\n    n_layers = trial.suggest_int('n_layers', 2, 8, log=True)\n    hidden_size = trial.suggest_int('hidden_size', 128, 4096, log=True)\n    activation = eval(trial.suggest_categorical('activation', ['torch.nn.SiLU', 'torch.nn.GELU', 'torch.nn.ReLU']))\n    dropout = trial.suggest_categorical('dropout', [True, False])\n    skip_connection = trial.suggest_categorical('skip', [True, False])\n\n    ds_train, ds_test = next(iter(kfold_split(ds_data)))\n    cfg = copy.copy(CFG)\n    cfg.update({\n        'ADAMW': is_adamw,\n        'WEIGHT_DECAY': weight_decay,\n        'MAX_LR': max_lr,\n        'N_LAYERS': n_layers,\n        'HIDDEN_SIZE': hidden_size,\n        'ACTIVATION': activation,\n        'DROPOUT': dropout,\n        'SKIP_CONNECTION': skip_connection,\n        'EPOCHS': 5,\n    })\n    return train(cfg, ds_train, ds_test, wandb_run=None, batch_progress=False, store_best=False)\n\ndef run_optuna():\n    ds_data = load_dataset()\n    study = optuna.create_study(direction=StudyDirection.MINIMIZE)\n    study.optimize(lambda trial: objective(trial, ds_data), n_trials=OPTUNA_N_TRIALS)\n    print(study.best_params)\n\nif OPTUNA:\n    run_optuna()","metadata":{"collapsed":false,"pycharm":{"name":"#%%\n"},"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-07T07:32:57.106319Z","iopub.execute_input":"2022-09-07T07:32:57.106979Z","iopub.status.idle":"2022-09-07T07:32:57.119443Z","shell.execute_reply.started":"2022-09-07T07:32:57.106940Z","shell.execute_reply":"2022-09-07T07:32:57.118346Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Prediction","metadata":{"pycharm":{"name":"#%% md\n"}}},{"cell_type":"code","source":"def load_ds_test():\n    df_meta = load_meta()\n    if 'sparse.npz' in CFG['TEST_INPUTS_VALUES_NPZ']:\n        ds_test = MSCIDatasetSparse(\n            df_meta,\n            np.load(CFG['TEST_INPUTS_IDXCOL_NPZ'], allow_pickle=True)['index'],\n            sparse.load_npz(CFG['TEST_INPUTS_VALUES_NPZ']),\n        )\n    else:        \n        ds_test = MSCIDatasetSVD(\n            df_meta,\n            np.load(CFG['TEST_INPUTS_IDXCOL_NPZ'], allow_pickle=True)['index'],\n            np.load(CFG['TEST_INPUTS_VALUES_NPZ'])['values'],\n        )\n    return ds_test","metadata":{"collapsed":false,"pycharm":{"name":"#%%\n"},"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-07T07:32:57.122252Z","iopub.execute_input":"2022-09-07T07:32:57.123824Z","iopub.status.idle":"2022-09-07T07:32:57.133763Z","shell.execute_reply.started":"2022-09-07T07:32:57.123782Z","shell.execute_reply":"2022-09-07T07:32:57.132652Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def predict(models, ds):\n    with torch.no_grad():\n        dl_eval = torch.utils.data.DataLoader(ds, batch_size=64, shuffle=False, num_workers=1)\n        preds = []\n        with tqdm(dl_eval, miniters=100, desc='Predict') as progress:\n            for batch_idx, (X) in enumerate(progress):\n                pred = None\n                for model in models:\n                    model.eval()\n                    pred_fold = model.forward(X.to(DEVICE))\n                    if pred is None:\n                        pred = pred_fold / len(models)\n                    else:\n                        pred += pred_fold / len(models)\n                preds.append(pred)\n        preds = torch.concat(preds)                \n        return preds","metadata":{"collapsed":false,"pycharm":{"name":"#%%\n"},"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-07T07:32:57.137045Z","iopub.execute_input":"2022-09-07T07:32:57.137722Z","iopub.status.idle":"2022-09-07T07:32:57.146615Z","shell.execute_reply.started":"2022-09-07T07:32:57.137679Z","shell.execute_reply":"2022-09-07T07:32:57.145596Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def predict_nfold():\n    ds_test = load_ds_test()\n    models = [load_model(f'{MODELS_ROOT}/{CFG[\"WANDB_PROJECT\"]}-{fold}.pth').to(DEVICE) for fold in tqdm(range(N_FOLDS), desc='Loading models')]\n    preds = predict(models, ds_test)\n    del models\n    del ds_test\n    gc.collect()\n    return preds.cpu()\n    \n\ndef predict_and_save():\n    preds = predict_nfold()\n    np.save(CFG[\"TECHNOLOGY\"], preds)\n    \nif PREDICT:\n    predict_and_save()\n    gc.collect()    ","metadata":{"collapsed":false,"pycharm":{"name":"#%%\n"},"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-07T07:32:57.147980Z","iopub.execute_input":"2022-09-07T07:32:57.149067Z","iopub.status.idle":"2022-09-07T07:33:39.729064Z","shell.execute_reply.started":"2022-09-07T07:32:57.149028Z","shell.execute_reply":"2022-09-07T07:33:39.727900Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Submission","metadata":{"pycharm":{"name":"#%% md\n"}}},{"cell_type":"code","source":"def gen_preds(df_eval, pred_file, train_targets_idxcol, test_inputs_idxcol):\n    cite_pred = np.load(pred_file)\n    cols = np.load(train_targets_idxcol, allow_pickle=True)['columns']\n    cols_idx = dict(zip(cols, range(len(cols))))\n    cells = np.load(test_inputs_idxcol, allow_pickle=True)['index']\n    cells_idx = dict(zip(cells, range(len(cells))))\n    return cite_pred[df_eval.cell_id.map(cells_idx), df_eval.gene_id.map(cols_idx)]","metadata":{"collapsed":false,"pycharm":{"name":"#%%\n"},"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-07T07:33:39.730749Z","iopub.execute_input":"2022-09-07T07:33:39.731486Z","iopub.status.idle":"2022-09-07T07:33:39.738642Z","shell.execute_reply.started":"2022-09-07T07:33:39.731420Z","shell.execute_reply":"2022-09-07T07:33:39.737485Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def gen_submission():\n    df_meta = pd.read_parquet(META_FILE)\n    df_evaluation = pd.read_parquet(f'{SPARSE_ROOT}/evaluation.parquet')\n    print('Loaded', f'{SPARSE_ROOT}/evaluation.parquet', df_evaluation.shape)\n\n    # citeseq\n    df_cite_eval = df_evaluation[df_evaluation.cell_id.isin(df_meta.query('technology == \"citeseq\"').cell_id)]\n    if os.path.exists(f'CITEseq.npy'):\n        cite_eval_pred = gen_preds(df_cite_eval, f'{PREDICTIONS_ROOT}/CITEseq.npy', f'{SPARSE_ROOT}/train_cite_targets_idxcol.npz', f'{SPARSE_ROOT}/test_cite_inputs_idxcol.npz')\n        print('Loaded citeseq prediction', cite_eval_pred.shape)\n    else:\n        cite_eval_pred = np.zeros(df_cite_eval.shape[0])\n        print('Empty citeseq prediction', cite_eval_pred.shape)\n    del df_cite_eval\n    gc.collect()\n\n    # multiome\n    df_multi_eval = df_evaluation[df_evaluation.cell_id.isin(df_meta.query('technology == \"multiome\"').cell_id)]\n    if os.path.exists(f'Multiome.npy'):\n        multi_eval_pred = gen_preds(df_multi_eval, f'{PREDICTIONS_ROOT}/Multiome.npy', f'{SPARSE_ROOT}/train_multi_targets_idxcol.npz', f'{SPARSE_ROOT}/test_multi_inputs_idxcol.npz')\n        print(\"Loaded Multiome predictions\", multi_eval_pred.shape)\n    else:\n        multi_eval_pred = np.zeros(df_multi_eval.shape[0])\n        print(\"Empty Multiome predictions\", multi_eval_pred.shape)\n    del df_multi_eval, df_meta\n    gc.collect()\n\n    if SUBMISSION:    \n        print('Generating pure submission')\n        df_evaluation['target'] = np.concatenate([cite_eval_pred, multi_eval_pred])\n        df_evaluation[['row_id', 'target']].to_csv('submission.csv', index=False)\n    elif SUBMISSION_FOR_PATCHING is not None:\n        del df_evaluation\n        gc.collect()\n        df_sub = pd.read_csv(SUBMISSION_FOR_PATCHING)\n        print('Generating patched submission from', SUBMISSION_FOR_PATCHING, df_sub.shape)\n        if PATCH_CITESEQ:\n            preds = np.concatenate([cite_eval_pred, df_sub['target'].tail(len(multi_eval_pred)).values])\n            print(\"Patching CITEseq data\", preds.shape)\n            df_sub['target'] = preds\n        else:\n            multi_patch = np.concatenate([df_sub['target'].head(len(cite_eval_pred)).values, multi_eval_pred])\n            print(\"Patching Multiome data\", multi_patch.shape)\n            df_sub['target'] = multi_patch\n        df_sub.to_csv('submission.csv', index=False)\n    print('Done')\n                \n\ngen_submission()","metadata":{"collapsed":false,"pycharm":{"name":"#%%\n"},"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-07T07:33:39.740182Z","iopub.execute_input":"2022-09-07T07:33:39.740813Z","iopub.status.idle":"2022-09-07T07:37:23.248844Z","shell.execute_reply.started":"2022-09-07T07:33:39.740774Z","shell.execute_reply":"2022-09-07T07:37:23.247692Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<div class=\"alert alert-block alert-danger\" style=\"text-align:center; font-size:20px;\">\n    ❤️ Dont forget to ▲upvote▲ if you find this notebook usefull!  ❤️\n</div>","metadata":{"execution":{"iopub.status.busy":"2022-09-06T19:09:33.212613Z","iopub.execute_input":"2022-09-06T19:09:33.213316Z","iopub.status.idle":"2022-09-06T19:09:33.219733Z","shell.execute_reply.started":"2022-09-06T19:09:33.213282Z","shell.execute_reply":"2022-09-06T19:09:33.218271Z"}}}]}