{"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":"## The Predict process of the stacking model ","metadata":{}},{"cell_type":"code","source":"import sys\nsys.path.append(\"../input/cite-stacking\")\nimport gc\nimport numpy as np\nimport pandas as pd\nimport random\nfrom stacking_model import ModelStacking\nfrom tqdm.notebook import tqdm\nfrom sklearn.preprocessing import LabelEncoder,OneHotEncoder","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-11-20T08:26:04.774626Z","iopub.execute_input":"2022-11-20T08:26:04.774998Z","iopub.status.idle":"2022-11-20T08:26:04.780588Z","shell.execute_reply.started":"2022-11-20T08:26:04.774966Z","shell.execute_reply":"2022-11-20T08:26:04.779616Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Since the path have changed so we need to change the predict function\nclass ModelStacking_predict(ModelStacking):\n    def predict(self,dir_n = 0):\n        self.surfix = f\"../input/stacking/results/{dir_n}/\"\n        self.predict_first_layer()\n        self.predict_second_layer()\n        self.predict_third_layer()\n        return self.third_predict\n","metadata":{"execution":{"iopub.status.busy":"2022-11-20T08:32:33.862782Z","iopub.execute_input":"2022-11-20T08:32:33.863190Z","iopub.status.idle":"2022-11-20T08:32:33.869455Z","shell.execute_reply.started":"2022-11-20T08:32:33.863157Z","shell.execute_reply":"2022-11-20T08:32:33.868373Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"feature_path = '/kaggle/input/single-cell-features/'\n# train\ntrain = np.load(feature_path+'train_cite_X.npy')\ntrain_index = np.load(\"../input/multimodal-single-cell-as-sparse-matrix/train_cite_inputs_idxcol.npz\",allow_pickle=True)\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[\"index\"])\ntrain_meta = meta_train[\"cell_type\"].values.reshape(-1, 1)\nohe = OneHotEncoder(sparse=False)\ntrain_meta = ohe.fit_transform(train_meta)\ntrain = np.concatenate([train,train_meta],axis= -1)\n\n# target\ntarget = np.load(feature_path+'train_cite_targets.npy') \ntarget -= target.mean(axis=1).reshape(-1, 1)\ntarget /= target.std(axis=1).reshape(-1, 1)\n\n# test\ntest = np.load(feature_path+'test_cite_X.npy')\n\ntest_index = np.load(\"../input/multimodal-single-cell-as-sparse-matrix/test_cite_inputs_idxcol.npz\",allow_pickle=True)\nmeta_test = meta.reindex(test_index[\"index\"])\ntest_meta = meta_test[\"cell_type\"].values.reshape(-1, 1)\ntest_meta = ohe.transform(test_meta)\ntest = np.concatenate([test,test_meta],axis= -1)\n\n# all\nfea_columns = [f\"fea_{i}\" for i in range(train.shape[1])]\nlab_columns = [f\"lab_{i}\" for i in range(target.shape[1])]\ntrain = pd.DataFrame(train,columns=fea_columns)\ntest = pd.DataFrame(test,columns=fea_columns)\ntarget = pd.DataFrame(target,columns=lab_columns)\nall = pd.concat([train,target],axis=1)\ndel train,target\ngc.collect()\nprint(all.shape,test.shape)","metadata":{"execution":{"iopub.status.busy":"2022-11-20T08:27:02.356695Z","iopub.execute_input":"2022-11-20T08:27:02.357193Z","iopub.status.idle":"2022-11-20T08:27:13.992487Z","shell.execute_reply.started":"2022-11-20T08:27:02.357148Z","shell.execute_reply":"2022-11-20T08:27:13.991240Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"../input/stacking/results/0/layer_1/0/catboost/Fold/model_43.cbm","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"main_config = dict(\n    train_features = fea_columns,   \n    predict_label = lab_columns,\n    random_state = 42,\n    layer_1_model_list = [\"KNN\",\"CNN\",'ridge',\"rf\",\"catboost\",\"torch\"], # [\"KNN\",\"CNN\",'ridge',\"rf\",\"catboost\",\"torch\"],#,\"lgbm\",\"CNN\",\"KernelRidge\",\"ElasticNet\",'ridge',\"rf\",\"et\",\"catboost\",\"torch\",\"KernelRidge\"\n    layer_2_model_list = [\"CNN\",\"catboost\",\"torch\"], # [\"CNN\",\"catboost\",\"torch\"],#,\"lgbm\",,\"catboost\"\"CNN\",\"torch\"\n    layer_3_model = \"mlp\"\n)","metadata":{"execution":{"iopub.status.busy":"2022-11-20T08:27:13.994810Z","iopub.execute_input":"2022-11-20T08:27:13.995246Z","iopub.status.idle":"2022-11-20T08:27:14.001064Z","shell.execute_reply.started":"2022-11-20T08:27:13.995203Z","shell.execute_reply":"2022-11-20T08:27:14.000033Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pickle\nwith open(\"../input/stacking/fold_list.pkl\",\"rb\") as f:\n    fold_list = pickle.load(f)","metadata":{"execution":{"iopub.status.busy":"2022-11-20T08:27:17.349046Z","iopub.execute_input":"2022-11-20T08:27:17.350012Z","iopub.status.idle":"2022-11-20T08:27:17.378599Z","shell.execute_reply.started":"2022-11-20T08:27:17.349974Z","shell.execute_reply":"2022-11-20T08:27:17.377316Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"res_all = []\nfor id,part in enumerate(tqdm(fold_list)):\n    stacking_model = ModelStacking_predict(all,test,meta_train,main_config,part)\n    res = stacking_model.predict(id)\n    res_all.append(res)","metadata":{"execution":{"iopub.status.busy":"2022-11-20T08:33:03.668192Z","iopub.execute_input":"2022-11-20T08:33:03.668596Z","iopub.status.idle":"2022-11-20T08:52:30.889940Z","shell.execute_reply.started":"2022-11-20T08:33:03.668564Z","shell.execute_reply":"2022-11-20T08:52:30.887619Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport seaborn as sns\nplt.figure(figsize=(14,12))\nfor id,res in enumerate(tqdm(res_all)):\n    plt.subplot(3,3,id+1)\n    sns.heatmap(res)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-11-20T09:04:16.794113Z","iopub.execute_input":"2022-11-20T09:04:16.794773Z","iopub.status.idle":"2022-11-20T09:04:47.837681Z","shell.execute_reply.started":"2022-11-20T09:04:16.794739Z","shell.execute_reply":"2022-11-20T09:04:47.836686Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import seaborn as sns\nres_all_ = np.mean(res_all,axis=0)\nsns.heatmap(res_all_)","metadata":{"execution":{"iopub.status.busy":"2022-11-20T09:04:47.839279Z","iopub.execute_input":"2022-11-20T09:04:47.840168Z","iopub.status.idle":"2022-11-20T09:04:58.399658Z","shell.execute_reply.started":"2022-11-20T09:04:47.840127Z","shell.execute_reply":"2022-11-20T09:04:58.398597Z"},"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.csv\")\n    submission.to_csv('submission.csv')\n    print(\"submission.csv saved!\")","metadata":{"execution":{"iopub.status.busy":"2022-11-20T09:05:28.222073Z","iopub.execute_input":"2022-11-20T09:05:28.222510Z","iopub.status.idle":"2022-11-20T09:05:28.228964Z","shell.execute_reply.started":"2022-11-20T09:05:28.222475Z","shell.execute_reply":"2022-11-20T09:05:28.227836Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submit(res_all_,\"../input/4th-solution-ensemble/submission.zip\")","metadata":{"execution":{"iopub.status.busy":"2022-11-20T09:05:37.109017Z","iopub.execute_input":"2022-11-20T09:05:37.109441Z","iopub.status.idle":"2022-11-20T09:09:33.217564Z","shell.execute_reply.started":"2022-11-20T09:05:37.109405Z","shell.execute_reply":"2022-11-20T09:09:33.216491Z"},"trusted":true},"execution_count":null,"outputs":[]}]}