{"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":67356,"databundleVersionId":8006601,"sourceType":"competition"},{"sourceId":8073018,"sourceType":"datasetVersion","datasetId":4763782}],"dockerImageVersionId":30698,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# BELKA: GAT baseline [Train/Inference]\n\n## Introduction\n\nGNN is a very powerful method to predict molecular properties.  \nHere, GAT (Graph Attention Network) is trained to predict the binds in BELKA.\n\nNote hyperparameters in this notebook were not optimized.","metadata":{}},{"cell_type":"code","source":"# CPU\n! pip install  dgl -f https://data.dgl.ai/wheels/repo.html -q\n! pip install  dglgo -f https://data.dgl.ai/wheels-test/repo.html -q\n! pip install pip install dgllife -q\n! pip install duckdb -q","metadata":{"execution":{"iopub.status.busy":"2024-04-27T11:18:52.396030Z","iopub.execute_input":"2024-04-27T11:18:52.396491Z","iopub.status.idle":"2024-04-27T11:20:12.293144Z","shell.execute_reply.started":"2024-04-27T11:18:52.396432Z","shell.execute_reply":"2024-04-27T11:20:12.291657Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Preparation","metadata":{}},{"cell_type":"code","source":"import duckdb\nimport pandas as pd\nfrom sklearn.metrics import average_precision_score\n\nn_train = 3 * 10 ** 5 # the number of train data to sample (max.: 98415610)\ntask_names=['binds_BRD4', 'binds_HSA', 'binds_sEH']\nn_tasks = len(task_names)","metadata":{"execution":{"iopub.status.busy":"2024-04-27T11:20:12.296161Z","iopub.execute_input":"2024-04-27T11:20:12.296673Z","iopub.status.idle":"2024-04-27T11:20:13.484344Z","shell.execute_reply.started":"2024-04-27T11:20:12.296614Z","shell.execute_reply":"2024-04-27T11:20:13.483301Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# calculate the number of extracted data\n# around 0.5 % are positive\n# let's say, the max. number of positive data to extract is 1/3 of the total number of training data \nn_max_positive = int(98415610 * 0.008)\nn_max_extract = int(n_train / 3)\n\nn_train_bind = min(n_max_positive, n_max_extract)\nn_train_nonbind = max(n_train - n_train_bind * 3, 1)","metadata":{"execution":{"iopub.status.busy":"2024-04-27T11:20:13.485720Z","iopub.execute_input":"2024-04-27T11:20:13.486180Z","iopub.status.idle":"2024-04-27T11:20:13.493523Z","shell.execute_reply.started":"2024-04-27T11:20:13.486152Z","shell.execute_reply":"2024-04-27T11:20:13.492524Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Train dataset loading","metadata":{}},{"cell_type":"code","source":"train_path = '/kaggle/input/processed-shrunken-dataset/train.parquet'\n\ncon = duckdb.connect()\n\ndf = con.query(f\"\"\"\n(SELECT * FROM parquet_scan('{train_path}')\n WHERE {task_names[0]} = 1\n ORDER BY random()\n LIMIT {n_train_bind})\nUNION ALL\n(SELECT * FROM parquet_scan('{train_path}')\n WHERE {task_names[1]} = 1\n ORDER BY random()\n LIMIT {n_train_bind})\nUNION ALL\n(SELECT * FROM parquet_scan('{train_path}')\n WHERE {task_names[2]} = 1\n ORDER BY random()\n LIMIT {n_train_bind})\nUNION ALL\n(SELECT * FROM parquet_scan('{train_path}')\n ORDER BY random()\n LIMIT {n_train_nonbind})\n\"\"\").df()\n\ncon.close()","metadata":{"execution":{"iopub.status.busy":"2024-04-27T11:20:13.495775Z","iopub.execute_input":"2024-04-27T11:20:13.496100Z","iopub.status.idle":"2024-04-27T11:20:37.582313Z","shell.execute_reply.started":"2024-04-27T11:20:13.496073Z","shell.execute_reply":"2024-04-27T11:20:37.581266Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # check training data\nn_pos = df.sum()\nn_all =len(df)\npos_ratio = [n_pos[c]/n_all for c in task_names]\npos_ratio","metadata":{"execution":{"iopub.status.busy":"2024-04-27T11:20:37.583709Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Train model","metadata":{}},{"cell_type":"code","source":"import os\nos.environ[\"DGLBACKEND\"] = \"pytorch\"\nfrom pathlib import Path\nimport shutil\nfrom dgllife.utils import CanonicalAtomFeaturizer, SMILESToBigraph, ScaffoldSplitter, Meter, EarlyStopping\nfrom dgllife.data import MoleculeCSVDataset\nimport json\nfrom torch.utils.data import DataLoader\nimport dgl\nimport torch\nfrom dgllife.model import GATPredictor\nimport torch.nn as nn\nfrom torch.optim import Adam\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport warnings\n\nwarnings.filterwarnings(\"ignore\", category=RuntimeWarning)\nwarnings.filterwarnings(\"ignore\", category=FutureWarning)\nrefresh_data = True # True to run from scratch\nreuse_featurized_graph = True\nreuse_model = True\ninference = True\nif refresh_data:\n    reuse_featurized_graph = False\n    reuse_model = False\ndevice = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\nresult_path = \"classification_results\"\nnode_featurizer = CanonicalAtomFeaturizer()\nedge_featurizer = None\nnum_workers = 1\nsplit_ratio = \"0.8,0.1,0.1\"\ntrain_ratio, val_ratio, test_ratio = map(float, split_ratio.split(','))\nexp_config = json.loads('''\n{\n  \"alpha\": 0.02,\n  \"batch_size\": 512,\n  \"dropout\": 0.0085,\n  \"gnn_hidden_feats\": 64,\n  \"lr\": 0.144,\n  \"num_gnn_layers\": 1,\n  \"num_heads\": 8,\n  \"patience\": 100,\n  \"predictor_hidden_feats\": 64,\n  \"residual\": false,\n  \"weight_decay\": 2.54e-06\n}\n''')\n\ndef collate_molgraphs_masked(data):\n    mapped_data = map(list, zip(*data))\n    smiles, graphs, labels, masks = mapped_data\n\n    bg = dgl.batch(graphs)\n    bg.set_n_initializer(dgl.init.zero_initializer)\n    bg.set_e_initializer(dgl.init.zero_initializer)\n    labels = torch.stack(labels, dim=0)\n\n    if masks is None:\n        masks = torch.ones(labels.shape)\n    else:\n        masks = torch.stack(masks, dim=0)\n\n    return smiles, bg, labels, masks\n\nin_node_feats = node_featurizer.feat_size()\nmetric = \"roc_auc_score\"\nnum_epochs = 50\nprint_every = 10 ** 4\n\ndef predict(model, bg, device):\n    bg = bg.to(device)\n    node_feats = bg.ndata.pop('h').to(device)\n    return model(bg, node_feats)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"setup directory","metadata":{}},{"cell_type":"code","source":"def setup_path_to_results(result_path):\n    Path(result_path).mkdir(exist_ok=True)\n    trial_id = 0\n    path_exists = True\n    path_to_results_next = None\n    while path_exists:\n        path_to_results = path_to_results_next\n        trial_id += 1\n        path_to_results_next = Path(result_path) / str(trial_id)\n        path_exists = path_to_results_next.exists()\n    return path_to_results_next, path_to_results\n\n\nif refresh_data:\n    current_dir = Path.cwd()\n    for item in current_dir.iterdir():\n        if item.is_file():\n            item.unlink()\n        elif item.is_dir():\n            shutil.rmtree(item)\n\n    trial_path, _ = setup_path_to_results(result_path)\n    trial_path.mkdir(exist_ok=True)\nelse:\n    _, trial_path = setup_path_to_results(result_path)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"featurize","metadata":{}},{"cell_type":"code","source":"%%time\nsmiles_to_g = SMILESToBigraph(\n    add_self_loop=True,\n    node_featurizer=node_featurizer,\n    edge_featurizer=edge_featurizer\n)\nif not reuse_featurized_graph:\n    dataset = MoleculeCSVDataset(df=df,\n                                 smiles_to_graph=smiles_to_g,\n                                 smiles_column=\"molecule_smiles\",\n                                 cache_file_path=result_path + \"/graph.bin\",\n                                 task_names=task_names,\n                                 n_jobs=num_workers\n                                )\nelse:\n    dataset = MoleculeCSVDataset(df=df,\n                                 smiles_to_graph=smiles_to_g,\n                                 smiles_column=\"molecule_smiles\",\n                                 cache_file_path=result_path + \"/graph.bin\",\n                                 task_names=task_names,\n                                 n_jobs=num_workers,\n                                 load=True\n                                )","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del df","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"split dataset","metadata":{}},{"cell_type":"code","source":"%%time\ntrain_set, val_set, test_set = ScaffoldSplitter.train_val_test_split(\n      dataset, frac_train=train_ratio, frac_val=val_ratio, frac_test=test_ratio,\n      scaffold_func='smiles'\n)\n\ntrain_loader = DataLoader(dataset=train_set, batch_size=exp_config['batch_size'], shuffle=True,\n                          collate_fn=collate_molgraphs_masked, num_workers=num_workers)\nval_loader = DataLoader(dataset=val_set, batch_size=exp_config['batch_size'],\n                        collate_fn=collate_molgraphs_masked, num_workers=num_workers)\ntest_loader = DataLoader(dataset=test_set, batch_size=exp_config['batch_size'],\n                          collate_fn=collate_molgraphs_masked, num_workers=num_workers)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"define predictor, loss criterion, optimizer...","metadata":{}},{"cell_type":"code","source":"model = GATPredictor(\n    in_feats=in_node_feats,\n    hidden_feats=[exp_config['gnn_hidden_feats']] * exp_config['num_gnn_layers'],\n    num_heads=[exp_config['num_heads']] * exp_config['num_gnn_layers'],\n    feat_drops=[exp_config['dropout']] * exp_config['num_gnn_layers'],\n    attn_drops=[exp_config['dropout']] * exp_config['num_gnn_layers'],\n    alphas=[exp_config['alpha']] * exp_config['num_gnn_layers'],\n    residuals=[exp_config['residual']] * exp_config['num_gnn_layers'],\n    predictor_hidden_feats=exp_config['predictor_hidden_feats'],\n    predictor_dropout=exp_config['dropout'],\n    n_tasks=n_tasks\n).to(device)\n# TODO: use weight option\npos_weight = torch.tensor([(1-r)/r for r in pos_ratio])\nloss_criterion = nn.BCEWithLogitsLoss(reduction='none', pos_weight=pos_weight)\noptimizer = Adam(model.parameters(), lr=exp_config['lr'],\n                  weight_decay=exp_config['weight_decay'])\nstopper = EarlyStopping(patience=exp_config['patience'],\n                        filename=trial_path / 'model.pth',\n                        metric=metric)\nif reuse_model:\n    model.load_state_dict(torch.load(trial_path / 'model.pth')[\"model_state_dict\"])","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"train!","metadata":{}},{"cell_type":"code","source":"%%time\nlosses = []\ntrain_scores = []\nval_scores = []\ntrain_scores_each = []\nval_scores_each = []\nfor epoch in range(num_epochs):\n    # Train\n    model.train()\n    train_meter = Meter()\n    for batch_id, batch_data in enumerate(train_loader):\n        smiles, bg, labels, masks = batch_data\n        if len(smiles) == 1:\n            # Avoid potential issues with batch normalization\n            continue\n\n        labels, masks = labels.to(device), masks.to(device)\n        logits = predict(model, bg, device)\n        # Mask non-existing labels\n        loss = (loss_criterion(logits, labels) * (masks != 0).float()).mean()\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n        train_meter.update(logits, labels, masks)\n        if batch_id % print_every == 0:\n            losses.append(loss.item())\n            print('epoch {:d}/{:d}, batch {:d}/{:d}, loss {:.4f}'.format(\n                epoch + 1, num_epochs, batch_id + 1, len(train_loader), loss.item()))\n    train_score_each = train_meter.compute_metric(metric)\n    train_scores_each.append(train_score_each)\n    train_score = np.mean(train_score_each)\n    train_scores.append(train_score)\n    print('epoch {:d}/{:d}, training {} {:.4f}'.format(\n        epoch + 1, num_epochs, metric, train_score))\n\n    # Validation and early stop\n    model.eval()\n    eval_meter = Meter()\n    with torch.no_grad():\n        for batch_id, batch_data in enumerate(val_loader):\n            smiles, bg, labels, masks = batch_data\n            labels = labels.to(device)\n            bg = bg.to(device)\n            prediction = predict(model, bg, device)\n            eval_meter.update(prediction, labels, masks)\n    val_score_each = eval_meter.compute_metric(metric)\n    val_scores_each.append(val_score_each)\n    val_score = np.mean(val_score_each)\n    val_scores.append(val_score)\n\n    early_stop = stopper.step(val_score, model)\n    print('epoch {:d}/{:d}, validation {} {:.4f}, best validation {} {:.4f}'.format(\n        epoch + 1, num_epochs, metric, val_score,\n        metric, stopper.best_score))\n\n    if early_stop:\n        break","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.plot(losses, marker=\"o\", label=\"loss\")\nplt.legend()\nplt.show()\n\nplt.plot(train_scores, marker=\"o\", label=\"train (all)\")\nplt.plot(val_scores, marker=\"o\", label=\"val (all)\")\nplt.xlabel(\"epoch\")\nplt.ylabel(metric)\nplt.legend()\nplt.show()\n\nfor t, ts, vs in zip(task_names, np.array(train_scores_each).T, np.array(val_scores_each).T):\n    plt.plot(ts, marker=\"o\", label=f\"train {t}\")\n    plt.plot(vs, marker=\"o\", label=f\"val {t}\")\n    plt.xlabel(\"epoch\")\n    plt.ylabel(metric)\n    plt.legend()\n    plt.show()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"stopper.load_checkpoint(model)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"evaluate!","metadata":{}},{"cell_type":"code","source":"%%time\nmodel.eval()\neval_meter = Meter()\nwith torch.no_grad():\n    for batch_id, batch_data in enumerate(test_loader):\n        smiles, bg, labels, masks = batch_data\n        labels = labels.to(device)\n        prediction = predict(model, bg, device)\n        eval_meter.update(prediction, labels, masks)\ntest_score = np.mean(eval_meter.compute_metric(metric))\nprint('test {} {:.4f}'.format(metric, test_score))","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Test Dataset","metadata":{}},{"cell_type":"code","source":"%%time\nif inference:\n    import os\n\n    # Process the test.parquet file chunk by chunk\n    test_smiles_file = '/kaggle/input/processed-shrunken-dataset/test_smiles.csv'\n    test_is_protein_file = '/kaggle/input/processed-shrunken-dataset/test_is_protein.csv'\n    test_id_file = '/kaggle/input/processed-shrunken-dataset/test_id.csv'\n    output_temp_file = 'submission_temp.csv'\n    output_file = 'submission.csv'  # Specify the path and filename for the output file\n    chunksize = exp_config['batch_size']\n\n    pd.DataFrame(columns=task_names).to_csv(output_temp_file, index=False, mode='w')\n\n    # Read the test.parquet file into a pandas DataFrame\n    for df_test_smiles, df_test_is_protein in zip(\n        pd.read_csv(test_smiles_file, chunksize=chunksize),\n        pd.read_csv(test_is_protein_file, chunksize=chunksize)\n    ):\n    #     # only for test\n    #     prediction_df = pd.DataFrame(np.ones((len(df_test_smiles), 3)), columns=task_names)\n    #     prediction_df.to_csv(output_temp_file, mode='a', header=False, index=False)\n    #     continue\n        dataset_test = MoleculeCSVDataset(df=df_test_smiles,\n                                     smiles_to_graph=smiles_to_g,\n                                     smiles_column=\"molecule_smiles\",\n                                     cache_file_path=result_path + \"/graph_test.bin\",\n                                     n_jobs=num_workers\n                                    )\n        actualtest_loader = DataLoader(dataset=dataset_test, batch_size=exp_config['batch_size'],\n                                  collate_fn=collate_molgraphs_masked, num_workers=num_workers)\n\n        prediction_df = pd.DataFrame(columns=[\"binds\"])\n        model.eval()\n        with torch.no_grad():\n            for batch_id, batch_data in enumerate(actualtest_loader):\n                smiles, bg, _, _ = batch_data\n                prediction = predict(model, bg, device)\n                prediction_logits = torch.sigmoid(prediction)\n                prediction_arr = prediction_logits.numpy()\n                prediction_df = pd.DataFrame(prediction_arr, columns=task_names)\n\n                prediction_df.to_csv(output_temp_file, mode='a', header=False, index=False)\n","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\nif inference:\n    prediction_arr_all = pd.read_csv(output_temp_file).to_numpy().reshape(-1)\n    arr_is_protein = pd.read_csv(test_is_protein_file).to_numpy().reshape(-1)\n\n    # # check they have the same length\n    # print(len(prediction_arr_all) == len(arr_is_protein))\n\n    df_test_id = pd.read_csv(test_id_file, dtype={'id': np.int64})\n    df_test_id[\"binds\"] = prediction_arr_all[arr_is_protein]\n    # # df_test_id = pd.read_csv(test_id_file, dtype={'id': np.int64}).loc[:exp_config['batch_size']*3+1,:]\n    df_test_id.to_csv(output_file, index=False)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}