{"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":8042988,"sourceType":"datasetVersion","datasetId":4740586},{"sourceId":8519891,"sourceType":"datasetVersion","datasetId":4784530}],"dockerImageVersionId":30732,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-06-24T05:20:57.217437Z","iopub.execute_input":"2024-06-24T05:20:57.218062Z","iopub.status.idle":"2024-06-24T05:20:57.856179Z","shell.execute_reply.started":"2024-06-24T05:20:57.217993Z","shell.execute_reply":"2024-06-24T05:20:57.854894Z"},"trusted":true},"execution_count":1,"outputs":[{"name":"stdout","text":"/kaggle/input/belka-shrunken-train-set/train.parquet\n/kaggle/input/belka-shrunken-train-set/test.parquet\n/kaggle/input/belka-shrunken-train-set/train.csv\n/kaggle/input/belka-shrunken-train-set/test.csv\n/kaggle/input/belka-shrunken-train-set/test_dicts/BBs_dict_1_test.p\n/kaggle/input/belka-shrunken-train-set/test_dicts/molecule_smiles_unique.p\n/kaggle/input/belka-shrunken-train-set/test_dicts/BBs_dict_2_test.p\n/kaggle/input/belka-shrunken-train-set/test_dicts/BBs_dict_reverse_2_test.p\n/kaggle/input/belka-shrunken-train-set/test_dicts/BBs_dict_reverse_3_test.p\n/kaggle/input/belka-shrunken-train-set/test_dicts/BBs_dict_reverse_1_test.p\n/kaggle/input/belka-shrunken-train-set/test_dicts/BBs_dict_3_test.p\n/kaggle/input/belka-shrunken-train-set/train_dicts/BBs_dict_reverse_3.p\n/kaggle/input/belka-shrunken-train-set/train_dicts/BBs_dict_reverse_2.p\n/kaggle/input/belka-shrunken-train-set/train_dicts/BBs_dict_2.p\n/kaggle/input/belka-shrunken-train-set/train_dicts/BBs_dict_3.p\n/kaggle/input/belka-shrunken-train-set/train_dicts/BBs_dict_1.p\n/kaggle/input/belka-shrunken-train-set/train_dicts/BBs_dict_reverse_1.p\n/kaggle/input/leash-BELKA/sample_submission.csv\n/kaggle/input/leash-BELKA/train.parquet\n/kaggle/input/leash-BELKA/test.parquet\n/kaggle/input/leash-BELKA/train.csv\n/kaggle/input/leash-BELKA/test.csv\n/kaggle/input/leash-bio-processed-dataset/test.ecfp4.packed.npz\n/kaggle/input/leash-bio-processed-dataset/train-replace-c.smiles.bytestring.bz2\n/kaggle/input/leash-bio-processed-dataset/train-replace-c.sub-train03.conforge.sdf.bz2\n/kaggle/input/leash-bio-processed-dataset/protein.csv\n/kaggle/input/leash-bio-processed-dataset/train-replace-c.sub-train04.conforge.sdf.bz2\n/kaggle/input/leash-bio-processed-dataset/train-replace-c.sub-train05.conforge.status.parquet\n/kaggle/input/leash-bio-processed-dataset/test-replace-c.conforge.status.parquet\n/kaggle/input/leash-bio-processed-dataset/train-replace-c.sub-train02.conforge.sdf.bz2\n/kaggle/input/leash-bio-processed-dataset/train-replace-c.sub-train.conforge.sdf.bz2\n/kaggle/input/leash-bio-processed-dataset/train-replace-c.sub-train01.conforge.status.parquet\n/kaggle/input/leash-bio-processed-dataset/all_buildingblock.csv\n/kaggle/input/leash-bio-processed-dataset/train-replace-c-30m.graph.pickle.b2z\n/kaggle/input/leash-bio-processed-dataset/train-replace-c.sub-valid.conforge.sdf.bz2\n/kaggle/input/leash-bio-processed-dataset/train.reduced.parquet\n/kaggle/input/leash-bio-processed-dataset/train-replace-c.sub-train08.conforge.sdf.bz2\n/kaggle/input/leash-bio-processed-dataset/train-replace-c.sub-train04.conforge.status.parquet\n/kaggle/input/leash-bio-processed-dataset/train-replace-c.sub-train08.conforge.status.parquet\n/kaggle/input/leash-bio-processed-dataset/train-replace-c.sub-train06.conforge.status.parquet\n/kaggle/input/leash-bio-processed-dataset/train-replace-c.sub-nonshare.conforge.sdf.bz2\n/kaggle/input/leash-bio-processed-dataset/fold0.parquet\n/kaggle/input/leash-bio-processed-dataset/train-replace-c.sub-train03.conforge.status.parquet\n/kaggle/input/leash-bio-processed-dataset/train-replace-c.sub-train07.conforge.sdf.bz2\n/kaggle/input/leash-bio-processed-dataset/train-replace-c-30m.graph.pickle.02.b2z\n/kaggle/input/leash-bio-processed-dataset/train.reduced.small100.csv\n/kaggle/input/leash-bio-processed-dataset/train-replace-c.sub-nonshare.conforge.status.parquet\n/kaggle/input/leash-bio-processed-dataset/train-replace-c.sub-train09.conforge.status.parquet\n/kaggle/input/leash-bio-processed-dataset/test-replace-c.smiles.bytestring.bz2\n/kaggle/input/leash-bio-processed-dataset/test.reduced.parquet\n/kaggle/input/leash-bio-processed-dataset/train-replace-c.sub-train.conforge.status.parquet\n/kaggle/input/leash-bio-processed-dataset/train-replace-c.sub-train02.conforge.status.parquet\n/kaggle/input/leash-bio-processed-dataset/train-replace-c.sub-train06.conforge.sdf.bz2\n/kaggle/input/leash-bio-processed-dataset/train-replace-c.sub-valid.conforge.status.parquet\n/kaggle/input/leash-bio-processed-dataset/train.bind.npz\n/kaggle/input/leash-bio-processed-dataset/train.ecfp4.packed.npz\n/kaggle/input/leash-bio-processed-dataset/test.reduced.small100.csv\n/kaggle/input/leash-bio-processed-dataset/train-replace-c.sub-train07.conforge.status.parquet\n/kaggle/input/leash-bio-processed-dataset/train-replace-c.sub-train05.conforge.sdf.bz2\n/kaggle/input/leash-bio-processed-dataset/train-replace-c-30m.graph.pickle.01.b2z\n/kaggle/input/leash-bio-processed-dataset/conformer-3d-gnn-example-code/readme\n/kaggle/input/leash-bio-processed-dataset/conformer-3d-gnn-example-code/log.train.txt\n/kaggle/input/leash-bio-processed-dataset/conformer-3d-gnn-example-code/log.valid.txt\n/kaggle/input/leash-bio-processed-dataset/conformer-3d-gnn-example-code/dig-comeNet-sphereNet-06/configure.py\n/kaggle/input/leash-bio-processed-dataset/conformer-3d-gnn-example-code/dig-comeNet-sphereNet-06/model.py\n/kaggle/input/leash-bio-processed-dataset/conformer-3d-gnn-example-code/dig-comeNet-sphereNet-06/run_train.py\n/kaggle/input/leash-bio-processed-dataset/conformer-3d-gnn-example-code/dig-comeNet-sphereNet-06/run_submit.py\n/kaggle/input/leash-bio-processed-dataset/conformer-3d-gnn-example-code/dig-comeNet-sphereNet-06/dataset.py\n/kaggle/input/leash-bio-processed-dataset/conformer-3d-gnn-example-code/dig-comeNet-sphereNet-06/run_valid.py\n","output_type":"stream"}]},{"cell_type":"code","source":"!pip install torch torch-geometric rdkit duckdb","metadata":{"execution":{"iopub.status.busy":"2024-06-23T08:27:31.957228Z","iopub.execute_input":"2024-06-23T08:27:31.958018Z","iopub.status.idle":"2024-06-23T08:27:51.560846Z","shell.execute_reply.started":"2024-06-23T08:27:31.957971Z","shell.execute_reply":"2024-06-23T08:27:51.559538Z"},"trusted":true},"execution_count":2,"outputs":[{"name":"stdout","text":"Requirement already satisfied: torch in /opt/conda/lib/python3.10/site-packages (2.1.2+cpu)\nCollecting torch-geometric\n  Downloading torch_geometric-2.5.3-py3-none-any.whl.metadata (64 kB)\n\u001b[2K     \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m64.2/64.2 kB\u001b[0m \u001b[31m1.5 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0ma \u001b[36m0:00:01\u001b[0m\n\u001b[?25hCollecting rdkit\n  Downloading rdkit-2023.9.6-cp310-cp310-manylinux_2_17_x86_64.manylinux2014_x86_64.whl.metadata (3.9 kB)\nCollecting duckdb\n  Downloading duckdb-1.0.0-cp310-cp310-manylinux_2_17_x86_64.manylinux2014_x86_64.whl.metadata (762 bytes)\nRequirement already satisfied: filelock in /opt/conda/lib/python3.10/site-packages (from torch) (3.13.1)\nRequirement already satisfied: typing-extensions in /opt/conda/lib/python3.10/site-packages (from torch) (4.9.0)\nRequirement already satisfied: sympy in /opt/conda/lib/python3.10/site-packages (from torch) (1.12.1)\nRequirement already satisfied: networkx in /opt/conda/lib/python3.10/site-packages (from torch) (3.2.1)\nRequirement already satisfied: jinja2 in /opt/conda/lib/python3.10/site-packages (from torch) (3.1.2)\nRequirement already satisfied: fsspec in /opt/conda/lib/python3.10/site-packages (from torch) (2024.3.1)\nRequirement already satisfied: tqdm in /opt/conda/lib/python3.10/site-packages (from torch-geometric) (4.66.4)\nRequirement already satisfied: numpy in /opt/conda/lib/python3.10/site-packages (from torch-geometric) (1.26.4)\nRequirement already satisfied: scipy in /opt/conda/lib/python3.10/site-packages (from torch-geometric) (1.11.4)\nRequirement already satisfied: aiohttp in /opt/conda/lib/python3.10/site-packages (from torch-geometric) (3.9.1)\nRequirement already satisfied: requests in /opt/conda/lib/python3.10/site-packages (from torch-geometric) (2.32.3)\nRequirement already satisfied: pyparsing in /opt/conda/lib/python3.10/site-packages (from torch-geometric) (3.1.1)\nRequirement already satisfied: scikit-learn in /opt/conda/lib/python3.10/site-packages (from torch-geometric) (1.2.2)\nRequirement already satisfied: psutil>=5.8.0 in /opt/conda/lib/python3.10/site-packages (from torch-geometric) (5.9.3)\nRequirement already satisfied: Pillow in /opt/conda/lib/python3.10/site-packages (from rdkit) (9.5.0)\nRequirement already satisfied: attrs>=17.3.0 in /opt/conda/lib/python3.10/site-packages (from aiohttp->torch-geometric) (23.2.0)\nRequirement already satisfied: multidict<7.0,>=4.5 in /opt/conda/lib/python3.10/site-packages (from aiohttp->torch-geometric) (6.0.4)\nRequirement already satisfied: yarl<2.0,>=1.0 in /opt/conda/lib/python3.10/site-packages (from aiohttp->torch-geometric) (1.9.3)\nRequirement already satisfied: frozenlist>=1.1.1 in /opt/conda/lib/python3.10/site-packages (from aiohttp->torch-geometric) (1.4.1)\nRequirement already satisfied: aiosignal>=1.1.2 in /opt/conda/lib/python3.10/site-packages (from aiohttp->torch-geometric) (1.3.1)\nRequirement already satisfied: async-timeout<5.0,>=4.0 in /opt/conda/lib/python3.10/site-packages (from aiohttp->torch-geometric) (4.0.3)\nRequirement already satisfied: MarkupSafe>=2.0 in /opt/conda/lib/python3.10/site-packages (from jinja2->torch) (2.1.3)\nRequirement already satisfied: charset-normalizer<4,>=2 in /opt/conda/lib/python3.10/site-packages (from requests->torch-geometric) (3.3.2)\nRequirement already satisfied: idna<4,>=2.5 in /opt/conda/lib/python3.10/site-packages (from requests->torch-geometric) (3.6)\nRequirement already satisfied: urllib3<3,>=1.21.1 in /opt/conda/lib/python3.10/site-packages (from requests->torch-geometric) (1.26.18)\nRequirement already satisfied: certifi>=2017.4.17 in /opt/conda/lib/python3.10/site-packages (from requests->torch-geometric) (2024.2.2)\nRequirement already satisfied: joblib>=1.1.1 in /opt/conda/lib/python3.10/site-packages (from scikit-learn->torch-geometric) (1.4.2)\nRequirement already satisfied: threadpoolctl>=2.0.0 in /opt/conda/lib/python3.10/site-packages (from scikit-learn->torch-geometric) (3.2.0)\nRequirement already satisfied: mpmath<1.4.0,>=1.1.0 in /opt/conda/lib/python3.10/site-packages (from sympy->torch) (1.3.0)\nDownloading torch_geometric-2.5.3-py3-none-any.whl (1.1 MB)\n\u001b[2K   \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m1.1/1.1 MB\u001b[0m \u001b[31m14.7 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m00:01\u001b[0m00:01\u001b[0m\n\u001b[?25hDownloading rdkit-2023.9.6-cp310-cp310-manylinux_2_17_x86_64.manylinux2014_x86_64.whl (34.9 MB)\n\u001b[2K   \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m34.9/34.9 MB\u001b[0m \u001b[31m32.3 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m00:01\u001b[0m00:01\u001b[0mm\n\u001b[?25hDownloading duckdb-1.0.0-cp310-cp310-manylinux_2_17_x86_64.manylinux2014_x86_64.whl (18.5 MB)\n\u001b[2K   \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m18.5/18.5 MB\u001b[0m \u001b[31m68.0 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m:00:01\u001b[0m00:01\u001b[0m\n\u001b[?25hInstalling collected packages: rdkit, duckdb, torch-geometric\nSuccessfully installed duckdb-1.0.0 rdkit-2023.9.6 torch-geometric-2.5.3\n","output_type":"stream"}]},{"cell_type":"code","source":"import _pickle as  cPickle\nimport bz2\ndef save_compressed_pickle(file, data):\n    with bz2.BZ2File(file , 'w') as f:\n        cPickle.dump(data, f)\n\ndef load_decompress_pickle(file_path, start, end):\n    with bz2.BZ2File(file_path, 'rb') as f:\n        data = cPickle.load(f)\n        # データを[start:end]の範囲で部分的に取得\n        return data[start:end]","metadata":{"execution":{"iopub.status.busy":"2024-06-24T05:36:53.275625Z","iopub.execute_input":"2024-06-24T05:36:53.276176Z","iopub.status.idle":"2024-06-24T05:36:53.284871Z","shell.execute_reply.started":"2024-06-24T05:36:53.276139Z","shell.execute_reply":"2024-06-24T05:36:53.283387Z"},"trusted":true},"execution_count":2,"outputs":[]},{"cell_type":"code","source":"%%time\ntrain1 = load_decompress_pickle('/kaggle/input/leash-bio-processed-dataset/train-replace-c-30m.graph.pickle.01.b2z', 0, 1)","metadata":{"execution":{"iopub.status.busy":"2024-06-24T05:39:54.740588Z","iopub.execute_input":"2024-06-24T05:39:54.741047Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"メモリクラッシュ。。不憫すぎる。。","metadata":{}},{"cell_type":"code","source":"import torch\nfrom torch.nn import Linear\nimport torch.nn.functional as F\nfrom torch_geometric.nn import TransformerConv, global_mean_pool\nfrom torch_geometric.data import Data\nfrom torch_geometric.loader import DataLoader\nfrom rdkit import Chem\nfrom rdkit.Chem import rdmolops\nimport networkx as nx\nimport duckdb\nimport gc\n\nsamples_per_category = 500000\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Using device: {device}\")","metadata":{"execution":{"iopub.status.busy":"2024-06-23T08:27:51.56259Z","iopub.execute_input":"2024-06-23T08:27:51.562962Z","iopub.status.idle":"2024-06-23T08:27:58.606624Z","shell.execute_reply.started":"2024-06-23T08:27:51.562928Z","shell.execute_reply":"2024-06-23T08:27:58.60503Z"},"trusted":true},"execution_count":3,"outputs":[{"name":"stdout","text":"Using device: cpu\n","output_type":"stream"}]},{"cell_type":"code","source":"def get_balanced_data_for_protein(file_path, protein, samples):\n    query = f\"\"\"(SELECT *\n                    FROM parquet_scan('{train_path}')\n                    WHERE binds = 0 AND protein_name = '{protein}'\n                    ORDER BY random()\n                    LIMIT {samples})\n                    UNION ALL\n                    (SELECT *\n                    FROM parquet_scan('{train_path}')\n                    WHERE binds = 1 AND protein_name = '{protein}'\n                    ORDER BY random()\n                    LIMIT {samples})\"\"\"\n\n    return con.query(query).df()\n\ndef get_data_for_protein(file_path, protein):\n    query = f\"\"\"(SELECT * \n                    FROM parquet_scan('{file_path}')\n                    WHERE protein_name = '{protein}')\"\"\"\n    return con.query(query).df()","metadata":{"execution":{"iopub.status.busy":"2024-06-16T16:32:37.788924Z","iopub.execute_input":"2024-06-16T16:32:37.789588Z","iopub.status.idle":"2024-06-16T16:32:37.795313Z","shell.execute_reply.started":"2024-06-16T16:32:37.789558Z","shell.execute_reply":"2024-06-16T16:32:37.794464Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\ntrain_path = '/kaggle/input/belka-shrunken-train-set/train.parquet'\ntest_path = '/kaggle/input/belka-shrunken-train-set/test.parquet'\ncon = duckdb.connect()  \n\nproteins = ['sEH', 'BRD4', 'HSA']\n\ndatasets = {}\nfor protein in proteins:\n    datasets[protein] = get_balanced_data_for_protein(train_path, protein, samples_per_category)\nseh_df = datasets['sEH']\nbrd4_df = datasets['BRD4']\nhsa_df = datasets['HSA']\n\nprint(seh_df.head())\n\ntest_datasets = {}\nfor protein in proteins:  \n    test_datasets[protein] = get_data_for_protein(test_path, protein)\nseh_test_df = test_datasets['sEH']\nbrd4_test_df = test_datasets['BRD4']\nhsa_test_df = test_datasets['HSA']\n\nprint(seh_test_df.head())\n\ndel datasets, test_datasets\ngc.collect()\n    \ncon.close()","metadata":{"execution":{"iopub.status.busy":"2024-06-16T16:32:37.796444Z","iopub.execute_input":"2024-06-16T16:32:37.796792Z","iopub.status.idle":"2024-06-16T16:32:37.825155Z","shell.execute_reply.started":"2024-06-16T16:32:37.79675Z","shell.execute_reply":"2024-06-16T16:32:37.824314Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Creating Graph dataset","metadata":{"execution":{"iopub.status.busy":"2024-06-11T06:18:59.363446Z","iopub.execute_input":"2024-06-11T06:18:59.364981Z","iopub.status.idle":"2024-06-11T06:18:59.376012Z","shell.execute_reply.started":"2024-06-11T06:18:59.36493Z","shell.execute_reply":"2024-06-11T06:18:59.374766Z"}}},{"cell_type":"code","source":"# Atom Featurisation\n## Auxiliary function for one-hot enconding transformation based on list of\n##permitted values\n\ndef one_hot_encoding(x, permitted_list):\n    \"\"\"\n    Maps input elements x which are not in the permitted list to the last element\n    of the permitted list.\n    \"\"\"\n    if x not in permitted_list:\n        x = permitted_list[-1]\n    binary_encoding = [int(boolean_value) for boolean_value in list(map(lambda s: x == s, permitted_list))]\n    return binary_encoding\n    \n    \n# Main atom feat. func\n\ndef get_atom_features(atom, use_chirality=True):\n    # Define a simplified list of atom types\n    permitted_atom_types = ['C', 'N', 'O', 'S', 'P', 'F', 'Cl', 'Br', 'I','Dy', 'Unknown']\n    atom_type = atom.GetSymbol() if atom.GetSymbol() in permitted_atom_types else 'Unknown'\n    atom_type_enc = one_hot_encoding(atom_type, permitted_atom_types)\n    \n    # Consider only the most impactful features: atom degree and whether the atom is in a ring\n    atom_degree = one_hot_encoding(atom.GetDegree(), [0, 1, 2, 3, 4, 'MoreThanFour'])\n    is_in_ring = [int(atom.IsInRing())]\n    \n    # Optionally include chirality\n    if use_chirality:\n        chirality_enc = one_hot_encoding(str(atom.GetChiralTag()), [\"CHI_UNSPECIFIED\", \"CHI_TETRAHEDRAL_CW\", \"CHI_TETRAHEDRAL_CCW\", \"CHI_OTHER\"])\n        atom_features = atom_type_enc + atom_degree + is_in_ring + chirality_enc\n    else:\n        atom_features = atom_type_enc + atom_degree + is_in_ring\n    \n    return np.array(atom_features, dtype=np.float32)\n\n# Bond featurization\n\ndef get_bond_features(bond):\n    # Simplified list of bond types\n    permitted_bond_types = [Chem.rdchem.BondType.SINGLE, Chem.rdchem.BondType.DOUBLE, Chem.rdchem.BondType.TRIPLE, Chem.rdchem.BondType.AROMATIC, 'Unknown']\n    bond_type = bond.GetBondType() if bond.GetBondType() in permitted_bond_types else 'Unknown'\n    \n    # Features: Bond type, Is in a ring\n    features = one_hot_encoding(bond_type, permitted_bond_types) \\\n               + [int(bond.IsInRing())]\n    \n    return np.array(features, dtype=np.float32)","metadata":{"execution":{"iopub.status.busy":"2024-06-16T16:32:37.82672Z","iopub.execute_input":"2024-06-16T16:32:37.827046Z","iopub.status.idle":"2024-06-16T16:32:37.839289Z","shell.execute_reply.started":"2024-06-16T16:32:37.827017Z","shell.execute_reply":"2024-06-16T16:32:37.838466Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def create_pytorch_geometric_graph_data_list_from_smiles_and_labels(x_smiles, ids, y=None):\n    data_list = []\n    \n    for index, smiles in enumerate(x_smiles):\n        mol = Chem.MolFromSmiles(smiles)\n        \n        if not mol:  # Skip invalid SMILES strings\n            continue\n        \n        # Node features\n        atom_features = [get_atom_features(atom) for atom in mol.GetAtoms()]\n        atom_features = np.array(atom_features)\n        x = torch.tensor(atom_features, dtype=torch.float)\n        \n        # Edge features\n        edge_index = []\n        edge_features = []\n        for bond in mol.GetBonds():\n            start, end = bond.GetBeginAtomIdx(), bond.GetEndAtomIdx()\n            edge_index += [(start, end), (end, start)]  # Undirected graph\n            bond_feature = get_bond_features(bond)\n            edge_features += [bond_feature, bond_feature]  # Same features in both directions\n        \n        edge_index = np.array(edge_index)\n        edge_index = torch.tensor(edge_index, dtype=torch.long).t().contiguous()\n        edge_features = np.array(edge_features)\n        edge_attr = torch.tensor(edge_features, dtype=torch.float)\n        \n        # Creating the Data object\n        data = Data(x=x, edge_index=edge_index, edge_attr=edge_attr)\n        data.molecule_id = ids[index]\n        if y is not None:\n            data.y = torch.tensor([y[index]], dtype=torch.float)\n        \n        data_list.append(data)\n    \n    return data_list","metadata":{"execution":{"iopub.status.busy":"2024-06-16T16:32:37.840406Z","iopub.execute_input":"2024-06-16T16:32:37.840732Z","iopub.status.idle":"2024-06-16T16:32:37.853947Z","shell.execute_reply.started":"2024-06-16T16:32:37.840703Z","shell.execute_reply":"2024-06-16T16:32:37.853237Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tqdm import tqdm\ndef featurize_data_in_batches(smiles_list, labels_list, batch_size):\n    data_list = []\n    # Define tqdm progress bar\n    pbar = tqdm(total=len(smiles_list), desc=\"Featurizing data\")\n    for i in range(0, len(smiles_list), batch_size):\n        smiles_batch = smiles_list[i:i+batch_size]\n        labels_batch = labels_list[i:i+batch_size]\n        ids_batch = ids_list[i:i+batch_size]\n        batch_data_list = create_pytorch_geometric_graph_data_list_from_smiles_and_labels(smiles_batch, ids_batch, labels_batch)\n        data_list.extend(batch_data_list)\n        pbar.update(len(smiles_batch))\n        \n    pbar.close()\n    return data_list","metadata":{"execution":{"iopub.status.busy":"2024-06-16T16:32:37.85503Z","iopub.execute_input":"2024-06-16T16:32:37.855307Z","iopub.status.idle":"2024-06-16T16:32:37.867222Z","shell.execute_reply.started":"2024-06-16T16:32:37.855284Z","shell.execute_reply":"2024-06-16T16:32:37.866414Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Define the batch size for featurization\nbatch_size = 2**8\n# List of proteins and their corresponding dataframes\nproteins_data = {\n    'sEH': seh_df,\n    'BRD4': brd4_df,\n    'HSA': hsa_df\n}\n# Dictionary to store the featurized data for each protein\nfeaturized_data = {}\n# Loop over each protein and its dataframe\nfor protein_name, df in proteins_data.items():\n    if protein_name == 'sEH':\n        print(f\"Processing {protein_name}...\")\n        smiles_list = df['molecule_smiles'].tolist()\n        ids_list = df['id'].tolist()\n        labels_list = df['binds'].tolist()\n        # Featurize the data\n        featurized_data[protein_name] = featurize_data_in_batches(smiles_list, labels_list, batch_size)\n\nseh_train_data = featurized_data['sEH']\nbrd4_train_data = featurized_data['BRD4']\nhsa_train_data = featurized_data['HSA']\nprint(seh_train_data[:5])\n\ndel seh_df, brd4_df, hsa_df\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2024-06-16T16:32:37.868444Z","iopub.execute_input":"2024-06-16T16:32:37.868721Z","iopub.status.idle":"2024-06-16T16:32:37.878811Z","shell.execute_reply.started":"2024-06-16T16:32:37.868698Z","shell.execute_reply":"2024-06-16T16:32:37.877997Z"},"trusted":true},"execution_count":null,"outputs":[]}]}