{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":67356,"databundleVersionId":8006601,"sourceType":"competition"},{"sourceId":10916495,"sourceType":"datasetVersion","datasetId":6786369}],"dockerImageVersionId":30918,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"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","trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:19:31.283294Z","iopub.execute_input":"2025-03-05T09:19:31.283653Z","iopub.status.idle":"2025-03-05T09:19:32.219822Z","shell.execute_reply.started":"2025-03-05T09:19:31.283604Z","shell.execute_reply":"2025-03-05T09:19:32.218878Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 1. 安裝 PyTorch\n!pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu121\n\n# 2. 安裝 DGL（選擇適合的 CUDA 版本）\n!pip install dgl -f https://data.dgl.ai/wheels/cu121.html\n\n# 3. 安裝 RDKit（用於處理分子）\n!pip install rdkit-pypi\n\n# 4. 下載 MAT 的原始碼\n!git clone https://github.com/ardigen/MAT.git /kaggle/working/MAT\n#%cd MAT\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:19:32.221099Z","iopub.execute_input":"2025-03-05T09:19:32.221571Z","iopub.status.idle":"2025-03-05T09:19:47.318861Z","shell.execute_reply.started":"2025-03-05T09:19:32.221546Z","shell.execute_reply":"2025-03-05T09:19:47.31804Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile requirements.txt\neasydict\nfuture\nmatplotlib\nnumpy\nopencv-python\nscikit-image\nscipy\nclick\nrequests\ntqdm\npyspng\nninja\nimageio-ffmpeg==0.4.3\ntimm\npsutil\nscikit-learn","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:19:47.320476Z","iopub.execute_input":"2025-03-05T09:19:47.320728Z","iopub.status.idle":"2025-03-05T09:19:47.326566Z","shell.execute_reply.started":"2025-03-05T09:19:47.320708Z","shell.execute_reply":"2025-03-05T09:19:47.325845Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 5. 安裝 MAT 依賴\n!pip install -r requirements.txt\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:19:47.32771Z","iopub.execute_input":"2025-03-05T09:19:47.327979Z","iopub.status.idle":"2025-03-05T09:19:52.666566Z","shell.execute_reply.started":"2025-03-05T09:19:47.327952Z","shell.execute_reply":"2025-03-05T09:19:52.665719Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nimport pandas as pd\n\ndata = {\n    \"smiles\": [\"CCO\", \"C1=CC=CC=C1\", \"CCN(CC)C(=O)C1=CC=CC=C1\"],\n    \"BRD4\": [1, 0, 1],  # 這是可選的標籤數據\n    \"HSA\": [0, 0, 1],\n    \"sEH\": [1, 1, 0],\n}\ndf = pd.DataFrame(data)\n# 存成 CSV\ndf.to_csv(\"/kaggle/working/example_data.csv\", index=False)\n\nprint(\"example_data.csv 已成功建立！\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:19:52.667624Z","iopub.execute_input":"2025-03-05T09:19:52.667926Z","iopub.status.idle":"2025-03-05T09:19:52.688767Z","shell.execute_reply.started":"2025-03-05T09:19:52.667901Z","shell.execute_reply":"2025-03-05T09:19:52.688055Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip uninstall -y torch torchvision torchaudio fastai pylibcugraph-cu12 pylibraft-cu12 rmm-cu12 dgl torchdata\n!pip uninstall -y torch torchvision torchaudio torchdata triton pytorch-lightning\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:19:52.689489Z","iopub.execute_input":"2025-03-05T09:19:52.689726Z","iopub.status.idle":"2025-03-05T09:20:32.371184Z","shell.execute_reply.started":"2025-03-05T09:19:52.689707Z","shell.execute_reply":"2025-03-05T09:20:32.370331Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 1. 升級 PyTorch 和 torchvision\n#!pip install torch==2.1.0 torchvision==0.15.2 torchaudio==2.0.2 torchdata==0.6.1\n!pip install torch torchvision torchaudio torchdata --index-url https://download.pytorch.org/whl/cu118\n!pip install pytorch-lightning\n\n\nimport torch\nprint(torch.__version__)  # 目標是 2.1.0\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:20:32.372182Z","iopub.execute_input":"2025-03-05T09:20:32.372412Z","iopub.status.idle":"2025-03-05T09:22:54.686763Z","shell.execute_reply.started":"2025-03-05T09:20:32.372392Z","shell.execute_reply":"2025-03-05T09:22:54.685968Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nprint(torch.__version__)  # 目標是 2.1.0\n!nvcc --version  # 如果已經安裝了CUDA\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:22:54.687666Z","iopub.execute_input":"2025-03-05T09:22:54.688103Z","iopub.status.idle":"2025-03-05T09:22:54.833484Z","shell.execute_reply.started":"2025-03-05T09:22:54.688073Z","shell.execute_reply":"2025-03-05T09:22:54.832717Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!python --version","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:22:54.8361Z","iopub.execute_input":"2025-03-05T09:22:54.836325Z","iopub.status.idle":"2025-03-05T09:22:54.966169Z","shell.execute_reply.started":"2025-03-05T09:22:54.836303Z","shell.execute_reply":"2025-03-05T09:22:54.965399Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 3. 檢查安裝是否成功\nimport torch\nimport torchdata\n\n\nprint(\"Torch version:\", torch.__version__)\nprint(\"Torchdata version:\", torchdata.__version__)\n\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:22:54.968161Z","iopub.execute_input":"2025-03-05T09:22:54.968433Z","iopub.status.idle":"2025-03-05T09:22:54.974943Z","shell.execute_reply.started":"2025-03-05T09:22:54.968412Z","shell.execute_reply":"2025-03-05T09:22:54.974209Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from rdkit import Chem\nfrom rdkit.Chem import rdmolops\nimport networkx as nx\nimport matplotlib.pyplot as plt\n\ndef smiles_to_molgraph(smiles):\n    # 使用 RDKit 轉換 SMILES 字符串為分子對象\n    mol = Chem.MolFromSmiles(smiles)\n    if mol is None:\n        return None\n    \n    # 創建空的 NetworkX 圖來表示分子圖\n    G = nx.Graph()\n    \n    # 添加原子節點到圖中\n    for atom in mol.GetAtoms():\n        G.add_node(atom.GetIdx(), element=atom.GetSymbol())\n    \n    # 添加鍵邊到圖中\n    for bond in mol.GetBonds():\n        start_idx = bond.GetBeginAtomIdx()\n        end_idx = bond.GetEndAtomIdx()\n        bond_type = bond.GetBondTypeAsDouble()  # 鍵的類型（單鍵、雙鍵等）\n        G.add_edge(start_idx, end_idx, bond_type=bond_type)\n    \n    return G\n\ndef draw_molgraph(G):\n    # 使用 matplotlib 繪製分子圖\n    pos = nx.spring_layout(G)  # 使用 spring layout 算法來排布節點\n    labels = nx.get_node_attributes(G, 'element')  # 標籤為原子符號\n    \n    # 繪製圖形\n    nx.draw(G, pos, with_labels=True, labels=labels, node_size=700, node_color='skyblue', font_size=10)\n    edge_labels = nx.get_edge_attributes(G, 'bond_type')  # 鍵的類型標籤\n    nx.draw_networkx_edge_labels(G, pos, edge_labels=edge_labels)\n    plt.show()\n\n# 測試：轉換 SMILES 並顯示 MolGraph\nsmiles = \"CCO\"\nmol_graph = smiles_to_molgraph(smiles)\ndraw_molgraph(mol_graph)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:22:54.975981Z","iopub.execute_input":"2025-03-05T09:22:54.976266Z","iopub.status.idle":"2025-03-05T09:22:55.700518Z","shell.execute_reply.started":"2025-03-05T09:22:54.976238Z","shell.execute_reply":"2025-03-05T09:22:55.699713Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from rdkit import Chem\nfrom rdkit.Chem.Pharm2D import Gobbi_Pharm2D, Generate\n\n# 讀取分子\nsmiles = \"O=C(c1ccccc1S(=O)(=O)N2CCN(CC2)c3nnc(s3)C(F)(F)F)\"\nmol = Chem.MolFromSmiles(smiles)\n\n# 生成 Pharmacophore Fingerprint\nfp = Generate.Gen2DFingerprint(mol, Gobbi_Pharm2D.factory)\nprint(fp)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:22:55.701427Z","iopub.execute_input":"2025-03-05T09:22:55.701896Z","iopub.status.idle":"2025-03-05T09:22:55.738916Z","shell.execute_reply.started":"2025-03-05T09:22:55.70187Z","shell.execute_reply":"2025-03-05T09:22:55.738248Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install torch-scatter torch-sparse torch-cluster torch-spline-conv -f https://data.pyg.org/whl/torch-2.1.0+cu118.html\n\n!pip install torch-geometric","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:22:55.739578Z","iopub.execute_input":"2025-03-05T09:22:55.739808Z","iopub.status.idle":"2025-03-05T09:23:05.548499Z","shell.execute_reply.started":"2025-03-05T09:22:55.739789Z","shell.execute_reply":"2025-03-05T09:23:05.547583Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nfrom rdkit import Chem\nfrom torch_geometric.data import Data\n\ndef smiles_to_pyg_graph(smiles):\n    \"\"\"將 SMILES 轉換成 PyTorch Geometric 的圖結構\"\"\"\n    mol = Chem.MolFromSmiles(smiles)\n    if mol is None:\n        return None\n\n    # 原子特徵\n    atom_features = [atom.GetAtomicNum() for atom in mol.GetAtoms()]\n    x = torch.tensor(atom_features, dtype=torch.float32).view(-1, 1)\n\n    # 建立鍵結（Edge）資訊\n    edge_index = []\n    for bond in mol.GetBonds():\n        i, j = bond.GetBeginAtomIdx(), bond.GetEndAtomIdx()\n        edge_index.append([i, j])\n        edge_index.append([j, i])  # 雙向邊\n\n    edge_index = torch.tensor(edge_index, dtype=torch.long).t().contiguous()\n\n    return Data(x=x, edge_index=edge_index)\n\n# 測試\nsmiles = \"CCO\"  # 乙醇\ngraph = smiles_to_pyg_graph(smiles)\nprint(graph)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:23:05.549551Z","iopub.execute_input":"2025-03-05T09:23:05.549914Z","iopub.status.idle":"2025-03-05T09:23:09.110359Z","shell.execute_reply.started":"2025-03-05T09:23:05.549884Z","shell.execute_reply":"2025-03-05T09:23:09.109432Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport pandas as pd\nimport torch\n%cd MAT/src\n#os.chdir('src')\nfrom featurization.data_utils import load_data_from_df, construct_loader\n\nbatch_size = 64\n\n# Formal charges are one-hot encoded to keep compatibility with the pre-trained weights.\n# If you do not plan to use the pre-trained weights, we recommend to set one_hot_formal_charge to False.\nX, y = load_data_from_df('/kaggle/working/MAT/data/freesolv/freesolv.csv', one_hot_formal_charge=True)\ndata_loader = construct_loader(X, y, batch_size)\n\n\npd.read_csv('../data/freesolv/freesolv.csv').head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:23:09.111308Z","iopub.execute_input":"2025-03-05T09:23:09.111823Z","iopub.status.idle":"2025-03-05T09:23:15.4445Z","shell.execute_reply.started":"2025-03-05T09:23:09.111795Z","shell.execute_reply":"2025-03-05T09:23:15.443801Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from transformer import make_model\nd_atom = X[0][0].shape[1]  # It depends on the used featurization.\n\nmodel_params = {\n    'd_atom': d_atom,\n    'd_model': 1024,\n    'N': 8,\n    'h': 16,\n    'N_dense': 1,\n    'lambda_attention': 0.33, \n    'lambda_distance': 0.33,\n    'leaky_relu_slope': 0.1, \n    'dense_output_nonlinearity': 'relu', \n    'distance_matrix_kernel': 'exp', \n    'dropout': 0.0,\n    'aggregation_type': 'mean'\n}\n\nmodel = make_model(**model_params)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:23:15.44557Z","iopub.execute_input":"2025-03-05T09:23:15.445916Z","iopub.status.idle":"2025-03-05T09:23:15.844592Z","shell.execute_reply.started":"2025-03-05T09:23:15.445885Z","shell.execute_reply":"2025-03-05T09:23:15.843736Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pretrained_name = '/kaggle/input/pretrain/pretrained_weights.pt'  \npretrained_state_dict = torch.load(pretrained_name)\nmodel_state_dict = model.state_dict()\nfor name, param in pretrained_state_dict.items():\n    if 'generator' in name:\n         continue\n    if isinstance(param, torch.nn.Parameter):\n        param = param.data\n    model_state_dict[name].copy_(param)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:23:15.845388Z","iopub.execute_input":"2025-03-05T09:23:15.845605Z","iopub.status.idle":"2025-03-05T09:23:17.914587Z","shell.execute_reply.started":"2025-03-05T09:23:15.845586Z","shell.execute_reply":"2025-03-05T09:23:17.913812Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.to(torch.device('cuda'))\n\nfor batch in data_loader:\n    adjacency_matrix, node_features, distance_matrix, y = batch\n    batch_mask = torch.sum(torch.abs(node_features), dim=-1) != 0\n    output = model(node_features, batch_mask, adjacency_matrix, distance_matrix, None)\n\n    print(output)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:23:17.915484Z","iopub.execute_input":"2025-03-05T09:23:17.915812Z","iopub.status.idle":"2025-03-05T09:23:18.904342Z","shell.execute_reply.started":"2025-03-05T09:23:17.915783Z","shell.execute_reply":"2025-03-05T09:23:18.903237Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install pyarrow  # 或者你也可以選擇 fastparquet\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:23:18.905856Z","iopub.execute_input":"2025-03-05T09:23:18.906109Z","iopub.status.idle":"2025-03-05T09:23:22.303486Z","shell.execute_reply.started":"2025-03-05T09:23:18.906088Z","shell.execute_reply":"2025-03-05T09:23:22.302583Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n#data = pd.read_parquet('/kaggle/input/leash-BELKA/train.parquet')\n#test_df = pd.read_parquet('/kaggle/input/leash-BELKA/test.parquet')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:23:22.304539Z","iopub.execute_input":"2025-03-05T09:23:22.304824Z","iopub.status.idle":"2025-03-05T09:23:22.308679Z","shell.execute_reply.started":"2025-03-05T09:23:22.304788Z","shell.execute_reply":"2025-03-05T09:23:22.307844Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport gc  # 引入垃圾回收模組\n\nfilename = \"/kaggle/input/leash-BELKA/train.csv\"\ncolumns_to_read = [\"molecule_smiles\", \"protein_name\", \"binds\"]\n\nbatch_size = 60000       # 每次讀取 60,000 筆\ntarget_rows = 3600000    # 每輪儲存 1,200,000 筆\ntotal_batches = 80       # 總共執行 30 次\ntotal_rows = 0           # 記錄當前累積筆數\n\n# 用 chunksize 讀取 CSV\ncsv_reader = pd.read_csv(filename, usecols=columns_to_read, chunksize=batch_size)\n\nbatch_count = 0  # 記錄當前批次數\nchunks = []  # 暫存當前批次的資料\n\nfor chunk in csv_reader:\n    chunks.append(chunk)\n    total_rows += len(chunk)\n\n    # 當累積的資料量達到 `target_rows`，就存檔一次\n    if total_rows >= target_rows * (batch_count + 1):\n        batch_df = pd.concat(chunks, ignore_index=True)\n\n        # **將每 3 列合併成 1 列**\n        batch_df[\"row_idx\"] = batch_df.index // 3  # 每 3 列分組\n        batch_pivot = batch_df.pivot(index=\"row_idx\", columns=\"protein_name\", values=\"binds\").reset_index()\n\n        # **合併 molecule_smiles**\n        smiles_df = batch_df.groupby(\"row_idx\")[\"molecule_smiles\"].first().reset_index()\n        final_df = smiles_df.merge(batch_pivot, on=\"row_idx\").drop(columns=[\"row_idx\"])\n\n        # 存成 CSV，每次都存不同的檔案\n        output_filename = f\"/kaggle/working/sampled_train_part{batch_count+1}.csv\"\n        final_df.to_csv(output_filename, index=False)\n\n        print(f\"✅ 第 {batch_count+1} 次存檔：{len(final_df)} 筆，已累積 {total_rows} 筆\")\n\n        # 清理記憶體\n        del batch_df, batch_pivot, smiles_df, final_df\n        chunks = []  # 清空暫存區\n        gc.collect()  # 執行垃圾回收\n        batch_count += 1\n\n    if batch_count >= total_batches:\n        break  # 達到最大批次數就停止\n\n# 最後如果還有未存的資料，則再存一次\nif chunks:\n    batch_df = pd.concat(chunks, ignore_index=True)\n\n    batch_df[\"row_idx\"] = batch_df.index // 3\n    batch_pivot = batch_df.pivot(index=\"row_idx\", columns=\"protein_name\", values=\"binds\").reset_index()\n    smiles_df = batch_df.groupby(\"row_idx\")[\"molecule_smiles\"].first().reset_index()\n    final_df = smiles_df.merge(batch_pivot, on=\"row_idx\").drop(columns=[\"row_idx\"])\n\n \n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:23:22.309416Z","iopub.execute_input":"2025-03-05T09:23:22.309728Z","iopub.status.idle":"2025-03-05T09:44:01.810404Z","shell.execute_reply.started":"2025-03-05T09:23:22.309695Z","shell.execute_reply":"2025-03-05T09:44:01.809447Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"data1name = \"/kaggle/working/sampled_train_part1.csv\"\ndata1= pd.read_csv(data1name)\ndata1.head(10)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:44:01.811384Z","iopub.execute_input":"2025-03-05T09:44:01.811742Z","iopub.status.idle":"2025-03-05T09:44:03.580385Z","shell.execute_reply.started":"2025-03-05T09:44:01.81171Z","shell.execute_reply":"2025-03-05T09:44:03.579692Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"data1.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:44:03.581077Z","iopub.execute_input":"2025-03-05T09:44:03.581294Z","iopub.status.idle":"2025-03-05T09:44:03.586651Z","shell.execute_reply.started":"2025-03-05T09:44:03.581275Z","shell.execute_reply":"2025-03-05T09:44:03.585733Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\n\ndef load_data_from_df(file_path, one_hot_formal_charge=True):\n    # 讀取 CSV\n    df = pd.read_csv(file_path)\n\n    # 定義 X（輸入特徵）和 y（目標變數）\n    X = df[\"molecule_smiles\"].tolist()  # SMILES 字串作為輸入\n    y = df[[\"BRD4\", \"HSA\", \"sEH\"]].values  # 3 欄數值作為輸出\n\n    return X, y\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T10:10:34.832251Z","iopub.execute_input":"2025-03-05T10:10:34.832589Z","iopub.status.idle":"2025-03-05T10:10:34.837235Z","shell.execute_reply.started":"2025-03-05T10:10:34.832567Z","shell.execute_reply":"2025-03-05T10:10:34.836247Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nfrom featurization.data_utils import load_data_from_df\n\n# 讀取 CSV\n#df = pd.read_csv('/kaggle/working/sampled_train_part1.csv')\n\nbatch_size = 1\n\n\n# 轉換成 MAT 格式\nX, y = load_data_from_df('/kaggle/working/example_data.csv', one_hot_formal_charge=True)\ndata_loader = construct_loader(X, y, batch_size)\n\nmodel.to(torch.device('cuda'))\n\nfor batch in data_loader:\n    adjacency_matrix, node_features, distance_matrix, y = batch\n    batch_mask = torch.sum(torch.abs(node_features), dim=-1) != 0\n    output = model(node_features, batch_mask, adjacency_matrix, distance_matrix, None)\n\n    print(output)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T10:10:36.012161Z","iopub.execute_input":"2025-03-05T10:10:36.012495Z","iopub.status.idle":"2025-03-05T10:10:36.041757Z","shell.execute_reply.started":"2025-03-05T10:10:36.012466Z","shell.execute_reply":"2025-03-05T10:10:36.040561Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"STOP","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:44:04.095001Z","iopub.status.idle":"2025-03-05T09:44:04.095262Z","shell.execute_reply":"2025-03-05T09:44:04.095159Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\n\ndef make_model(d_atom, d_model, N, h, N_dense, lambda_attention, lambda_distance,\n               leaky_relu_slope, dense_output_nonlinearity, distance_matrix_kernel,\n               dropout, aggregation_type):\n    # 假設這部分是你的 MAT 模型結構\n    # 這裡包含了 Transformer 層或其他前置層\n    \n    class MATModel(nn.Module):\n        def __init__(self):\n            super(MATModel, self).__init__()\n            # Transformer 部分\n            self.transformer = nn.Transformer(d_atom, nhead=h, num_encoder_layers=N)\n            \n            # 其他層設置（你可以根據模型的設計來調整）\n            # 例如，Dense 層\n            self.dense = nn.Linear(d_atom, d_model)\n            self.fc = nn.Linear(d_model, 3)  # 假設最後需要3個輸出（binds_BRD4, binds_HSA, binds_sEH）\n            \n            # 激活函數\n            self.sigmoid = nn.Sigmoid()\n\n        def forward(self, x):\n            # Transformer 層\n            x = self.transformer(x)\n            \n            # Dense 層\n            x = self.dense(x)\n            \n            # 全連接層\n            x = self.fc(x)\n            \n            # Sigmoid 激活\n            x = self.sigmoid(x)  # 多標籤分類\n            return x\n\n    return MATModel()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:44:03.587664Z","iopub.execute_input":"2025-03-05T09:44:03.58797Z","iopub.status.idle":"2025-03-05T09:44:03.603865Z","shell.execute_reply.started":"2025-03-05T09:44:03.587936Z","shell.execute_reply":"2025-03-05T09:44:03.602876Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model_params = {\n    'd_atom': d_atom,\n    'd_model': 1024,\n    'N': 8,\n    'h': 4,\n    'N_dense': 2,\n    'lambda_attention': 0.33, \n    'lambda_distance': 0.33,\n    'leaky_relu_slope': 0.1, \n    'dense_output_nonlinearity': 'relu', \n    'distance_matrix_kernel': 'exp', \n    'dropout': 0.0,\n    'aggregation_type': 'mean'\n}\n\nmodel = make_model(**model_params)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:44:03.604814Z","iopub.execute_input":"2025-03-05T09:44:03.60506Z","iopub.status.idle":"2025-03-05T09:44:03.731385Z","shell.execute_reply.started":"2025-03-05T09:44:03.605039Z","shell.execute_reply":"2025-03-05T09:44:03.730495Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport pandas as pd\nimport torch\nfrom featurization.data_utils import load_data_from_df, construct_loader\n\n# 設定 CUDA\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nmodel.to(device)\n\n# 參數設定\nbatch_size = 128\nnum_echoes = 1  # 每個檔案跑 10 輪\n\n# 取得所有檔案列表\nfile_dir = \"/kaggle/working\"\nfile_pattern = \"sampled_train_part{}.csv\"\nfile_list = [file_pattern.format(i) for i in range(1, 81)]  # 1 到 80\n\n# 開始處理每個檔案\nfor file_idx, file_name in enumerate(file_list, start=1):\n    file_path = os.path.join(file_dir, file_name)\n    \n    if not os.path.exists(file_path):\n        print(f\"⚠️ 檔案 {file_name} 不存在，跳過\")\n        continue\n\n    print(f\"🚀 正在處理：{file_name} ({file_idx}/80)\")\n\n    # 讀取 CSV 並轉換為 data loader\n    X, y = load_data_from_df(file_path, one_hot_formal_charge=True)\n    data_loader = construct_loader(X, y, batch_size)\n\n    # 進行 10 輪 Echo\n    for echo in range(num_echoes):\n        print(f\"🔁 Echo {echo + 1}/{num_echoes} 開始...\")\n\n        for batch in data_loader:\n            adjacency_matrix, node_features, distance_matrix, y = batch\n            \n            # 移動到 GPU（如果可用）\n            adjacency_matrix = adjacency_matrix.to(device)\n            node_features = node_features.to(device)\n            distance_matrix = distance_matrix.to(device)\n            y = y.to(device)\n\n            batch_mask = torch.sum(torch.abs(node_features), dim=-1) != 0\n            output = model(node_features, batch_mask, adjacency_matrix, distance_matrix, None)\n\n            print(f\"✅ Echo {echo + 1}, 檔案 {file_name}, Batch size: {len(y)}, 輸出範例: {output[:3]}\")\n\n        print(f\"🔄 Echo {echo + 1}/{num_echoes} 完成！\")\n\n    print(f\"🎯 檔案 {file_name} 處理完成！\\n\")\n\nprint(\"🎉 所有檔案處理完畢！\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:44:04.096146Z","iopub.status.idle":"2025-03-05T09:44:04.096523Z","shell.execute_reply":"2025-03-05T09:44:04.096357Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"只要smiles轉graph\ngraph = x\nx->mat\nfp \nbind = y\nmodel.fit(fp,y)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:44:04.097252Z","iopub.status.idle":"2025-03-05T09:44:04.09758Z","shell.execute_reply":"2025-03-05T09:44:04.097425Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pyarrow.parquet as pq\nimport pandas as pd\nimport gc  # 引入垃圾回收模組\n\nfilename = \"/kaggle/input/leash-BELKA/train.parquet\"\ncolumns_to_read = [\"molecule_smiles\", \"protein_name\", \"binds\"]\n\nbatch_size = 60000    # 每次讀取 60,000 筆\ntarget_rows = 12000000  # 每輪儲存 1,200,0000 筆\ntotal_batches = 30    # 總共執行 30 次\ntotal_rows = 0        # 記錄當前累積筆數\n\nparquet_file = pq.ParquetFile(filename)\n\n# 獲取總行組數量\nnum_row_groups = parquet_file.num_row_groups\nprint(f\"Total row groups in file: {num_row_groups}\")\n\n# 計算每次讀取多少行組\nrow_groups_per_batch = target_rows // batch_size\n\n\n# 開始進行批次處理\nfor i in range(total_batches):\n    chunks = []\n    current_rows = 0  # 每輪的計數器\n    batch_start_row_group = i * row_groups_per_batch  # 直接依序取\n    batch_end_row_group = min(batch_start_row_group + row_groups_per_batch, num_row_groups)\n\n    if batch_end_row_group >= num_row_groups:\n        batch_end_row_group = num_row_groups  # 避免超出行組範圍\n\n    print(f\"✅ 第 {i+1} 次處理：從行組 {batch_start_row_group} 到行組 {batch_end_row_group}\")\n\n    # 使用 pyarrow 的 ParquetFile 直接讀取指定範圍的行組\n    for row_group_idx in range(batch_start_row_group, batch_end_row_group):\n        try:\n            batch = parquet_file.read_row_groups([row_group_idx], columns=columns_to_read)\n            chunk = batch.to_pandas()\n\n            # 檢查是否有資料\n            if not chunk.empty:\n                chunks.append(chunk)\n                current_rows += len(chunk)\n                total_rows += len(chunk)\n\n            # 如果讀取到指定範圍的資料，就停止\n            if total_rows >= target_rows * (i + 1):\n                break  # 如果已經讀到該批次的結尾就停止\n\n        except Exception as e:\n            print(f\"⚠️ 讀取行組 {row_group_idx} 時發生錯誤: {e}\")\n\n    if chunks:  # 確保有資料才進行合併\n        # 合併 DataFrame\n        batch_df = pd.concat(chunks, ignore_index=True)\n\n        # **將每 3 列合併成 1 列**\n        batch_df[\"row_idx\"] = batch_df.index // 3  # 每 3 列分組\n        batch_pivot = batch_df.pivot(index=\"row_idx\", columns=\"protein_name\", values=\"binds\").reset_index()\n\n        # **合併 molecule_smiles**\n        smiles_df = batch_df.groupby(\"row_idx\")[\"molecule_smiles\"].first().reset_index()\n        final_df = smiles_df.merge(batch_pivot, on=\"row_idx\").drop(columns=[\"row_idx\"])\n\n        # 存成 parquet，每次都存不同的檔案\n        output_filename = f\"/kaggle/working/sampled_train_part{i+1}.parquet\"\n        final_df.to_parquet(output_filename, index=False)\n\n        print(f\"✅ 第 {i+1} 次存檔：{len(final_df)} 筆，已累積 {total_rows} 筆\")\n\n        # 清理無用的變數，釋放記憶體\n        del batch_df, batch_pivot, smiles_df, final_df\n        gc.collect()  # 執行垃圾回收\n    else:\n        print(f\"⚠️ 第 {i+1} 次處理未讀取到任何資料，跳過該批次。\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:44:04.098701Z","iopub.status.idle":"2025-03-05T09:44:04.098984Z","shell.execute_reply":"2025-03-05T09:44:04.098876Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pyarrow.parquet as pq\nimport pandas as pd\nimport gc  # 引入垃圾回收模組\n\nfilename = \"/kaggle/input/leash-BELKA/train.parquet\"\ncolumns_to_read = [\"molecule_smiles\", \"protein_name\", \"binds\"]\n\nbatch_size = 60000       # 每次讀取 60,000 筆\ntarget_rows = 12000000   # 每輪儲存 12,000,000 筆\ntotal_batches = 30       # 總共執行 30 次\ntotal_rows = 0           # 記錄當前累積筆數\n\nparquet_file = pq.ParquetFile(filename)\n\n# 獲取總行組數量\nnum_row_groups = parquet_file.num_row_groups\nprint(f\"Total row groups in file: {num_row_groups}\")\n\n# 計算每次讀取多少行組\nrow_groups_per_batch = target_rows // batch_size  # 計算每次應該讀取多少個行組\n\n# 開始進行批次處理\nfor i in range(total_batches):\n    chunks = []\n    current_rows = 0  # 每輪的計數器\n    batch_start_row_group = total_rows // batch_size  # 計算起始行組索引\n    batch_end_row_group = batch_start_row_group + row_groups_per_batch  # 計算結束行組索引\n\n    # 確保不超過行組數量\n    if batch_end_row_group >= num_row_groups:\n        batch_end_row_group = num_row_groups\n\n    print(f\"✅ 第 {i+1} 次處理：從行組 {batch_start_row_group} 到行組 {batch_end_row_group}\")\n\n    # 依序讀取行組\n    for row_group_idx in range(batch_start_row_group, batch_end_row_group):\n        try:\n            batch = parquet_file.read_row_groups([row_group_idx], columns=columns_to_read)\n            chunk = batch.to_pandas()\n\n            # 確保有資料才加入\n            if not chunk.empty:\n                chunks.append(chunk)\n                current_rows += len(chunk)\n                total_rows += len(chunk)\n\n            # 若本批次已達到 `target_rows`，則停止\n            if total_rows >= target_rows * (i + 1):\n                break\n\n        except Exception as e:\n            print(f\"⚠️ 讀取行組 {row_group_idx} 時發生錯誤: {e}\")\n\n    if chunks:  # 確保有資料才進行合併\n        # 合併 DataFrame\n        batch_df = pd.concat(chunks, ignore_index=True)\n\n        # **將每 3 列合併成 1 列**\n        batch_df[\"row_idx\"] = batch_df.index // 3  # 每 3 列分組\n        batch_pivot = batch_df.pivot(index=\"row_idx\", columns=\"protein_name\", values=\"binds\").reset_index()\n\n        # **合併 molecule_smiles**\n        smiles_df = batch_df.groupby(\"row_idx\")[\"molecule_smiles\"].first().reset_index()\n        final_df = smiles_df.merge(batch_pivot, on=\"row_idx\").drop(columns=[\"row_idx\"])\n\n        # 存成 parquet，每次都存不同的檔案\n        output_filename = f\"/kaggle/working/sampled_train_part{i+1}.parquet\"\n        final_df.to_parquet(output_filename, index=False)\n\n        print(f\"✅ 第 {i+1} 次存檔：{len(final_df)} 筆，已累積 {total_rows} 筆\")\n\n        # 清理記憶體\n        del batch_df, batch_pivot, smiles_df, final_df\n        gc.collect()  # 執行垃圾回收\n    else:\n        print(f\"⚠️ 第 {i+1} 次處理未讀取到任何資料，跳過該批次。\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:44:04.099737Z","iopub.status.idle":"2025-03-05T09:44:04.100059Z","shell.execute_reply":"2025-03-05T09:44:04.099947Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"data1name = \"/kaggle/working/sampled_train_part1.parquet\"\ndata1= pd.read_parquet(data1name)\ndata1.head(10)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:44:04.10101Z","iopub.status.idle":"2025-03-05T09:44:04.10131Z","shell.execute_reply":"2025-03-05T09:44:04.101182Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"data1.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:44:04.102149Z","iopub.status.idle":"2025-03-05T09:44:04.10245Z","shell.execute_reply":"2025-03-05T09:44:04.102347Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"data2name = \"/kaggle/working/sampled_train_part2.parquet\"\ndata2= pd.read_parquet(data2name)\ndata2.head(10)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:44:04.10348Z","iopub.status.idle":"2025-03-05T09:44:04.103869Z","shell.execute_reply":"2025-03-05T09:44:04.103678Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"data2.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:44:04.104588Z","iopub.status.idle":"2025-03-05T09:44:04.104914Z","shell.execute_reply":"2025-03-05T09:44:04.104801Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"filename = \"/kaggle/input/leash-BELKA/train.parquet\"\ndata3= pd.read_parquet(filename)\ndata3.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:44:04.105815Z","iopub.status.idle":"2025-03-05T09:44:04.106132Z","shell.execute_reply":"2025-03-05T09:44:04.105982Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nfrom rdkit import Chem\nfrom rdkit.Chem import rdmolops\nimport torch\nfrom torch_geometric.data import Data\n\n# 定義將 SMILES 轉換為圖結構的函數\ndef smiles_to_pyg_graph(smiles):\n    mol = Chem.MolFromSmiles(smiles)\n    if mol is not None:\n        # 生成分子圖\n        atoms = [atom.GetAtomicNum() for atom in mol.GetAtoms()]\n        bonds = [(bond.GetBeginAtomIdx(), bond.GetEndAtomIdx()) for bond in mol.GetBonds()]\n\n        # 創建邊列表\n        edge_index = torch.tensor(bonds, dtype=torch.long).t().contiguous()\n\n        # 創建特徵矩陣\n        x = torch.tensor(atoms, dtype=torch.float).view(-1, 1)  # 每個原子的特徵\n\n        # 創建 PyG 格式的數據對象\n        data = Data(x=x, edge_index=edge_index)\n\n        return data\n    else:\n        return None\n\n# 設定讀取 Parquet 文件的參數\nfilename = '/kaggle/input/leash-BELKA/train.parquet'\nchunk_size = 100000  # 每次讀取的行數\nnrows = 500000  # 需要處理的總行數\n\n# 讀取並處理每一個 chunk\nchunks = []\nfor chunk in pd.read_parquet(filename, chunksize=chunk_size, usecols=['molecule_smiles', 'protein', 'bind'], nrows=nrows):\n    # 將 molecule_smiles 轉換為圖結構\n    chunk['pyg_graph'] = chunk['molecule_smiles'].apply(smiles_to_pyg_graph)\n    chunks.append(chunk)\n\n# 合併所有的 chunks 以形成最終的 DataFrame\ntrain_df = pd.concat(chunks, ignore_index=True)\n\n# 顯示合併後的 DataFrame\nprint(train_df[['molecule_smiles', 'protein', 'bind', 'pyg_graph']].head())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:44:04.107105Z","iopub.status.idle":"2025-03-05T09:44:04.107376Z","shell.execute_reply":"2025-03-05T09:44:04.107269Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"只要smiles轉graph\ngraph = x\nx->mat\nfp \nbind = y\nmodel.fit(fp,y)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:44:04.108271Z","iopub.status.idle":"2025-03-05T09:44:04.108583Z","shell.execute_reply":"2025-03-05T09:44:04.108475Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"batch_size = 64\n\n# Formal charges are one-hot encoded to keep compatibility with the pre-trained weights.\n# If you do not plan to use the pre-trained weights, we recommend to set one_hot_formal_charge to False.\nX, y = train_df\ndata_loader = construct_loader(X, y, batch_size)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:44:04.109622Z","iopub.status.idle":"2025-03-05T09:44:04.109973Z","shell.execute_reply":"2025-03-05T09:44:04.10984Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.to(torch.device('cuda'))\n\nfor batch in data_loader:\n    adjacency_matrix, node_features, distance_matrix, y = batch\n    batch_mask = torch.sum(torch.abs(node_features), dim=-1) != 0\n    output = model(node_features, batch_mask, adjacency_matrix, distance_matrix, None)\n\n    print(output)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:44:04.1106Z","iopub.status.idle":"2025-03-05T09:44:04.11092Z","shell.execute_reply":"2025-03-05T09:44:04.110813Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#測試\n\nclass GNNFingerprint(nn.Module):\n    def __init__(self, in_dim=1, hidden_dim=64, out_dim=128, num_layers=3):\n        super(GNNFingerprint, self).__init__()\n        self.conv_layers = nn.ModuleList([GCNConv(in_dim if i == 0 else hidden_dim, hidden_dim) for i in range(num_layers)])\n        self.fc = nn.Linear(hidden_dim, out_dim)  # 最終輸出分子指紋\n\n    def forward(self, data):\n        x, edge_index = data.x, data.edge_index\n        for conv in self.conv_layers:\n            x = F.relu(conv(x, edge_index))\n        x = global_mean_pool(x, data.batch)  # 聚合成全局分子指紋\n        return self.fc(x)\n\n# 測試 GNN 指紋模型\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel = GNNFingerprint().to(device)\ngraph = graph.to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:44:04.112123Z","iopub.status.idle":"2025-03-05T09:44:04.112408Z","shell.execute_reply":"2025-03-05T09:44:04.112299Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nfrom rdkit import Chem\nfrom rdkit.Chem.Pharm2D import Gobbi_Pharm2D, Generate\n\n# 定義一個將 SMILES 轉換為 Pharmacophore Fingerprint 的函數\ndef smiles_to_fingerprint(smiles):\n    mol = Chem.MolFromSmiles(smiles)\n    if mol is not None:\n        fp = Generate.Gen2DFingerprint(mol, Gobbi_Pharm2D.factory)\n        return fp\n    else:\n        return None\n\n# 設定讀取 Parquet 文件的參數\nfilename = '/kaggle/input/leash-BELKA/train.parquet'\nchunk_size = 100000  # 每次讀取的行數\nnrows = 500000  # 需要處理的總行數\n\n# 讀取並處理每一個 chunk\nchunks = []\nfor chunk in pd.read_parquet(filename, chunksize=chunk_size, nrows=nrows):\n    chunk['pyg_graph'] = chunk['molecule_smiles'].apply(smiles_to_pyg_graph)\n    chunks.append(chunk)\n\n# 合併所有的 chunks 以形成最終的 DataFrame\ntrain_df = pd.concat(chunks, ignore_index=True)\n\n# 顯示合併後的 DataFrame\nprint(train_df.head())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:44:04.113116Z","iopub.status.idle":"2025-03-05T09:44:04.113444Z","shell.execute_reply":"2025-03-05T09:44:04.113288Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport torch\nfrom torch.utils.data import Dataset, DataLoader\nfrom convert_smiles_to_graph import smiles_to_graph\n\nclass MolecularDataset(Dataset):\n    def __init__(self, csv_file):\n        self.data = pd.read_csv(csv_file)\n\n    def __len__(self):\n        return len(self.data)\n\n    def __getitem__(self, idx):\n        smiles = self.data.iloc[idx][\"SMILES\"]\n        label = torch.tensor(self.data.iloc[idx][\"Property\"], dtype=torch.float32)\n        graph = smiles_to_graph(smiles)\n\n        return graph, label\n\n# 建立 DataLoader\ndef get_data_loader(csv_file, batch_size=32):\n    dataset = MolecularDataset(csv_file)\n    return DataLoader(dataset, batch_size=batch_size, shuffle=True, collate_fn=collate_fn)\n\n# 設定批量處理函數\ndef collate_fn(batch):\n    graphs, labels = zip(*batch)\n    graphs = dgl.batch(graphs)\n    labels = torch.stack(labels)\n    return graphs, labels\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:44:04.114211Z","iopub.status.idle":"2025-03-05T09:44:04.114488Z","shell.execute_reply":"2025-03-05T09:44:04.114369Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom models import MAT\nfrom dataset import get_data_loader\n\n# 設定 GPU\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# 1. 載入 MAT 模型\nmodel = MAT()\nmodel.load_state_dict(torch.load(\"pretrained/mat.pth\", map_location=device))  # 載入預訓練模型\nmodel.to(device)\n\n# 2. 設定損失函數和優化器\ncriterion = nn.MSELoss()  # 均方誤差（適用於回歸）\noptimizer = optim.Adam(model.parameters(), lr=0.0001)\n\n# 3. 讀取數據集\ntrain_loader = get_data_loader(\"data/custom_dataset.csv\", batch_size=16)\n\n# 4. 訓練模型\nnum_epochs = 20\nfor epoch in range(num_epochs):\n    model.train()\n    total_loss = 0\n\n    for graphs, labels in train_loader:\n        graphs, labels = graphs.to(device), labels.to(device)\n\n        optimizer.zero_grad()\n        outputs = model(graphs)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n\n        total_loss += loss.item()\n\n    print(f\"Epoch [{epoch+1}/{num_epochs}], Loss: {total_loss:.4f}\")\n\n# 5. 儲存微調後的模型\ntorch.save(model.state_dict(), \"fine_tuned_mat.pth\")\nprint(\"微調完成，模型已保存！\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:44:04.115245Z","iopub.status.idle":"2025-03-05T09:44:04.11565Z","shell.execute_reply":"2025-03-05T09:44:04.11547Z"}},"outputs":[],"execution_count":null}]}