{"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":"# SCP Quickstart\n\nThis notebook shows how to cross-validate a model for the *Open Problems – Single-Cell Perturbations* competition.","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19"}},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport matplotlib.patches as mpatches\nimport seaborn as sns\n\nfrom sklearn.preprocessing import OneHotEncoder, OrdinalEncoder\nfrom sklearn.decomposition import TruncatedSVD\nfrom sklearn.compose import ColumnTransformer\nfrom sklearn.pipeline import make_pipeline\nfrom sklearn.linear_model import Ridge\n","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-09-26T08:41:22.303713Z","iopub.execute_input":"2023-09-26T08:41:22.304125Z","iopub.status.idle":"2023-09-26T08:41:23.520783Z","shell.execute_reply.started":"2023-09-26T08:41:22.304091Z","shell.execute_reply":"2023-09-26T08:41:23.519459Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Reading the data","metadata":{}},{"cell_type":"code","source":"id_map = pd.read_csv('/kaggle/input/open-problems-single-cell-perturbations/id_map.csv',\n                     index_col='id')\n\nde_train = pd.read_parquet('/kaggle/input/open-problems-single-cell-perturbations/de_train.parquet')\ngenes = de_train.columns[5:] # 18211 genes\nde_train","metadata":{"execution":{"iopub.status.busy":"2023-09-26T08:41:23.523180Z","iopub.execute_input":"2023-09-26T08:41:23.523746Z","iopub.status.idle":"2023-09-26T08:41:25.067811Z","shell.execute_reply.started":"2023-09-26T08:41:23.523711Z","shell.execute_reply":"2023-09-26T08:41:25.066431Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Cross-validation\n\nWe can see this competition as a multi-output regression task with 18211 targets and only two features (cell_type and sm_name). Both features are categorical.\n\nThe diagram shows in red which combinations of the two categorical features (cell_type, sm_name) are in the test set.\n\n**Insight:**\n- We want to predict the 18211 differential gene expressions for unseen (cell_type, sm_name) combinations. Cross-validation must simulate this setting. A possible cross-validation strategy makes four folds, of which the diagram shows the first fold in blue.\n","metadata":{}},{"cell_type":"code","source":"\ndef plot_cv_diagram(val_cell_type):\n    cv_diagram = pd.concat([de_train[['cell_type', 'sm_name']],\n                            id_map[['cell_type', 'sm_name']]], axis=0, keys=[2, 1])\n    cv_diagram = cv_diagram.reset_index().drop(columns=['level_1'])\n    cv_diagram = cv_diagram.pivot(index='cell_type', columns='sm_name')\n    cv_diagram.fillna(0, inplace=True)\n    cv_diagram = cv_diagram.droplevel(level=0, axis=1)\n    cv_diagram = cv_diagram.sort_values('Myeloid cells', axis=1)\n    cv_diagram = cv_diagram.sort_index(ascending=False)\n    cv_diagram.loc[val_cell_type] *= 1 + (cv_diagram.loc[['Myeloid cells', 'B cells']] == 1).any().astype(float) * 0.5\n    # 0=Missing\n    # 1=Test\n    # 2=Training\n    # 3=Validation\n    \n    _, (ax1, ax2) = plt.subplots(1, 2, width_ratios=(20, 1), figsize=(36, 3))\n    sns.heatmap(cv_diagram, cbar=False, linewidths=1, ax=ax1, cmap=['k', 'r', 'g', 'b'])\n    \n    ax2.legend(handles=[mpatches.Patch(color='g', label='Training'),\n                        mpatches.Patch(color='b', label='Validation (fold 0)'),\n                        mpatches.Patch(color='r', label='Test'),\n                        mpatches.Patch(color='k', label='Missing')])\n    ax2.axis('off')\n    plt.show()\n    \nplot_cv_diagram('NK cells')\n# plot_cv_diagram('T cells CD8+')","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-09-26T08:41:25.069434Z","iopub.execute_input":"2023-09-26T08:41:25.069807Z","iopub.status.idle":"2023-09-26T08:41:26.893032Z","shell.execute_reply.started":"2023-09-26T08:41:25.069772Z","shell.execute_reply":"2023-09-26T08:41:26.891852Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Our model is simple:\n- We denoise the targets by applying a singular value decomposition. As a consequence of this transformation, we must inverse-tranform the predictions.\n- We use sm_name as our unique feature, one-hot encode it and then use ridge regression to predict the transformed targets.\n- The two hyperparameters `n_components` and `alpha` were tuned for the best cross-validation score.\n","metadata":{}},{"cell_type":"code","source":"temp = de_train.groupby(['cell_type']).size() > 20\nvalidation_cell_types = temp[temp].index # 4 cell types\n\ntrain_sm_names = de_train.query(\"cell_type == 'B cells'\").sm_name.values # 17 compounds including the two control compounds\n\nfeatures = ['cell_type', 'sm_name']\n\ndef cross_val_svd(model, label, n_components=5):\n    mrrmse_list = []\n    for fold, val_cell_type in enumerate(validation_cell_types):\n        mask_va = (de_train.cell_type == val_cell_type) & ~de_train.sm_name.isin(train_sm_names)\n        mask_tr = ~mask_va # 485 or 487 training rows\n        \n        train = de_train[mask_tr]\n        val = de_train[mask_va]\n        y_true = val[genes]\n        \n        svd = TruncatedSVD(n_components=n_components, random_state=1)\n        z_tr = svd.fit_transform(train[genes])\n        \n        model.fit(train[features], z_tr)\n        y_pred = svd.inverse_transform(model.predict(val[features]))\n        \n        mrrmse = np.sqrt(np.square(y_true - y_pred).mean(axis=1)).mean()\n        print(f\"# Fold {fold}: {mrrmse:5.3f} val='{val_cell_type}'\")\n        mrrmse_list.append(mrrmse)\n    mrrmse = np.array(mrrmse_list).mean()\n    print(f\"# Overall {mrrmse:5.3f} {label}\")\n    return","metadata":{"execution":{"iopub.status.busy":"2023-09-26T08:44:01.097239Z","iopub.execute_input":"2023-09-26T08:44:01.097664Z","iopub.status.idle":"2023-09-26T08:44:01.683704Z","shell.execute_reply.started":"2023-09-26T08:44:01.097632Z","shell.execute_reply":"2023-09-26T08:44:01.682389Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"n_components = 100\nalpha = 5\nmodel = make_pipeline(ColumnTransformer([('ohe', OneHotEncoder(), ['sm_name'])]),\n                      Ridge(alpha=alpha, fit_intercept=False))\ncross_val_svd(model, f'svd {n_components=} ridge {alpha=}',\n              n_components=n_components)\n# Fold 0: 1.089 val='NK cells'\n# Fold 1: 0.971 val='T cells CD4+'\n# Fold 2: 0.791 val='T cells CD8+'\n# Fold 3: 1.008 val='T regulatory cells'\n# Overall 0.965 svd n_components=100 ridge alpha=5","metadata":{"execution":{"iopub.status.busy":"2023-09-26T08:44:01.688380Z","iopub.execute_input":"2023-09-26T08:44:01.688779Z","iopub.status.idle":"2023-09-26T08:44:13.583674Z","shell.execute_reply.started":"2023-09-26T08:44:01.688745Z","shell.execute_reply":"2023-09-26T08:44:13.579339Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Submission\n\nWe retrain the model on the full training data and create a submission file.","metadata":{}},{"cell_type":"code","source":"svd = TruncatedSVD(n_components=n_components, random_state=1)\nz_tr = svd.fit_transform(de_train[genes])\nmodel.fit(de_train[features], z_tr)\ny_pred = svd.inverse_transform(model.predict(id_map[features]))\n","metadata":{"execution":{"iopub.status.busy":"2023-09-26T08:44:26.394679Z","iopub.execute_input":"2023-09-26T08:44:26.395062Z","iopub.status.idle":"2023-09-26T08:44:29.538768Z","shell.execute_reply.started":"2023-09-26T08:44:26.395032Z","shell.execute_reply":"2023-09-26T08:44:29.536764Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission = pd.DataFrame(y_pred, columns=genes, index=id_map.index)\ndisplay(submission)\nsubmission.to_csv('submission.csv')\n","metadata":{"execution":{"iopub.status.busy":"2023-09-26T08:41:43.159025Z","iopub.execute_input":"2023-09-26T08:41:43.159701Z","iopub.status.idle":"2023-09-26T08:41:56.197824Z","shell.execute_reply.started":"2023-09-26T08:41:43.159645Z","shell.execute_reply":"2023-09-26T08:41:56.196538Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}