{"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 with Torch\nThis notebooks is to help competitors that would like to apply Deep Neural Networks to the MSCI data.\nIt is focused on the more challenging Multiome part of the data (but it is trivial to adapt it to to the CITEseq data)\n\nThe main challenge here is that the Multiome data is very large while Kaggle  GPU machines only have 13GB RAM + 16GB GPU Memory.\n\nI found it is actually possible to store all of the dataset in GPU memory using sparse tensor formats. This uses ~12GB on the GPU, leaving only ~4GB for the model parameters and the forward/backward computation. Given that we have only ~100K training examples, I do not expect we will need very large models, so I feel 4GB is actually enough.\n\nIf 4GB is not enough, the other option is to leave the dataset in RAM and load the batches on demand to the GPU (which is what is more classically done). In that case, however, we have only ~1GB RAM left, and will suffer a small performance penalty from having to load the batches to the GPU. But we will have the whole 16GB available for training a complex model. Yet another option is to apply dimensionality reduction to the data beforehand (e.g. with PCA/TruncatedSVD), although I like more the idea of using the raw data and letting the network do its own dimensionality reduction.\n\nThe competition data is pre-encoded as sparse matrices in [this dataset](https://www.kaggle.com/datasets/fabiencrom/multimodal-single-cell-as-sparse-matrix) generated by [this notebook](https://www.kaggle.com/code/fabiencrom/multimodal-single-cell-creating-sparse-data/).\n\nThe model used here is just a very simple MLP. In the current version, I add a `Softplus` activation at the end, considering the values we have to predict are all positives (although I am not sure it will really work better that way).\n\nIn the current version, I also directly optimize the competition metric (row-wise Pearson correlation). Although it does not seem to perform much better than using a simpler Mean Square Error Loss.\n\nThis notebook will train 5 models over 5 folds. The final submission is created in [this notebook](https://www.kaggle.com/fabiencrom/msci-multiome-torch-quickstart-submission).\n\nSo far I did not get results better than the one obtained by the much simpler PCA+Ridge Regression method (that you can find in [this notebook](https://www.kaggle.com/code/ambrosm/msci-multiome-quickstart) as initially proposed by AmbrosM or in [this notebook](https://www.kaggle.com/code/fabiencrom/msci-multiome-quickstart-w-sparse-matrices) for a version using sparse matrices for better results). But I expect it will perform better after working on the architecture/hyperparameters. In any case, I think a deep learning model will be a part of any winning submission.","metadata":{"execution":{"iopub.status.busy":"2022-09-04T11:18:50.155344Z","iopub.execute_input":"2022-09-04T11:18:50.155840Z","iopub.status.idle":"2022-09-04T11:18:56.188046Z","shell.execute_reply.started":"2022-09-04T11:18:50.155762Z","shell.execute_reply":"2022-09-04T11:18:56.187022Z"}}},{"cell_type":"code","source":"import os\nimport copy\nimport gc\nimport math\nimport itertools\nimport pickle\nimport glob\nimport joblib\nimport json\nimport random\nimport re\nimport operator\n\nimport collections\nfrom collections import defaultdict\nfrom operator import itemgetter, attrgetter\n\nfrom tqdm.notebook import tqdm\n\nimport torch\nimport torch.nn as nn\n\nimport numpy as np\nimport pandas as pd\nimport plotly.express as px\n\nimport scipy\n\nimport sklearn\nimport sklearn.cluster\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.model_selection import KFold\nimport sklearn.preprocessing\n\nimport copy\n","metadata":{"execution":{"iopub.status.busy":"2022-09-04T11:33:40.829328Z","iopub.execute_input":"2022-09-04T11:33:40.829804Z","iopub.status.idle":"2022-09-04T11:33:40.839868Z","shell.execute_reply.started":"2022-09-04T11:33:40.829762Z","shell.execute_reply":"2022-09-04T11:33:40.837814Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Score and loss functions\nWe can use either a classic Mean Square Error loss (nn.MSELoss) or use a loss that will optimize directly the competition metric.","metadata":{}},{"cell_type":"code","source":"def partial_correlation_score_torch_faster(y_true, y_pred):\n    \"\"\"Compute the correlation between each rows of the y_true and y_pred tensors.\n    Compatible with backpropagation.\n    \"\"\"\n    y_true_centered = y_true - torch.mean(y_true, dim=1)[:,None]\n    y_pred_centered = y_pred - torch.mean(y_pred, dim=1)[:,None]\n    cov_tp = torch.sum(y_true_centered*y_pred_centered, dim=1)/(y_true.shape[1]-1)\n    var_t = torch.sum(y_true_centered**2, dim=1)/(y_true.shape[1]-1)\n    var_p = torch.sum(y_pred_centered**2, dim=1)/(y_true.shape[1]-1)\n    return cov_tp/torch.sqrt(var_t*var_p)\n\ndef correl_loss(pred, tgt):\n    \"\"\"Loss for directly optimizing the correlation.\n    \"\"\"\n    return -torch.mean(partial_correlation_score_torch_faster(tgt, pred))","metadata":{"execution":{"iopub.status.busy":"2022-09-04T11:33:40.841378Z","iopub.execute_input":"2022-09-04T11:33:40.842138Z","iopub.status.idle":"2022-09-04T11:33:40.852757Z","shell.execute_reply.started":"2022-09-04T11:33:40.842100Z","shell.execute_reply":"2022-09-04T11:33:40.851533Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Config\nWe put the configuration dict at the beginning of the notebook, so that it is easier to find and modify","metadata":{}},{"cell_type":"code","source":"config = dict(\n    layers = [128, 128, 128],\n    patience = 4,\n    max_epochs = 20,\n    criterion = correl_loss, #nn.MSELoss(),\n    \n    n_folds = 5,\n    folds_to_train = [0, 1, 2, 3, 4],\n    kfold_random_state = 42,\n    \n    optimizerparams = dict(\n     lr=1e-3, \n     weight_decay=1e-2\n    ),\n    \n    head=\"softplus\"\n    \n)\n\nINPUT_SIZE = 228942\nOUTPUT_SIZE = 23418\n","metadata":{"execution":{"iopub.status.busy":"2022-09-04T11:33:40.854413Z","iopub.execute_input":"2022-09-04T11:33:40.854954Z","iopub.status.idle":"2022-09-04T11:33:40.863456Z","shell.execute_reply.started":"2022-09-04T11:33:40.854916Z","shell.execute_reply":"2022-09-04T11:33:40.862404Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Utility functions for loading and batching the sparse data in device memory\nThere are a few challenges here:\n- If we directly try to create a torch sparse tensor before moving it to memory, we will get an OOM error\n- Torch CSR tensors cannot be moved to the gpu; so we make our own TorchCSR class that will contain the csr format information\n- torch gpu operations are only compatible with COO tensors (not CSR), so we need some functions to create batches of COO tensors from the TorchCSR objects","metadata":{}},{"cell_type":"code","source":"# Strangely, current torch implementation of csr tensor do not accept to be moved to the gpu. \n# So we make our own equivalent class\nTorchCSR = collections.namedtuple(\"TrochCSR\", \"data indices indptr shape\")\n\ndef load_csr_data_to_gpu(train_inputs):\n    \"\"\"Move a scipy csr sparse matrix to the gpu as a TorchCSR object\n    This try to manage memory efficiently by creating the tensors and moving them to the gpu one by one\n    \"\"\"\n    th_data = torch.from_numpy(train_inputs.data).to(device)\n    th_indices = torch.from_numpy(train_inputs.indices).to(device)\n    th_indptr = torch.from_numpy(train_inputs.indptr).to(device)\n    th_shape = train_inputs.shape\n    return TorchCSR(th_data, th_indices, th_indptr, th_shape)\n\ndef make_coo_batch(torch_csr, indx):\n    \"\"\"Make a coo torch tensor from a TorchCSR object by taking the rows indicated by the indx tensor\n    \"\"\"\n    th_data, th_indices, th_indptr, th_shape = torch_csr\n    start_pts = th_indptr[indx]\n    end_pts = th_indptr[indx+1]\n    coo_data = torch.cat([th_data[start_pts[i]: end_pts[i]] for i in range(len(start_pts))], dim=0)\n    coo_col = torch.cat([th_indices[start_pts[i]: end_pts[i]] for i in range(len(start_pts))], dim=0)\n    coo_row = torch.repeat_interleave(torch.arange(indx.shape[0], device=device), th_indptr[indx+1] - th_indptr[indx])\n    coo_batch = torch.sparse_coo_tensor(torch.vstack([coo_row, coo_col]), coo_data, [indx.shape[0], th_shape[1]])\n    return coo_batch\n\n\ndef make_coo_batch_slice(torch_csr, start, end):\n    \"\"\"Make a coo torch tensor from a TorchCSR object by taking the rows within the (start, end) slice\n    \"\"\"\n    th_data, th_indices, th_indptr, th_shape = torch_csr\n    start_pts = th_indptr[start]\n    end_pts = th_indptr[end]\n    coo_data = th_data[start_pts: end_pts]\n    coo_col = th_indices[start_pts: end_pts]\n    coo_row = torch.repeat_interleave(torch.arange(end-start, device=device), th_indptr[start+1:end+1] - th_indptr[start:end])\n    coo_batch = torch.sparse_coo_tensor(torch.vstack([coo_row, coo_col]), coo_data, [end-start, th_shape[1]])\n    return coo_batch\n","metadata":{"execution":{"iopub.status.busy":"2022-09-04T11:33:40.865401Z","iopub.execute_input":"2022-09-04T11:33:40.865813Z","iopub.status.idle":"2022-09-04T11:33:40.881843Z","shell.execute_reply.started":"2022-09-04T11:33:40.865724Z","shell.execute_reply":"2022-09-04T11:33:40.880784Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# GPU memory DataLoader\nWe create a dataloader that will work with the in-device TorchCSR tensor.\nThis should ensure the fastest training speed.","metadata":{}},{"cell_type":"code","source":"class DataLoaderCOO:\n    \"\"\"Torch compatible DataLoader. Works with in-device TorchCSR tensors.\n    Args:\n         - train_inputs, train_targets: TorchCSR tensors\n         - train_idx: tensor containing the indices of the rows of train_inputs and train_targets that should be used\n         - batch_size, shuffle, drop_last: as in torch.utils.data.DataLoader\n    \"\"\"\n    def __init__(self, train_inputs, train_targets, train_idx=None, \n                 *,\n                batch_size=512, shuffle=False, drop_last=False):\n        self.batch_size = batch_size\n        self.shuffle = shuffle\n        self.drop_last = drop_last\n        \n        self.train_inputs = train_inputs\n        self.train_targets = train_targets\n        \n        self.train_idx = train_idx\n        \n        self.nb_examples = len(self.train_idx) if self.train_idx is not None else len(train_inputs)\n        \n        self.nb_batches = self.nb_examples//batch_size\n        if not drop_last and not self.nb_examples%batch_size==0:\n            self.nb_batches +=1\n        \n    def __iter__(self):\n        if self.shuffle:\n            shuffled_idx = torch.randperm(self.nb_examples, device=device)\n            if self.train_idx is not None:\n                idx_array = self.train_idx[shuffled_idx]\n            else:\n                idx_array = shuffled_idx\n        else:\n            if self.train_idx is not None:\n                idx_array = self.train_idx\n            else:\n                idx_array = None\n            \n        for i in range(self.nb_batches):\n            slc = slice(i*self.batch_size, (i+1)*self.batch_size)\n            if idx_array is None:\n                inp_batch = make_coo_batch_slice(self.train_inputs, i*self.batch_size, (i+1)*self.batch_size)\n                tgt_batch = make_coo_batch_slice(self.train_targets, i*self.batch_size, (i+1)*self.batch_size)\n            else:\n                idx_batch = idx_array[slc]\n                inp_batch = make_coo_batch(self.train_inputs, idx_batch)\n                tgt_batch = make_coo_batch(self.train_targets, idx_batch)\n            yield inp_batch, tgt_batch\n            \n            \n    def __len__(self):\n        return self.nb_batches","metadata":{"execution":{"iopub.status.busy":"2022-09-04T11:33:40.885672Z","iopub.execute_input":"2022-09-04T11:33:40.886335Z","iopub.status.idle":"2022-09-04T11:33:40.899230Z","shell.execute_reply.started":"2022-09-04T11:33:40.886307Z","shell.execute_reply":"2022-09-04T11:33:40.898101Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Simple Model: MLP","metadata":{}},{"cell_type":"code","source":"class MLP(nn.Module):\n    def __init__(self, layer_size_lst, add_final_activation=False):\n        super().__init__()\n        \n        assert len(layer_size_lst) > 2\n        \n        layer_lst = []\n        for i in range(len(layer_size_lst)-1):\n            sz1 = layer_size_lst[i]\n            sz2 = layer_size_lst[i+1]\n            layer_lst += [nn.Linear(sz1, sz2)]\n            if i != len(layer_size_lst)-2 or add_final_activation:\n                 layer_lst += [nn.ReLU()]\n        self.mlp = nn.Sequential(*layer_lst)\n        \n    def forward(self, x):\n        return self.mlp(x)\n    \ndef build_model():\n    model = MLP([INPUT_SIZE] + config[\"layers\"] + [OUTPUT_SIZE])\n    if config[\"head\"] == \"softplus\":\n        model = nn.Sequential(model, nn.Softplus())\n    else:\n        assert config[\"head\"] is None\n    return model","metadata":{"execution":{"iopub.status.busy":"2022-09-04T11:33:40.900887Z","iopub.execute_input":"2022-09-04T11:33:40.901296Z","iopub.status.idle":"2022-09-04T11:33:40.913306Z","shell.execute_reply.started":"2022-09-04T11:33:40.901261Z","shell.execute_reply":"2022-09-04T11:33:40.912121Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training functions","metadata":{}},{"cell_type":"code","source":"def train_fn(model, optimizer, criterion, dl_train):\n\n    loss_list = []\n    model.train()\n    for inpt, tgt in tqdm(dl_train):\n        mb_size = inpt.shape[0]\n        tgt = tgt.to_dense()\n\n        optimizer.zero_grad()\n        pred = model(inpt)\n\n        loss = criterion(pred, tgt)\n        loss_list.append(loss.detach())\n        loss.backward()\n        optimizer.step()\n    avg_loss = sum(loss_list).cpu().item()/len(loss_list)\n    \n    return {\"loss\":avg_loss}\n","metadata":{"execution":{"iopub.status.busy":"2022-09-04T11:33:40.915036Z","iopub.execute_input":"2022-09-04T11:33:40.915461Z","iopub.status.idle":"2022-09-04T11:33:40.924778Z","shell.execute_reply.started":"2022-09-04T11:33:40.915422Z","shell.execute_reply":"2022-09-04T11:33:40.923777Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def valid_fn(model, criterion, dl_valid):\n    loss_list = []\n    all_preds = []\n    all_tgts = []\n    partial_correlation_scores = []\n    model.eval()\n    for inpt, tgt in tqdm(dl_valid):\n        mb_size = inpt.shape[0]\n        tgt = tgt.to_dense()\n        with torch.no_grad():\n            pred = model(inpt)\n        loss = criterion(pred, tgt)\n        loss_list.append(loss.detach())\n        \n        partial_correlation_scores.append(partial_correlation_score_torch_faster(tgt, pred))\n\n    avg_loss = sum(loss_list).cpu().item()/len(loss_list)\n    \n    partial_correlation_scores = torch.cat(partial_correlation_scores)\n\n    score = torch.sum(partial_correlation_scores).cpu().item()/len(partial_correlation_scores) #correlation_score_torch(all_tgts, all_preds)\n    \n    return {\"loss\":avg_loss, \"score\":score}\n","metadata":{"execution":{"iopub.status.busy":"2022-09-04T11:33:40.927411Z","iopub.execute_input":"2022-09-04T11:33:40.928206Z","iopub.status.idle":"2022-09-04T11:33:40.936535Z","shell.execute_reply.started":"2022-09-04T11:33:40.928122Z","shell.execute_reply":"2022-09-04T11:33:40.935306Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_model(model, optimizer, dl_train, dl_valid, save_prefix):\n\n    criterion = config[\"criterion\"]\n    \n    save_params_filename = save_prefix+\"_best_params.pth\"\n    save_config_filename = save_prefix+\"_config.pkl\"\n    best_score = None\n\n    for epoch in range(config[\"max_epochs\"]):\n        log_train = train_fn(model, optimizer, criterion, dl_train)\n        log_valid = valid_fn(model, criterion, dl_valid)\n\n        print(log_train)\n        print(log_valid)\n        \n        score = log_valid[\"score\"]\n        if best_score is None or score > best_score:\n            best_score = score\n            patience = config[\"patience\"]\n            best_params = copy.deepcopy(model.state_dict())\n        else:\n            patience -= 1\n        \n        if patience < 0:\n            print(\"out of patience\")\n            break\n\n\n    torch.save(best_params, save_params_filename)\n    pickle.dump(config,open(save_config_filename, \"wb\"))\n    \n","metadata":{"execution":{"iopub.status.busy":"2022-09-04T11:33:40.943602Z","iopub.execute_input":"2022-09-04T11:33:40.944667Z","iopub.status.idle":"2022-09-04T11:33:40.953098Z","shell.execute_reply.started":"2022-09-04T11:33:40.944629Z","shell.execute_reply":"2022-09-04T11:33:40.952002Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_one_fold(num_fold):\n    \n    train_idx, valid_idx = FOLDS_LIST[num_fold]\n    \n    train_idx = torch.from_numpy(train_idx).to(device)\n    valid_idx = torch.from_numpy(valid_idx).to(device)\n    \n    \n    dl_train = DataLoaderCOO(train_inputs, train_targets, train_idx=train_idx,\n                batch_size=512, shuffle=True, drop_last=True)\n    dl_valid = DataLoaderCOO(train_inputs, train_targets, train_idx=valid_idx,\n                batch_size=512, shuffle=False, drop_last=False)\n    \n    model =  build_model()\n    model.to(device)\n    \n    optimizer = torch.optim.AdamW(model.parameters(), **config[\"optimizerparams\"])\n    \n    train_model(model, optimizer, dl_train, dl_valid, save_prefix=\"f%i\"%num_fold)","metadata":{"execution":{"iopub.status.busy":"2022-09-04T11:33:40.955029Z","iopub.execute_input":"2022-09-04T11:33:40.955510Z","iopub.status.idle":"2022-09-04T11:33:40.964281Z","shell.execute_reply.started":"2022-09-04T11:33:40.955452Z","shell.execute_reply":"2022-09-04T11:33:40.963174Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load Data","metadata":{}},{"cell_type":"code","source":"if torch.cuda.is_available():\n    device = torch.device(\"cuda:0\")\n    print(f\"machine has {torch.cuda.device_count()} cuda devices\")\n    print(f\"model of first cuda device is {torch.cuda.get_device_name(0)}\")\nelse:\n    device = torch.device(\"cpu\")","metadata":{"execution":{"iopub.status.busy":"2022-09-04T11:33:40.966044Z","iopub.execute_input":"2022-09-04T11:33:40.966381Z","iopub.status.idle":"2022-09-04T11:33:40.978377Z","shell.execute_reply.started":"2022-09-04T11:33:40.966348Z","shell.execute_reply":"2022-09-04T11:33:40.977394Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\ntrain_inputs = scipy.sparse.load_npz(\n    \"../input/multimodal-single-cell-as-sparse-matrix/train_multi_inputs_values.sparse.npz\")","metadata":{"execution":{"iopub.status.busy":"2022-09-04T11:33:40.979551Z","iopub.execute_input":"2022-09-04T11:33:40.982961Z","iopub.status.idle":"2022-09-04T11:34:33.824440Z","shell.execute_reply.started":"2022-09-04T11:33:40.982920Z","shell.execute_reply":"2022-09-04T11:34:33.823446Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"We will normalize the input by dividing each column by its max value. This is the simplest reasonable option. Centering the data (i.e. substracting the mean, would destroy the sparsity here)","metadata":{}},{"cell_type":"code","source":"max_inputs = train_inputs.max(axis=0)\nmax_inputs = max_inputs.todense()+1e-10\nnp.savez(\"max_inputs.npz\", max_inputs = max_inputs)\nmax_inputs = torch.from_numpy(max_inputs)[0].to(device)","metadata":{"execution":{"iopub.status.busy":"2022-09-04T11:34:33.825788Z","iopub.execute_input":"2022-09-04T11:34:33.826409Z","iopub.status.idle":"2022-09-04T11:35:09.015112Z","shell.execute_reply.started":"2022-09-04T11:34:33.826371Z","shell.execute_reply":"2022-09-04T11:35:09.012053Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\ntrain_inputs = load_csr_data_to_gpu(train_inputs)\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2022-09-04T11:35:09.019622Z","iopub.execute_input":"2022-09-04T11:35:09.020383Z","iopub.status.idle":"2022-09-04T11:35:10.227008Z","shell.execute_reply.started":"2022-09-04T11:35:09.020344Z","shell.execute_reply":"2022-09-04T11:35:10.225854Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_inputs.data[...] /= max_inputs[train_inputs.indices.long()]","metadata":{"execution":{"iopub.status.busy":"2022-09-04T11:35:10.228523Z","iopub.execute_input":"2022-09-04T11:35:10.229340Z","iopub.status.idle":"2022-09-04T11:35:10.847601Z","shell.execute_reply.started":"2022-09-04T11:35:10.229298Z","shell.execute_reply":"2022-09-04T11:35:10.846426Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.max(train_inputs.data)","metadata":{"execution":{"iopub.status.busy":"2022-09-04T11:35:10.848703Z","iopub.status.idle":"2022-09-04T11:35:10.849861Z","shell.execute_reply.started":"2022-09-04T11:35:10.849605Z","shell.execute_reply":"2022-09-04T11:35:10.849630Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\ntrain_targets = scipy.sparse.load_npz(\n    \"../input/multimodal-single-cell-as-sparse-matrix/train_multi_targets_values.sparse.npz\")","metadata":{"execution":{"iopub.status.busy":"2022-09-04T11:35:10.851025Z","iopub.status.idle":"2022-09-04T11:35:10.851964Z","shell.execute_reply.started":"2022-09-04T11:35:10.851710Z","shell.execute_reply":"2022-09-04T11:35:10.851734Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\ntrain_targets = load_csr_data_to_gpu(train_targets)\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2022-09-04T11:35:10.853231Z","iopub.status.idle":"2022-09-04T11:35:10.854133Z","shell.execute_reply.started":"2022-09-04T11:35:10.853875Z","shell.execute_reply":"2022-09-04T11:35:10.853903Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"assert INPUT_SIZE == train_inputs.shape[1]\nassert OUTPUT_SIZE == train_targets.shape[1]\n\nNB_EXAMPLES = train_inputs.shape[0]\nassert NB_EXAMPLES == train_targets.shape[0]\n\nprint(INPUT_SIZE, OUTPUT_SIZE, NB_EXAMPLES)","metadata":{"execution":{"iopub.status.busy":"2022-09-04T11:30:14.393036Z","iopub.execute_input":"2022-09-04T11:30:14.393880Z","iopub.status.idle":"2022-09-04T11:30:14.403462Z","shell.execute_reply.started":"2022-09-04T11:30:14.393838Z","shell.execute_reply":"2022-09-04T11:30:14.402203Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training\nWe use a rather naive kfold split here, which might not be optimal for this competition.","metadata":{}},{"cell_type":"code","source":"kfold = KFold(n_splits=config[\"n_folds\"], shuffle=True, random_state=config[\"kfold_random_state\"])\nFOLDS_LIST = list(kfold.split(range(train_inputs.shape[0])))","metadata":{"execution":{"iopub.status.busy":"2022-09-04T11:30:14.405270Z","iopub.execute_input":"2022-09-04T11:30:14.405822Z","iopub.status.idle":"2022-09-04T11:30:14.428158Z","shell.execute_reply.started":"2022-09-04T11:30:14.405726Z","shell.execute_reply":"2022-09-04T11:30:14.426967Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for num_fold in config[\"folds_to_train\"]:\n    train_one_fold(num_fold)","metadata":{"execution":{"iopub.status.busy":"2022-09-04T11:30:14.429962Z","iopub.execute_input":"2022-09-04T11:30:14.430455Z","iopub.status.idle":"2022-09-04T11:31:19.728733Z","shell.execute_reply.started":"2022-09-04T11:30:14.430418Z","shell.execute_reply":"2022-09-04T11:31:19.726542Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}