{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":59094,"databundleVersionId":7010844,"sourceType":"competition"}],"dockerImageVersionId":30664,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport torch\nimport random\nimport numpy as np\nimport pandas as pd\nimport torch.nn as nn\n\nfrom glob import glob\nfrom tqdm import tqdm\nfrom copy import deepcopy\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.model_selection import StratifiedKFold","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-03-18T05:33:46.229527Z","iopub.execute_input":"2024-03-18T05:33:46.229989Z","iopub.status.idle":"2024-03-18T05:33:49.557342Z","shell.execute_reply.started":"2024-03-18T05:33:46.229940Z","shell.execute_reply":"2024-03-18T05:33:49.556060Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Stats Features","metadata":{}},{"cell_type":"code","source":"df_train = pd.read_parquet('../input/open-problems-single-cell-perturbations/de_train.parquet')\nid_map = pd.read_csv('../input/open-problems-single-cell-perturbations/id_map.csv')\nde_cell_type = df_train.iloc[:, [0] + list(range(5, df_train.shape[1]))]\nde_sm_name = df_train.iloc[:, [1] + list(range(5, df_train.shape[1]))]","metadata":{"execution":{"iopub.status.busy":"2024-03-18T05:33:49.560157Z","iopub.execute_input":"2024-03-18T05:33:49.560880Z","iopub.status.idle":"2024-03-18T05:33:52.168143Z","shell.execute_reply.started":"2024-03-18T05:33:49.560826Z","shell.execute_reply":"2024-03-18T05:33:52.166840Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"median_cell_type = de_cell_type.groupby('cell_type').median()\nmedian_sm_name = de_sm_name.groupby('sm_name').median()\n\nct_dict = dict()\nmedian_values_ct = median_cell_type.values\nfor k, v in zip(median_cell_type.index, median_values_ct):\n    ct_dict[k] = torch.tensor(v).float()\n    \nsm_dict = dict()\nmedian_values_sm = median_sm_name.values\nfor k, v in zip(median_sm_name.index, median_values_sm):\n    sm_dict[k] = torch.tensor(v).float()\n    \ntotal_dict = dict()\ntotal_dict['sm_name'] = sm_dict\ntotal_dict['cell_type'] = ct_dict\ntorch.save(total_dict, 'median_values.pt')","metadata":{"execution":{"iopub.status.busy":"2024-03-18T05:33:52.169727Z","iopub.execute_input":"2024-03-18T05:33:52.170096Z","iopub.status.idle":"2024-03-18T05:33:52.723895Z","shell.execute_reply.started":"2024-03-18T05:33:52.170065Z","shell.execute_reply":"2024-03-18T05:33:52.722591Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"code","source":"def seed_everything(seed_value):\n    random.seed(seed_value) # Python\n    np.random.seed(seed_value) # cpu vars\n    torch.manual_seed(seed_value) # cpu vars    \n    if torch.cuda.is_available(): \n        torch.cuda.manual_seed(seed_value)\n        torch.cuda.manual_seed_all(seed_value) # gpu vars if use multi-GPU\n    if torch.backends.cudnn.is_available:\n        torch.backends.cudnn.deterministic = True\n        torch.backends.cudnn.benchmark = False\n    return print('# SEEDING DONE')\n\ndef kfold(df, fold, n_split, random_state=2023, input_col=None):\n    skf = StratifiedKFold(n_splits=n_split, random_state=random_state, shuffle=True)\n    for idx, (train_index, valid_index) in enumerate(skf.split(df[input_col], df[input_col], df['flag'])):\n        if idx == fold:\n            return train_index, valid_index\n\ndef calculate_mae_and_mrrmse(y_pred_original, y_true):\n    rowwise_rmse = np.sqrt(np.mean(np.square(y_true - y_pred_original), axis=1))\n    mrrmse_score = np.mean(rowwise_rmse)\n    return mrrmse_score","metadata":{"execution":{"iopub.status.busy":"2024-03-18T05:33:52.726809Z","iopub.execute_input":"2024-03-18T05:33:52.727210Z","iopub.status.idle":"2024-03-18T05:33:52.739214Z","shell.execute_reply.started":"2024-03-18T05:33:52.727177Z","shell.execute_reply":"2024-03-18T05:33:52.737483Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CustomModel(nn.Module):\n    def __init__(self, emb_size=512, exp_ratio=4, labels=18211):\n        super().__init__()\n        self.ce_emb = nn.Linear(18211, emb_size)\n        self.sm_emb = nn.Linear(18211, emb_size)\n\n        hidden_size = emb_size * 2\n        self.fc1 = nn.Linear(hidden_size, hidden_size*exp_ratio)\n        self.fc2 = nn.Linear(hidden_size*exp_ratio, hidden_size)\n        \n        self.act = nn.GELU()\n        self.out = nn.Linear(hidden_size, labels)\n \n    def forward(self, x1, x2):\n        x1 = self.ce_emb(x1)\n        x2 = self.sm_emb(x2)\n        \n        x1 = self.act(x1)\n        x2 = self.act(x2)\n\n        x = torch.concat([x1, x2], dim=-1)\n\n        x = self.fc1(x)\n        x = self.act(x)\n        x = self.fc2(x)\n        x = self.act(x)\n\n        x = self.out(x)\n        return x\n\nclass CustomDataset(Dataset):\n    def __init__(self, features, labels, stats_dict):\n        super().__init__()\n        self.x = features\n        self.y = labels\n        self.stats_dict = stats_dict\n        \n    def __len__(self):\n        return len(self.x)\n    \n    def __getitem__(self, idx):\n        x = self.x[idx]\n        x1 = self.stats_dict['cell_type'][x[0]]\n        x2 = self.stats_dict['sm_name'][x[1]]\n        x = {'cell_type': x1, 'sm_name': x2}\n        y = self.y[idx]\n        return x, torch.tensor(y).float()\n\ndef train_epoch(model, optimizer, loss_fn, loader, device, scheduler=None):\n    model.train()\n    \n    losses, metrics = [], []\n    for i, (x, y) in enumerate(tqdm(loader)):\n        x1, x2 = x['cell_type'].to(device), x['sm_name'].to(device)\n        y = y.to(device)\n\n        output = model(x1, x2) \n        loss = loss_fn(output, y)\n\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n\n        if scheduler is not None:\n            scheduler.step()\n\n        losses.append(loss.detach().cpu().item())\n        metrics.append(calculate_mae_and_mrrmse(output.detach().cpu().numpy(), y.detach().cpu().numpy()))\n    return np.mean(losses), np.mean(metrics)\n\ndef validate(model, loss_fn, loader, device):\n    model.eval()\n\n    losses, metrics = [], []\n    with torch.no_grad():\n        for x, y in tqdm(loader):\n            x1, x2 = x['cell_type'].to(device), x['sm_name'].to(device)\n            y = y.to(device)\n\n            output = model(x1, x2) \n            loss = loss_fn(output, y)\n  \n            losses.append(loss.detach().cpu().item())\n            metrics.append(calculate_mae_and_mrrmse(output.detach().cpu().numpy(), y.detach().cpu().numpy()))\n    return np.mean(losses), np.mean(metrics)","metadata":{"execution":{"iopub.status.busy":"2024-03-18T05:33:52.741666Z","iopub.execute_input":"2024-03-18T05:33:52.742051Z","iopub.status.idle":"2024-03-18T05:33:52.767151Z","shell.execute_reply.started":"2024-03-18T05:33:52.742021Z","shell.execute_reply":"2024-03-18T05:33:52.765931Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Config:\n    def __init__(self):\n        self.device = 0\n        self.train_batch_size = 16\n        self.test_batch_size = 16\n        self.workers = 0\n        self.fold = 0\n        self.split = 10\n        self.seed = 2023\n        self.split_seed = 2023\n        self.epoch = 100\n        self.lr = 3e-4\n        self.wd = 0.\n        self.emb_size = 512\n        self.exp_ratio = 4\n        \nargs = Config()\nprint(args.__dict__)","metadata":{"execution":{"iopub.status.busy":"2024-03-18T05:33:52.769024Z","iopub.execute_input":"2024-03-18T05:33:52.769540Z","iopub.status.idle":"2024-03-18T05:33:52.784881Z","shell.execute_reply.started":"2024-03-18T05:33:52.769502Z","shell.execute_reply":"2024-03-18T05:33:52.783989Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"seed_everything(args.seed)\ndevice = torch.device(f'cuda:{args.device}' if torch.cuda.is_available() else 'cpu')\nprint(device)","metadata":{"execution":{"iopub.status.busy":"2024-03-18T05:33:52.787286Z","iopub.execute_input":"2024-03-18T05:33:52.788062Z","iopub.status.idle":"2024-03-18T05:33:52.800585Z","shell.execute_reply.started":"2024-03-18T05:33:52.788000Z","shell.execute_reply":"2024-03-18T05:33:52.799707Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_parquet('../input/open-problems-single-cell-perturbations/de_train.parquet')\nlabels_columns=[\"cell_type\", \"sm_name\", \"sm_lincs_id\", \"SMILES\", \"control\", \"flag\"]\nfeatures_columns = [\"cell_type\", \"sm_name\"]","metadata":{"execution":{"iopub.status.busy":"2024-03-18T05:33:52.802146Z","iopub.execute_input":"2024-03-18T05:33:52.802539Z","iopub.status.idle":"2024-03-18T05:33:54.186158Z","shell.execute_reply.started":"2024-03-18T05:33:52.802507Z","shell.execute_reply":"2024-03-18T05:33:54.184852Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df['flag'] = df['cell_type']\ntrain_index, valid_index = kfold(df=df, fold=args.fold, n_split=args.split,\n                                random_state=args.split_seed, input_col='cell_type')\nx_data, y_data = df[features_columns].values, df.drop(columns=labels_columns).values\nprint(len(train_index), len(valid_index))","metadata":{"execution":{"iopub.status.busy":"2024-03-18T05:33:54.188480Z","iopub.execute_input":"2024-03-18T05:33:54.188992Z","iopub.status.idle":"2024-03-18T05:33:54.252934Z","shell.execute_reply.started":"2024-03-18T05:33:54.188935Z","shell.execute_reply":"2024-03-18T05:33:54.251712Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"x_train, y_train = x_data[train_index], y_data[train_index]\nx_valid, y_valid = x_data[valid_index], y_data[valid_index]\n\nstats_dict = torch.load('/kaggle/working/median_values.pt')\ntrain_set = CustomDataset(features=x_train, labels=y_train, stats_dict=stats_dict)\nvalid_set = CustomDataset(features=x_valid, labels=y_valid, stats_dict=stats_dict)\n\ntrain_loader = DataLoader(train_set, batch_size=args.train_batch_size, drop_last=True, shuffle=True, num_workers=args.workers)\nvalid_loader = DataLoader(valid_set, batch_size=args.test_batch_size, shuffle=False, num_workers=args.workers)\n\nmodel = CustomModel(emb_size=args.emb_size, exp_ratio=args.exp_ratio)\nmodel = model.to(device)\nloss_fn = nn.L1Loss()\n\noptimizer = torch.optim.AdamW(model.parameters(), lr=args.lr, weight_decay=args.wd)","metadata":{"execution":{"iopub.status.busy":"2024-03-18T05:33:54.256554Z","iopub.execute_input":"2024-03-18T05:33:54.256966Z","iopub.status.idle":"2024-03-18T05:33:58.478378Z","shell.execute_reply.started":"2024-03-18T05:33:54.256935Z","shell.execute_reply":"2024-03-18T05:33:58.476506Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"best_score = 100\nbest_epoch = 0\nbest_model = None\n\nfor epoch in range(args.epoch):\n    print(f\"epoch: {epoch}\")\n\n    train_loss, train_metric = train_epoch(model, optimizer, loss_fn, train_loader, device, None)\n    valid_loss, valid_metric = validate(model, loss_fn, valid_loader, device)\n\n    print(f\"train: loss {train_loss:.4f} metric {train_metric:.4f}\")\n    print(f\"valid: loss {valid_loss:.4f} metric {valid_metric:.4f}\")\n\n    if valid_metric < best_score:\n        best_score = valid_metric\n        best_model = deepcopy(model.state_dict())\n        best_epoch = epoch\n        print(f\"best score updated {best_score :.4f}\")\n        \ntorch.save(best_model, f\"{args.fold}fold_best_model.pt\")","metadata":{"execution":{"iopub.status.busy":"2024-03-18T05:33:58.480723Z","iopub.execute_input":"2024-03-18T05:33:58.481479Z","iopub.status.idle":"2024-03-18T05:34:14.296617Z","shell.execute_reply.started":"2024-03-18T05:33:58.481433Z","shell.execute_reply":"2024-03-18T05:34:14.295482Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Inference","metadata":{}},{"cell_type":"code","source":"test_df = pd.read_csv('/kaggle/input/open-problems-single-cell-perturbations/id_map.csv')\nss = pd.read_csv('/kaggle/input/open-problems-single-cell-perturbations/sample_submission.csv')\nfeatures_columns = [\"cell_type\", \"sm_name\"]","metadata":{"execution":{"iopub.status.busy":"2024-03-18T05:34:14.298329Z","iopub.execute_input":"2024-03-18T05:34:14.298819Z","iopub.status.idle":"2024-03-18T05:34:17.479164Z","shell.execute_reply.started":"2024-03-18T05:34:14.298773Z","shell.execute_reply":"2024-03-18T05:34:17.477684Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def inference(paths, loader, device):\n    res = []\n    for p in paths:\n        print(p)\n        model = CustomModel(emb_size=512, exp_ratio=4)\n        model = model.to(device)\n\n        model.load_state_dict(torch.load(p))\n        model.eval()\n        \n        tmp = []\n        with torch.no_grad():\n            for x, y in tqdm(loader):\n                x1, x2 = x['cell_type'].to(device), x['sm_name'].to(device)\n                y = y.to(device)\n                output = model(x1, x2) \n\n                tmp.extend(output.clone().detach().cpu())\n            output_stack = torch.stack(tmp, dim=0)\n            res.append(output_stack)\n    return torch.mean(torch.stack(res, dim=0), dim=0)","metadata":{"execution":{"iopub.status.busy":"2024-03-18T05:34:17.480802Z","iopub.execute_input":"2024-03-18T05:34:17.481264Z","iopub.status.idle":"2024-03-18T05:34:17.491244Z","shell.execute_reply.started":"2024-03-18T05:34:17.481214Z","shell.execute_reply":"2024-03-18T05:34:17.490340Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"x_data = test_df[features_columns].values\ntest_set = CustomDataset(features=x_data, labels=np.zeros(len(x_data)), stats_dict=stats_dict)\ntest_loader = DataLoader(test_set, batch_size=16, drop_last=False, shuffle=False, num_workers=0, pin_memory=False)","metadata":{"execution":{"iopub.status.busy":"2024-03-18T05:34:17.493339Z","iopub.execute_input":"2024-03-18T05:34:17.493944Z","iopub.status.idle":"2024-03-18T05:34:17.508304Z","shell.execute_reply.started":"2024-03-18T05:34:17.493837Z","shell.execute_reply":"2024-03-18T05:34:17.506857Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"paths = glob('/kaggle/working/*model*.pt')\npaths","metadata":{"execution":{"iopub.status.busy":"2024-03-18T05:34:17.510658Z","iopub.execute_input":"2024-03-18T05:34:17.511228Z","iopub.status.idle":"2024-03-18T05:34:17.524871Z","shell.execute_reply.started":"2024-03-18T05:34:17.511185Z","shell.execute_reply":"2024-03-18T05:34:17.523824Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"res = inference(paths, test_loader, device)","metadata":{"execution":{"iopub.status.busy":"2024-03-18T05:34:17.526006Z","iopub.execute_input":"2024-03-18T05:34:17.526424Z","iopub.status.idle":"2024-03-18T05:34:18.809088Z","shell.execute_reply.started":"2024-03-18T05:34:17.526349Z","shell.execute_reply":"2024-03-18T05:34:18.807808Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ss.iloc[:, 1:] = res.numpy()","metadata":{"execution":{"iopub.status.busy":"2024-03-18T05:34:18.810782Z","iopub.execute_input":"2024-03-18T05:34:18.811248Z","iopub.status.idle":"2024-03-18T05:34:19.237300Z","shell.execute_reply.started":"2024-03-18T05:34:18.811206Z","shell.execute_reply":"2024-03-18T05:34:19.236221Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ss","metadata":{"execution":{"iopub.status.busy":"2024-03-18T05:34:19.238532Z","iopub.execute_input":"2024-03-18T05:34:19.238853Z","iopub.status.idle":"2024-03-18T05:34:19.282028Z","shell.execute_reply.started":"2024-03-18T05:34:19.238826Z","shell.execute_reply":"2024-03-18T05:34:19.280563Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"save_name = 'submission'\nss.to_csv(f'{save_name}.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2024-03-18T05:34:19.286108Z","iopub.execute_input":"2024-03-18T05:34:19.286565Z","iopub.status.idle":"2024-03-18T05:34:33.462586Z","shell.execute_reply.started":"2024-03-18T05:34:19.286530Z","shell.execute_reply":"2024-03-18T05:34:33.461357Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}