{"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":"## A demo training and predicting process of the NN network. Some of the code are inspired by this [notebook](https://www.kaggle.com/code/fabiencrom/msci-multiome-torch-quickstart-w-sparse-tensors#Training-functions). ","metadata":{}},{"cell_type":"code","source":"import time\nimport torch\nimport numpy as np\nimport pandas as pd\nimport pickle\nimport glob\nimport copy \nimport os\nimport gc\nfrom sklearn.preprocessing import LabelEncoder,StandardScaler\nfrom sklearn.model_selection import KFold,GroupKFold\nfrom torch.utils import tensorboard\nfrom tqdm.notebook import tqdm","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-11-18T11:35:24.676263Z","iopub.execute_input":"2022-11-18T11:35:24.676607Z","iopub.status.idle":"2022-11-18T11:35:27.041698Z","shell.execute_reply.started":"2022-11-18T11:35:24.676578Z","shell.execute_reply":"2022-11-18T11:35:27.040776Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = np.load(\"../input/cite-final/new_cite_train_final.npz\")[\"arr_0\"]\ntarget = pd.read_hdf(\"../input/open-problems-multimodal/train_cite_targets.h5\").values\ntrain.shape,target.shape","metadata":{"execution":{"iopub.status.busy":"2022-11-18T11:42:22.415513Z","iopub.execute_input":"2022-11-18T11:42:22.415909Z","iopub.status.idle":"2022-11-18T11:42:23.501668Z","shell.execute_reply.started":"2022-11-18T11:42:22.415875Z","shell.execute_reply":"2022-11-18T11:42:23.500768Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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    if type(y_true) == np.ndarray: y_true = torch.tensor(y_true)\n    if type(y_pred) == np.ndarray: y_pred = torch.tensor(y_pred)\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-11-18T11:35:55.747136Z","iopub.execute_input":"2022-11-18T11:35:55.747566Z","iopub.status.idle":"2022-11-18T11:35:55.759941Z","shell.execute_reply.started":"2022-11-18T11:35:55.747528Z","shell.execute_reply":"2022-11-18T11:35:55.758707Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"config = dict(\n    atte_dims = 128,\n    output_num = target.shape[1],\n    input_num = train.shape[1],\n    dropout = 0.1,\n    \n    layers = 9,\n    patience = 5,\n    max_epochs = 100,\n    criterion = correl_loss,\n    batch_size = 128,\n    mlp_dims = [train.shape[1]*2,train.shape[1]],\n\n    n_folds = 3,\n    folds_to_train = [0,1,2],\n    kfold_random_state = 42,\n\n    tb_dir = \"./log/\",\n\n    optimizer = torch.optim.AdamW,\n    optimizerparams = dict(lr=1e-4, weight_decay=1e-2,amsgrad= True),\n    \n    scheduler = torch.optim.lr_scheduler.MultiStepLR,\n    schedulerparams = dict(milestones=[6,10,15,20,25,30], gamma=0.1,verbose  = False), #9,12,15,20,25,30\n    min_epoch = 11,\n)","metadata":{"execution":{"iopub.status.busy":"2022-11-18T11:42:35.051302Z","iopub.execute_input":"2022-11-18T11:42:35.051656Z","iopub.status.idle":"2022-11-18T11:42:35.058845Z","shell.execute_reply.started":"2022-11-18T11:42:35.051627Z","shell.execute_reply":"2022-11-18T11:42:35.057666Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_index = np.load(\"../input/multimodal-single-cell-as-sparse-matrix/train_cite_inputs_idxcol.npz\",allow_pickle=True)[\"index\"]\nmeta = pd.read_csv(\"../input/open-problems-multimodal/metadata.csv\",index_col = \"cell_id\")\nmeta = meta[meta.technology==\"citeseq\"]\nlbe = LabelEncoder()\nmeta[\"cell_type\"] = lbe.fit_transform(meta[\"cell_type\"])\nmeta[\"gender\"] = meta.apply(lambda x:0 if x[\"donor\"]==13176 else 1,axis =1)\nmeta_train = meta.reindex(train_index)\ntrain_meta = meta_train[\"gender\"].values.reshape(-1, 1)\ntrain = np.concatenate([train,train_meta],axis= -1)\ntrain_meta = meta_train[\"cell_type\"].values.reshape(-1, 1)\ntrain = np.concatenate([train,train_meta],axis= -1)\ntrain.shape","metadata":{"execution":{"iopub.status.busy":"2022-11-18T11:36:42.651953Z","iopub.execute_input":"2022-11-18T11:36:42.652308Z","iopub.status.idle":"2022-11-18T11:36:44.299461Z","shell.execute_reply.started":"2022-11-18T11:36:42.652276Z","shell.execute_reply":"2022-11-18T11:36:44.298440Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Cite_Trainer:\n    def __init__(self,device):\n        self.device = device\n        self.time_stamp = str(time.time())\n        \n    def train_fn(self,model, optimizer, criterion, dl_train):\n        loss_list = []\n        all_mae = []\n        all_mse = []\n        model.train()\n        for inpt, tgt in dl_train:\n            self.train_steps +=1\n            mb_size = inpt.shape[0]\n\n            optimizer.zero_grad()\n            inpt = inpt.to(self.device)\n            tgt = tgt.to(self.device)\n            pred = model(inpt)\n\n            loss = criterion(pred, tgt)\n            loss_list.append(loss.detach())\n            self.logger.add_scalar(\"train/loss_step\",loss.detach(),self.train_steps)\n            loss.backward()\n            optimizer.step()\n            mae = torch.nn.functional.l1_loss(pred,tgt)\n            mse = torch.nn.functional.mse_loss(pred,tgt)\n            self.logger.add_scalar(\"train/mae_step\",mae,self.train_steps)\n            self.logger.add_scalar(\"train/mse_step\",mse,self.train_steps)\n            all_mae.append(mae.detach())\n            all_mse.append(mse.detach())\n\n        avg_loss = sum(loss_list).cpu().item()/len(loss_list)\n        mae_score = sum(all_mae).cpu().item()/len(all_mae) \n        mse_score = sum(all_mse).cpu().item()/len(all_mse) \n        self.logger.add_scalar(\"train/loss_epoch\",avg_loss,self.train_epochs)\n        self.logger.add_scalar(\"train/mae\",mae_score,self.train_epochs)\n        self.logger.add_scalar(\"train/mse\",mse_score,self.train_epochs)\n\n        lr = optimizer.param_groups[0][\"lr\"]\n        self.logger.add_scalar(\"val/learning_rate\",lr,self.val_epochs)\n\n        return {\"loss\":avg_loss}\n\n    def valid_fn(self,model, criterion, dl_valid):\n        loss_list = []\n        all_mae = []\n        all_mse = []\n        partial_correlation_scores = []\n        model.eval()\n        for inpt, tgt in dl_valid:\n            self.val_steps += 1\n            mb_size = inpt.shape[0]\n            inpt = inpt.to(self.device)\n            tgt = tgt.to(self.device)\n            with torch.no_grad():\n                pred = model(inpt)\n            loss = criterion(pred, tgt)\n            mae = torch.nn.functional.l1_loss(pred,tgt)\n            mse = torch.nn.functional.mse_loss(pred,tgt)\n            self.logger.add_scalar(\"val/loss_step\",loss,self.val_steps)\n\n            self.logger.add_scalar(\"val/mae_step\",mae,self.val_steps)\n            self.logger.add_scalar(\"val/mse_step\",mse,self.val_steps)\n            all_mae.append(mae.detach())\n            all_mse.append(mse.detach())\n            loss_list.append(loss.detach())\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        partial_correlation_scores = torch.cat(partial_correlation_scores)\n        score = torch.sum(partial_correlation_scores).cpu().item()/len(partial_correlation_scores) #correlation_score_torch(all_tgts, all_preds)\n        mae_score = sum(all_mae).cpu().item()/len(all_mae) \n        mse_score = sum(all_mse).cpu().item()/len(all_mse) \n        self.logger.add_scalar(\"val/loss_epoch\",avg_loss,self.val_epochs)\n        self.logger.add_scalar(\"val/pearson\",score,self.val_epochs)\n        self.logger.add_scalar(\"val/mae\",mae_score,self.val_epochs)\n        self.logger.add_scalar(\"val/mse\",mse_score,self.val_epochs)\n\n        return {\"loss\":avg_loss, \"score\":score}\n\n    def train_model(self,model, optimizer, scheduler,dl_train, dl_valid, save_prefix):\n\n        criterion = self.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 tqdm(range(self.config[\"max_epochs\"])):\n            self.train_epochs += 1\n            self.val_epochs += 1\n            log_train = self.train_fn(model, optimizer, criterion, dl_train)\n            log_valid = self.valid_fn(model, criterion, dl_valid)\n            self.model = model\n            \n            scheduler.step()\n            print(f\"epoch-{epoch} train_loss:{log_train['loss']} val_loss:{log_valid['loss']} corr_score:{log_valid['score']}\")\n            score = log_valid[\"score\"]\n            if best_score is None or score > best_score:\n                best_score = score\n                patience = self.config[\"patience\"]\n                best_params = copy.deepcopy(model.state_dict())\n                torch.save(best_params, save_params_filename)\n                with open(save_config_filename, \"wb+\") as f:\n                    pickle.dump(self.config,f)      \n            else:\n                patience -= 1\n            \n            if patience < 0 and self.train_epochs > self.config[\"min_epoch\"]:\n                print(\"out of patience\")\n                break\n\n        return best_score\n\n    def train_one_fold(self,num_fold,FOLDS_LIST,train_inputs,train_targets,model,config):\n        self.config = config\n        self.logger = tensorboard.SummaryWriter(self.config[\"tb_dir\"]+f\"{self.time_stamp}_{num_fold}/\")\n        self.train_steps = 0\n        self.train_epochs = 0\n        self.val_steps = 0\n        self.val_epochs = 0\n\n        train_idx, valid_idx = FOLDS_LIST[num_fold]\n        \n        train_inputs = torch.tensor(train_inputs,dtype=torch.float)\n        train_targets = torch.tensor(train_targets,dtype=torch.float)\n        \n        train_data = train_inputs[train_idx]\n        valid_data = train_inputs[valid_idx]\n        train_target = train_targets[train_idx]\n        valid_target = train_targets[valid_idx]\n\n        ds_train = torch.utils.data.TensorDataset(train_data,train_target)\n        ds_valid = torch.utils.data.TensorDataset(valid_data,valid_target)\n\n        dl_train = torch.utils.data.DataLoader(ds_train,\n                    batch_size=self.config[\"batch_size\"], shuffle=True, drop_last=False)\n        dl_valid = torch.utils.data.DataLoader(ds_valid, \n                    batch_size=self.config[\"batch_size\"], shuffle=False, drop_last=False)\n        \n        model.to(self.device)\n        \n        optimizer = self.config[\"optimizer\"](model.parameters(), **self.config[\"optimizerparams\"])\n        scheduler = self.config[\"scheduler\"](optimizer, **self.config[\"schedulerparams\"])\n      \n        best_score = self.train_model(model, optimizer, scheduler,dl_train, dl_valid, save_prefix=\"f%i\"%num_fold)\n\n        return best_score","metadata":{"execution":{"iopub.status.busy":"2022-11-18T11:36:44.301608Z","iopub.execute_input":"2022-11-18T11:36:44.302266Z","iopub.status.idle":"2022-11-18T11:36:44.330884Z","shell.execute_reply.started":"2022-11-18T11:36:44.302217Z","shell.execute_reply":"2022-11-18T11:36:44.329564Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class cell(torch.nn.Module):\n    def __init__(self,input_dim,out_dim,dropout=0.1):\n        super().__init__()\n        self.weight_1 = torch.nn.Sequential(\n            torch.nn.Linear(input_dim,input_dim),\n            # torch.nn.Mish(),\n            torch.nn.Softmax(dim= -1),\n        )\n        self.linear_0 = torch.nn.Sequential(\n            torch.nn.Linear(input_dim,input_dim),\n            torch.nn.Dropout(dropout),\n        ) \n        self.linear_1 = torch.nn.Sequential(\n            torch.nn.Linear(input_dim,input_dim),\n            torch.nn.Mish(),\n            )\n\n        self.bn_1 = torch.nn.LayerNorm((input_dim))\n        self.bn_2 = torch.nn.LayerNorm((out_dim))\n        self.bn_3 = torch.nn.LayerNorm((out_dim))\n\n        self.linear_2 = torch.nn.Sequential(\n\n            torch.nn.Linear(input_dim,out_dim),\n            torch.nn.Dropout(dropout),\n            # torch.nn.Mish(),\n            torch.nn.Linear(out_dim,out_dim),\n            # torch.nn.Dropout(dropout),\n            torch.nn.Mish(),\n            \n        )\n        self.linear_3 = torch.nn.Sequential(\n\n            torch.nn.Linear(input_dim,out_dim),\n            torch.nn.Dropout(dropout),\n            # torch.nn.Mish(),\n            torch.nn.Linear(out_dim,out_dim),\n            # torch.nn.Dropout(dropout),\n            torch.nn.Mish(),\n            \n        )\n\n    def forward(self,x):\n        x_1 = self.linear_1(self.linear_0(x) * self.weight_1(x))\n        x = self.bn_1(x_1+x)\n        x = self.bn_2(x+self.linear_2(x))\n        x = self.bn_3(x+self.linear_3(x))\n        return x","metadata":{"execution":{"iopub.status.busy":"2022-11-18T11:36:06.249096Z","iopub.execute_input":"2022-11-18T11:36:06.249442Z","iopub.status.idle":"2022-11-18T11:36:06.259654Z","shell.execute_reply.started":"2022-11-18T11:36:06.249413Z","shell.execute_reply":"2022-11-18T11:36:06.258679Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class modules(torch.nn.Module):\n    def __init__(self,config):\n        super().__init__()\n        output_num = config[\"output_num\"]\n        input_num = config[\"input_num\"]\n        dropout = config[\"dropout\"]\n        mlp_dims = config[\"mlp_dims\"]\n        self.model = torch.nn.ModuleList()\n        # might be wrong\n        in_dim_array = [\n            [mlp_dims[1],mlp_dims[1],mlp_dims[1],mlp_dims[1],mlp_dims[1]],\n            [mlp_dims[1],mlp_dims[1],mlp_dims[1],mlp_dims[1],mlp_dims[1]],\n            [mlp_dims[1],mlp_dims[1],mlp_dims[1],mlp_dims[1],mlp_dims[1]],\n            # [mlp_dims[1],mlp_dims[1],mlp_dims[1],mlp_dims[1]],\n        ]\n        middle_dim_array = [\n            [256,256,256,256],\n            [256,256,256,256],\n            [256,256,256,256],\n            [256,256,256,256],\n        ]\n        self.out_dim_array = [\n            [mlp_dims[1],mlp_dims[1],mlp_dims[1],mlp_dims[1],mlp_dims[1]],\n            [mlp_dims[1],mlp_dims[1],mlp_dims[1],mlp_dims[1],mlp_dims[1]],\n            [mlp_dims[1],mlp_dims[1],mlp_dims[1],mlp_dims[1],mlp_dims[1]],\n            [mlp_dims[1],mlp_dims[1],mlp_dims[1],mlp_dims[1],mlp_dims[1]],\n        ]\n\n        for i in range(len(in_dim_array)): # 行\n            temp_model = torch.nn.ModuleList()\n            for j in range(len(in_dim_array[0])): # 列\n                dim_in = in_dim_array[i][j]\n                dim_out = self.out_dim_array[i][j]\n                temp_model.append(\n                    torch.nn.Sequential(\n                        cell(dim_in,dim_out,dropout),\n                        # cell(dim_in,dim_out,dropout),\n                    )\n                )\n            self.model.append(temp_model)","metadata":{"execution":{"iopub.status.busy":"2022-11-18T11:36:06.999450Z","iopub.execute_input":"2022-11-18T11:36:07.000524Z","iopub.status.idle":"2022-11-18T11:36:07.012113Z","shell.execute_reply.started":"2022-11-18T11:36:07.000476Z","shell.execute_reply":"2022-11-18T11:36:07.011076Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class MLP(torch.nn.Module):\n    def __init__(self,config):\n        super().__init__()\n        output_num = config[\"output_num\"]\n        self.input_num = config[\"input_num\"]\n        dropout = config[\"dropout\"]\n        mlp_dims = config[\"mlp_dims\"]\n\n\n        self.backbone = torch.nn.Linear(self.input_num ,self.input_num)\n        self.embedding_1 = torch.nn.Embedding(2,self.input_num)\n        self.embedding_2 = torch.nn.Embedding(7,self.input_num)\n\n        self.model = modules(config)\n\n        tail_input_dim = np.sum(np.array(self.model.out_dim_array)[-2:,-2:])\n        self.tail = torch.nn.Sequential(\n            torch.nn.Linear(tail_input_dim,mlp_dims[1]*4),\n            torch.nn.Mish(),\n            torch.nn.Dropout(dropout),\n            torch.nn.Linear(mlp_dims[1]*4,mlp_dims[1]*2),\n            torch.nn.Mish(),\n            torch.nn.Dropout(dropout),\n            torch.nn.Linear(mlp_dims[1]*2,output_num),\n            torch.nn.Mish(),\n        )\n        \n\n    def forward(self,xin):\n        xin = self.backbone(xin[:,:self.input_num]) + self.embedding_2(xin[:,-1].int())  # + self.embedding_1(xin[:,-2].int())# \n        # neck\n        temp_model_list = self.model.model[0]\n        res_array = []\n        res_list  = []\n        temp_out = xin\n        for id,i in enumerate(temp_model_list):\n            temp_out = i(\n                temp_out+xin\n                )\n            res_list.append(temp_out) \n        # res_array += res_list[-2:]\n\n        temp_model_list = self.model.model[1]\n        temp_out = xin\n        for i in range(len(temp_model_list)):\n            temp_out = temp_model_list[i](\n                res_list[i]+temp_out\n                )\n            res_list[i] = temp_out\n        res_array += res_list[-2:]\n\n        temp_model_list = self.model.model[2]\n        temp_out = xin\n        for i in range(len(temp_model_list)):\n            temp_out = temp_model_list[i](\n                res_list[i]+temp_out\n            )\n            res_list[i] = temp_out\n        res_array += res_list[-2:]\n\n        res_array = torch.concat(res_array,dim = -1)\n        res_array = self.tail(res_array)\n\n        return res_array\n","metadata":{"execution":{"iopub.status.busy":"2022-11-18T11:36:08.651642Z","iopub.execute_input":"2022-11-18T11:36:08.653503Z","iopub.status.idle":"2022-11-18T11:36:08.666422Z","shell.execute_reply.started":"2022-11-18T11:36:08.653465Z","shell.execute_reply":"2022-11-18T11:36:08.665126Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Train","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\")\n\ntrainer = Cite_Trainer(device)\nkfold = GroupKFold(n_splits=config[\"n_folds\"]) # , shuffle=True, random_state=config[\"kfold_random_state\"]\nFOLDS_LIST = list(kfold.split(range(train.shape[0]),groups= meta_train.donor)) #\nprint(\"Training started\")\nfold_scores = []\nfor num_fold in config[\"folds_to_train\"]:\n    model = MLP(config)\n    best_score = trainer.train_one_fold(num_fold,FOLDS_LIST,train,target,model,config)\n    fold_scores.append(best_score)\nprint(\"\\n\")\nprint(f\"Final average score is {sum(fold_scores)/len(fold_scores)}\")\n    ","metadata":{"execution":{"iopub.status.busy":"2022-11-18T11:42:44.414180Z","iopub.execute_input":"2022-11-18T11:42:44.414564Z","iopub.status.idle":"2022-11-18T12:00:21.781800Z","shell.execute_reply.started":"2022-11-18T11:42:44.414530Z","shell.execute_reply":"2022-11-18T12:00:21.780787Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fold_scores","metadata":{"execution":{"iopub.status.busy":"2022-11-18T12:00:26.657183Z","iopub.execute_input":"2022-11-18T12:00:26.657550Z","iopub.status.idle":"2022-11-18T12:00:26.664404Z","shell.execute_reply.started":"2022-11-18T12:00:26.657517Z","shell.execute_reply":"2022-11-18T12:00:26.663307Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Test","metadata":{}},{"cell_type":"code","source":"class Tester:\n    def __init__(self,device,config):\n        self.device = device\n        self.config = config\n\n    def std(self,x):\n        return (x - np.mean(x,axis=1).reshape(-1,1)) / np.std(x,axis=1).reshape(-1,1)\n    \n    def test_fn_ensemble(self,model_list, dl_test):\n        \n        res = np.zeros(\n            (self.len, self.config[\"output_num\"]), )\n        \n        for model in model_list:\n            model.eval()\n            \n        cur = 0\n        for inpt in tqdm(dl_test):\n            inpt = inpt[0]\n            mb_size = inpt.shape[0]\n\n            with torch.no_grad():\n                pred_list = []\n                inpt = inpt.to(self.device)\n                for id,model in enumerate(model_list):\n                    model.to(self.device)\n                    model.eval()\n                    pred = model(inpt)\n                    model.to(\"cpu\")\n                    pred = self.std(pred.cpu().numpy())* self.weight[id]\n                    pred_list.append(pred)\n                pred = sum(pred_list)/len(pred_list)\n                \n            res[cur:cur+pred.shape[0]] = pred\n            cur += pred.shape[0]\n                \n        return {\"preds\":res}\n\n    def load_model(self,path ):\n        model_list = []\n        for fn in tqdm(glob.glob(path)):\n            prefix = fn[:-len(\"_best_params.pth\")]\n            config_fn = prefix + \"_config.pkl\"\n            \n            config = pickle.load(open(config_fn, \"rb\"))\n\n            model = MLP(config)\n            model.to(\"cpu\")\n            \n            params = torch.load(fn)\n            model.load_state_dict(params)\n            \n            model_list.append(model)\n        print(\"model loaded\")\n        return model_list\n    \n    def load_data(self,test ):\n        print(\"test inputs loaded\")\n        print(test.shape)\n        self.len = test.shape[0]\n        test = torch.tensor(test,dtype = torch.float)\n        test = torch.utils.data.TensorDataset(test)\n        return test\n\n    def test(self,test,model_path = \"./*_best_params.pth\",weight = fold_scores):\n        self.weight = weight\n        model_list = self.load_model(model_path)\n        test_inputs = self.load_data(test)\n        gc.collect()\n        dl_test = torch.utils.data.DataLoader(test_inputs, batch_size=4096, shuffle=False, drop_last=False)\n        test_pred = self.test_fn_ensemble(model_list, dl_test)[\"preds\"]\n        del model_list\n        del dl_test\n        del test_inputs\n        gc.collect()\n        print(test_pred.shape)\n        np.save(\"test_pred.npy\",test_pred)\n        return test_pred\n        ","metadata":{"execution":{"iopub.status.busy":"2022-11-18T12:00:30.356802Z","iopub.execute_input":"2022-11-18T12:00:30.357207Z","iopub.status.idle":"2022-11-18T12:00:30.372895Z","shell.execute_reply.started":"2022-11-18T12:00:30.357173Z","shell.execute_reply":"2022-11-18T12:00:30.371886Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test = np.load(\"../input/cite-final/new_cite_test_final.npz\")[\"arr_0\"]\ntest_index = np.load(\"../input/multimodal-single-cell-as-sparse-matrix/test_cite_inputs_idxcol.npz\",allow_pickle=True)[\"index\"]\nmeta_test = meta.reindex(test_index)\ntest_meta = meta_test[\"gender\"].values.reshape(-1, 1)\ntest = np.concatenate([test,test_meta],axis= -1)\ntest_meta = meta_test[\"cell_type\"].values.reshape(-1, 1)\ntest = np.concatenate([test,test_meta],axis= -1)\ntest.shape","metadata":{"execution":{"iopub.status.busy":"2022-11-18T12:00:33.812915Z","iopub.execute_input":"2022-11-18T12:00:33.813267Z","iopub.status.idle":"2022-11-18T12:00:36.550688Z","shell.execute_reply.started":"2022-11-18T12:00:33.813236Z","shell.execute_reply":"2022-11-18T12:00:36.549669Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tester = Tester( device,config)\ntest_pred = tester.test(test)","metadata":{"execution":{"iopub.status.busy":"2022-11-18T12:00:47.583530Z","iopub.execute_input":"2022-11-18T12:00:47.583943Z","iopub.status.idle":"2022-11-18T12:01:02.197437Z","shell.execute_reply.started":"2022-11-18T12:00:47.583909Z","shell.execute_reply":"2022-11-18T12:01:02.196328Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import seaborn as sns\nsns.heatmap(test_pred)","metadata":{"execution":{"iopub.status.busy":"2022-11-18T12:01:02.656481Z","iopub.execute_input":"2022-11-18T12:01:02.657016Z","iopub.status.idle":"2022-11-18T12:01:10.870620Z","shell.execute_reply.started":"2022-11-18T12:01:02.656983Z","shell.execute_reply":"2022-11-18T12:01:10.869685Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def submit(test_pred,multi_path):\n    submission = pd.read_csv(multi_path,index_col = 0)\n    submission = submission[\"target\"]\n    print(\"data loaded\")\n    submission.iloc[:len(test_pred.ravel())] = test_pred.ravel()\n    assert not submission.isna().any()\n    # submission = submission.round(6) # reduce the size of the csv\n    print(\"start -> submission.zip\")\n    submission.to_csv('submission.zip')\n    print(\"submission.zip saved!\")","metadata":{"execution":{"iopub.status.busy":"2022-11-18T12:01:17.103230Z","iopub.execute_input":"2022-11-18T12:01:17.103987Z","iopub.status.idle":"2022-11-18T12:01:17.109955Z","shell.execute_reply.started":"2022-11-18T12:01:17.103949Z","shell.execute_reply":"2022-11-18T12:01:17.108925Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submit(test_pred,\"../input/4th-solution-ensemble/submission.zip\")","metadata":{"execution":{"iopub.status.busy":"2022-11-18T12:01:21.466255Z","iopub.execute_input":"2022-11-18T12:01:21.466631Z","iopub.status.idle":"2022-11-18T12:08:02.309839Z","shell.execute_reply.started":"2022-11-18T12:01:21.466597Z","shell.execute_reply":"2022-11-18T12:08:02.308637Z"},"trusted":true},"execution_count":null,"outputs":[]}]}