{"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":"code","source":"import torch\nfrom torch.utils.data import Dataset\nimport torch.nn as nn\nfrom torch.utils.data import DataLoader\nimport torch.optim.lr_scheduler as lr_scheduler\n\nfrom collections import OrderedDict\nfrom sklearn import preprocessing","metadata":{"execution":{"iopub.status.busy":"2023-10-13T15:40:07.175648Z","iopub.execute_input":"2023-10-13T15:40:07.176111Z","iopub.status.idle":"2023-10-13T15:40:07.181701Z","shell.execute_reply.started":"2023-10-13T15:40:07.176076Z","shell.execute_reply":"2023-10-13T15:40:07.180835Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Work with data","metadata":{}},{"cell_type":"code","source":"de_train = pd.read_parquet('../input/open-problems-single-cell-perturbations/de_train.parquet')\nde_train","metadata":{"execution":{"iopub.status.busy":"2023-10-13T15:40:08.301670Z","iopub.execute_input":"2023-10-13T15:40:08.302080Z","iopub.status.idle":"2023-10-13T15:40:09.496688Z","shell.execute_reply.started":"2023-10-13T15:40:08.302045Z","shell.execute_reply":"2023-10-13T15:40:09.495709Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"-genes A1BG, A1BG-AS1, …, ZZEF1 (numbering 18,211 in all) - Differential expression value (-log10(p-value) * sign(LFC)) for each gene. Here, LFC is the estimated log-fold change in expression between the treatment and control condition after shrinkage as calculated by Limma. Positive LFC means the gene goes up in the treatment condition relative to the control.","metadata":{}},{"cell_type":"code","source":"#de_train.pivot(columns=['sm_lincs_id', 'sm_name', 'cell_type'], values='control')\nprint(list(de_train['cell_type'].unique()))\nprint(len(list(de_train['sm_name'].unique())))","metadata":{"execution":{"iopub.status.busy":"2023-10-13T15:40:12.627650Z","iopub.execute_input":"2023-10-13T15:40:12.628027Z","iopub.status.idle":"2023-10-13T15:40:12.634204Z","shell.execute_reply.started":"2023-10-13T15:40:12.628000Z","shell.execute_reply":"2023-10-13T15:40:12.633137Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"id_map = pd.read_csv('/kaggle/input/open-problems-single-cell-perturbations/id_map.csv',index_col=0)\nprint(list(id_map['cell_type'].unique()))\nprint(len(list(id_map['sm_name'].unique())))","metadata":{"execution":{"iopub.status.busy":"2023-10-13T15:40:15.496935Z","iopub.execute_input":"2023-10-13T15:40:15.497313Z","iopub.status.idle":"2023-10-13T15:40:15.521150Z","shell.execute_reply.started":"2023-10-13T15:40:15.497284Z","shell.execute_reply":"2023-10-13T15:40:15.520404Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Prepare data","metadata":{}},{"cell_type":"code","source":"num_df = pd.DataFrame()","metadata":{"execution":{"iopub.status.busy":"2023-10-13T15:40:17.040296Z","iopub.execute_input":"2023-10-13T15:40:17.040970Z","iopub.status.idle":"2023-10-13T15:40:17.046802Z","shell.execute_reply.started":"2023-10-13T15:40:17.040930Z","shell.execute_reply":"2023-10-13T15:40:17.045446Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"encoder = preprocessing.LabelEncoder()\nnum_df['cell_type'] = encoder.fit_transform(de_train['cell_type'])\nnum_df['sm_name'] = encoder.fit_transform(de_train['sm_name'])\nnum_df['sm_lincs_id'] =encoder.fit_transform(de_train['sm_lincs_id'])\nnum_df['smiles'] =encoder.fit_transform(de_train['SMILES'])\nnum_df['control'] = de_train['control'].fillna(False).astype('int')\n\nnum_df","metadata":{"execution":{"iopub.status.busy":"2023-10-13T15:40:17.516751Z","iopub.execute_input":"2023-10-13T15:40:17.517416Z","iopub.status.idle":"2023-10-13T15:40:17.535951Z","shell.execute_reply.started":"2023-10-13T15:40:17.517381Z","shell.execute_reply":"2023-10-13T15:40:17.534888Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = num_df.iloc[:,0:2]\noutput_names = de_train.iloc[:,5:].columns.values.tolist()","metadata":{"execution":{"iopub.status.busy":"2023-10-13T15:40:17.965469Z","iopub.execute_input":"2023-10-13T15:40:17.966514Z","iopub.status.idle":"2023-10-13T15:40:18.004391Z","shell.execute_reply.started":"2023-10-13T15:40:17.966465Z","shell.execute_reply":"2023-10-13T15:40:18.003286Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\n\nX_train, X_test, y_train, y_test = train_test_split(df, de_train.iloc[:,5:], test_size=0.3, shuffle=True)","metadata":{"execution":{"iopub.status.busy":"2023-10-13T15:40:18.504379Z","iopub.execute_input":"2023-10-13T15:40:18.504773Z","iopub.status.idle":"2023-10-13T15:40:18.594612Z","shell.execute_reply.started":"2023-10-13T15:40:18.504723Z","shell.execute_reply":"2023-10-13T15:40:18.593667Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# make training and test sets in torch\nX_train = torch.from_numpy(X_train.values).type(torch.Tensor)\nX_test = torch.from_numpy(X_test.values).type(torch.Tensor)\ny_train = torch.from_numpy(y_train.values).type(torch.Tensor)\ny_test = torch.from_numpy(y_test.values).type(torch.Tensor)","metadata":{"execution":{"iopub.status.busy":"2023-10-13T15:40:19.050111Z","iopub.execute_input":"2023-10-13T15:40:19.050461Z","iopub.status.idle":"2023-10-13T15:40:19.067782Z","shell.execute_reply.started":"2023-10-13T15:40:19.050434Z","shell.execute_reply":"2023-10-13T15:40:19.066743Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(X_train.shape)\nprint(y_train.shape)","metadata":{"execution":{"iopub.status.busy":"2023-10-13T15:40:21.283438Z","iopub.execute_input":"2023-10-13T15:40:21.283827Z","iopub.status.idle":"2023-10-13T15:40:21.288339Z","shell.execute_reply.started":"2023-10-13T15:40:21.283795Z","shell.execute_reply":"2023-10-13T15:40:21.287473Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Create model","metadata":{}},{"cell_type":"code","source":"class MLR(nn.Module):\n    def __init__(self, input_dim, hidden_dim, num_layers, output_dim):\n        super (MLR, self).__init__()\n        \n        self.hidden_dim = hidden_dim\n        self.num_layers = num_layers\n        \n        #self.lstm = nn.LSTM(input_dim, hidden_dim, num_layers,dropout=0.25)\n        \n        self.gru = nn.GRU(input_dim, hidden_dim, num_layers, dropout=0.25)\n        self.ln = nn.Linear(hidden_dim, 1)\n        self.end = nn.Linear(1, output_dim)\n\n    def forward(self, x):\n        \n        h0 = torch.zeros(self.num_layers, self.hidden_dim).requires_grad_()\n        #out,_ = self.lstm(x)\n        out,_ = self.gru(x, h0.detach())\n        out = self.ln(out)\n        out = self.end(out)\n        return out","metadata":{"execution":{"iopub.status.busy":"2023-10-13T15:40:22.310883Z","iopub.execute_input":"2023-10-13T15:40:22.311276Z","iopub.status.idle":"2023-10-13T15:40:22.318409Z","shell.execute_reply.started":"2023-10-13T15:40:22.311247Z","shell.execute_reply":"2023-10-13T15:40:22.317624Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = MLR(X_train.shape[1], 16, 2, 18211)","metadata":{"execution":{"iopub.status.busy":"2023-10-13T15:40:23.808861Z","iopub.execute_input":"2023-10-13T15:40:23.810064Z","iopub.status.idle":"2023-10-13T15:40:23.816624Z","shell.execute_reply.started":"2023-10-13T15:40:23.810016Z","shell.execute_reply":"2023-10-13T15:40:23.815606Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(model)","metadata":{"execution":{"iopub.status.busy":"2023-10-13T15:40:24.197876Z","iopub.execute_input":"2023-10-13T15:40:24.199061Z","iopub.status.idle":"2023-10-13T15:40:24.203953Z","shell.execute_reply.started":"2023-10-13T15:40:24.199015Z","shell.execute_reply":"2023-10-13T15:40:24.203276Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.save(model, 'single_cell_perturdations.pth')","metadata":{"execution":{"iopub.status.busy":"2023-10-13T15:40:24.678392Z","iopub.execute_input":"2023-10-13T15:40:24.679073Z","iopub.status.idle":"2023-10-13T15:40:24.686349Z","shell.execute_reply.started":"2023-10-13T15:40:24.679038Z","shell.execute_reply":"2023-10-13T15:40:24.685238Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train model","metadata":{}},{"cell_type":"code","source":"import torch.optim as optim\noptimizer = optim.Adam(model.parameters(), lr = 0.1)\nloss = torch.nn.MSELoss()","metadata":{"execution":{"iopub.status.busy":"2023-10-13T15:40:32.915065Z","iopub.execute_input":"2023-10-13T15:40:32.915437Z","iopub.status.idle":"2023-10-13T15:40:32.921218Z","shell.execute_reply.started":"2023-10-13T15:40:32.915409Z","shell.execute_reply":"2023-10-13T15:40:32.919927Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_model(x_train, y_train, model, loss_function, optimizer):\n    num_batches = len(x_train)\n    total_loss = 0\n    scheduler = lr_scheduler.ExponentialLR(optimizer, gamma = 0.9)#, start_factor=1.0, end_factor=0.5, total_iters=30)\n    model.train()\n\n    output = model(x_train)\n    loss = loss_function(output, y_train)\n\n    optimizer.zero_grad()\n    loss.backward()\n    optimizer.step()\n\n    total_loss += loss.item()\n\n    avg_loss = total_loss / num_batches\n    \n    #oprimize learning rate\n    before_lr = optimizer.param_groups[0][\"lr\"]\n    scheduler.step()\n    after_lr = optimizer.param_groups[0][\"lr\"]\n    print(f\"Train loss: {avg_loss}\\n Learning rate: {before_lr} -> {after_lr}\\n\")","metadata":{"execution":{"iopub.status.busy":"2023-10-13T15:40:34.162273Z","iopub.execute_input":"2023-10-13T15:40:34.162640Z","iopub.status.idle":"2023-10-13T15:40:34.169917Z","shell.execute_reply.started":"2023-10-13T15:40:34.162610Z","shell.execute_reply":"2023-10-13T15:40:34.168667Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def test_model(x_test, y_test, model):\n    num_batches = len(x_test)\n    total_loss = 0\n    model.eval()\n    with torch.no_grad():\n        output = model(x_test)\n        \n        total_loss += np.abs(output - y_test).mean()\n        # Calculate Mean Rowwise Root Mean Squared Error (MRRMSE)\n        output = output.numpy().reshape(x_test.size(0), -1)\n        y_test = y_test.numpy().reshape(x_test.size(0), -1)\n        mrrmse_score = np.sqrt(np.square(y_test - output).mean(axis=1)).mean()\n\n    avg_loss = total_loss / num_batches\n    print(f\"Test loss: {avg_loss}\\n MRRMSE = {mrrmse_score}\\n\")\n    return mrrmse_score","metadata":{"execution":{"iopub.status.busy":"2023-10-13T15:40:34.895271Z","iopub.execute_input":"2023-10-13T15:40:34.895669Z","iopub.status.idle":"2023-10-13T15:40:34.902408Z","shell.execute_reply.started":"2023-10-13T15:40:34.895632Z","shell.execute_reply":"2023-10-13T15:40:34.901672Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mrrmse_list = []\nfor ix_epoch in range(20):\n    print(f\"Epoch {ix_epoch}\\n---------\")\n    train_model(X_train, y_train, model, loss, optimizer)\n    mrrmse_score = test_model(X_test, y_test, model)\n    mrrmse_list.append(mrrmse_score)","metadata":{"execution":{"iopub.status.busy":"2023-10-13T15:40:36.327587Z","iopub.execute_input":"2023-10-13T15:40:36.327931Z","iopub.status.idle":"2023-10-13T15:40:42.665931Z","shell.execute_reply.started":"2023-10-13T15:40:36.327903Z","shell.execute_reply":"2023-10-13T15:40:42.664895Z"},"collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Predict","metadata":{}},{"cell_type":"code","source":"predict_df = pd.DataFrame()\npredict_df['cell_type'] =encoder.fit_transform(id_map['cell_type'])\npredict_df['sm_name'] =encoder.fit_transform(id_map['sm_name'])","metadata":{"execution":{"iopub.status.busy":"2023-10-13T15:40:43.518180Z","iopub.execute_input":"2023-10-13T15:40:43.518542Z","iopub.status.idle":"2023-10-13T15:40:43.526189Z","shell.execute_reply.started":"2023-10-13T15:40:43.518514Z","shell.execute_reply":"2023-10-13T15:40:43.525128Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"inputs = torch.from_numpy(predict_df.values).type(torch.Tensor)\ny_pred_tensor = model(inputs)\ny_pred = y_pred_tensor.detach().cpu().numpy()","metadata":{"execution":{"iopub.status.busy":"2023-10-13T15:40:43.917573Z","iopub.execute_input":"2023-10-13T15:40:43.918725Z","iopub.status.idle":"2023-10-13T15:40:43.963465Z","shell.execute_reply.started":"2023-10-13T15:40:43.918685Z","shell.execute_reply":"2023-10-13T15:40:43.962607Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"solution = pd.concat([id_map, pd.DataFrame(y_pred, index=id_map.index, columns=output_names)], axis=1)","metadata":{"execution":{"iopub.status.busy":"2023-10-13T15:40:45.318362Z","iopub.execute_input":"2023-10-13T15:40:45.318851Z","iopub.status.idle":"2023-10-13T15:40:45.354994Z","shell.execute_reply.started":"2023-10-13T15:40:45.318813Z","shell.execute_reply":"2023-10-13T15:40:45.354104Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(solution.info())","metadata":{"execution":{"iopub.status.busy":"2023-10-13T15:40:47.204466Z","iopub.execute_input":"2023-10-13T15:40:47.205283Z","iopub.status.idle":"2023-10-13T15:40:47.664473Z","shell.execute_reply.started":"2023-10-13T15:40:47.205243Z","shell.execute_reply":"2023-10-13T15:40:47.663360Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"solution.head(10)","metadata":{"execution":{"iopub.status.busy":"2023-10-13T15:40:47.666347Z","iopub.execute_input":"2023-10-13T15:40:47.666663Z","iopub.status.idle":"2023-10-13T15:40:47.692117Z","shell.execute_reply.started":"2023-10-13T15:40:47.666635Z","shell.execute_reply":"2023-10-13T15:40:47.691044Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"solution.to_csv('submission.csv')","metadata":{"execution":{"iopub.status.busy":"2023-10-13T15:41:13.577344Z","iopub.execute_input":"2023-10-13T15:41:13.577801Z","iopub.status.idle":"2023-10-13T15:41:19.701453Z","shell.execute_reply.started":"2023-10-13T15:41:13.577739Z","shell.execute_reply":"2023-10-13T15:41:19.700576Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Analyse the result","metadata":{}},{"cell_type":"code","source":"print(f\"We have: {solution.isnull().sum().unique()} NaN values\")","metadata":{"execution":{"iopub.status.busy":"2023-10-13T15:41:19.702965Z","iopub.execute_input":"2023-10-13T15:41:19.703729Z","iopub.status.idle":"2023-10-13T15:41:19.725958Z","shell.execute_reply.started":"2023-10-13T15:41:19.703699Z","shell.execute_reply":"2023-10-13T15:41:19.725030Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(solution['cell_type'].unique())\nde_train.drop(['sm_lincs_id','SMILES','control'], axis = 1,inplace=True)\nb_cells_df = de_train[de_train['cell_type'] == 'B cells'].reset_index(drop=True)\nb_cells_df = pd.concat([solution[solution['cell_type'] == 'B cells'].reset_index(drop=True),\n                      de_train[de_train['cell_type'] == 'B cells'].reset_index(drop=True)],\n                      sort = False,axis = 0)\nmyeloid_cells_df = pd.concat([solution[solution['cell_type'] == 'Myeloid cells'].reset_index(drop=True),\n                      de_train[de_train['cell_type'] == 'Myeloid cells'].reset_index(drop=True)],\n                      sort = False,axis = 0)","metadata":{"execution":{"iopub.status.busy":"2023-10-13T15:41:19.727218Z","iopub.execute_input":"2023-10-13T15:41:19.727558Z","iopub.status.idle":"2023-10-13T15:41:19.840648Z","shell.execute_reply.started":"2023-10-13T15:41:19.727532Z","shell.execute_reply":"2023-10-13T15:41:19.839902Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"myeloid_cells_df.info()","metadata":{"execution":{"iopub.status.busy":"2023-10-13T15:41:19.842750Z","iopub.execute_input":"2023-10-13T15:41:19.843551Z","iopub.status.idle":"2023-10-13T15:41:21.958542Z","shell.execute_reply.started":"2023-10-13T15:41:19.843521Z","shell.execute_reply":"2023-10-13T15:41:21.957536Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"b_cells_df.info()","metadata":{"execution":{"iopub.status.busy":"2023-10-13T15:41:21.959754Z","iopub.execute_input":"2023-10-13T15:41:21.960159Z","iopub.status.idle":"2023-10-13T15:41:22.411121Z","shell.execute_reply.started":"2023-10-13T15:41:21.960127Z","shell.execute_reply":"2023-10-13T15:41:22.409959Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predict_mean_values = []\ninput_mean_values = []\nfor i in output_names:\n    input_mean_values.append(de_train[i].mean())\n    predict_mean_values.append(solution[i].mean())","metadata":{"execution":{"iopub.status.busy":"2023-10-13T15:41:22.412159Z","iopub.execute_input":"2023-10-13T15:41:22.412412Z","iopub.status.idle":"2023-10-13T15:41:25.426452Z","shell.execute_reply.started":"2023-10-13T15:41:22.412389Z","shell.execute_reply":"2023-10-13T15:41:25.425565Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"corr_coef = np.corrcoef(np.asarray(input_mean_values), np.asarray(predict_mean_values))\nprint(f\"Corellation between input and predict mean values is {corr_coef[0,1]}\")","metadata":{"execution":{"iopub.status.busy":"2023-10-13T15:41:25.428430Z","iopub.execute_input":"2023-10-13T15:41:25.428720Z","iopub.status.idle":"2023-10-13T15:41:25.437566Z","shell.execute_reply.started":"2023-10-13T15:41:25.428695Z","shell.execute_reply":"2023-10-13T15:41:25.436487Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.DataFrame({'input_mean_values' : input_mean_values,\n                   'predict_mean_values' : predict_mean_values})\ndf.head(5)","metadata":{"execution":{"iopub.status.busy":"2023-10-13T15:41:25.438968Z","iopub.execute_input":"2023-10-13T15:41:25.439272Z","iopub.status.idle":"2023-10-13T15:41:25.465043Z","shell.execute_reply.started":"2023-10-13T15:41:25.439246Z","shell.execute_reply":"2023-10-13T15:41:25.463936Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Visualisation","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport seaborn as sns","metadata":{"execution":{"iopub.status.busy":"2023-10-13T15:41:31.334719Z","iopub.execute_input":"2023-10-13T15:41:31.335086Z","iopub.status.idle":"2023-10-13T15:41:31.340228Z","shell.execute_reply.started":"2023-10-13T15:41:31.335060Z","shell.execute_reply":"2023-10-13T15:41:31.338846Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"corr = sns.scatterplot(x=\"input_mean_values\", y=\"predict_mean_values\", data=df)\ncorr.set(title = 'Correlation between mean values of all 18211 genes')","metadata":{"execution":{"iopub.status.busy":"2023-10-13T15:41:31.684960Z","iopub.execute_input":"2023-10-13T15:41:31.685311Z","iopub.status.idle":"2023-10-13T15:41:32.048716Z","shell.execute_reply.started":"2023-10-13T15:41:31.685283Z","shell.execute_reply":"2023-10-13T15:41:32.047987Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.plot(mrrmse_list)\nplt.xlabel('Iterations')\nplt.ylabel('MRRMSE')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-10-13T15:41:33.202604Z","iopub.execute_input":"2023-10-13T15:41:33.203531Z","iopub.status.idle":"2023-10-13T15:41:33.429187Z","shell.execute_reply.started":"2023-10-13T15:41:33.203493Z","shell.execute_reply":"2023-10-13T15:41:33.428114Z"},"trusted":true},"execution_count":null,"outputs":[]}]}