{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":67356,"databundleVersionId":8006601},{"sourceType":"datasetVersion","sourceId":15710684,"datasetId":10064972,"databundleVersionId":16650631},{"sourceType":"modelInstanceVersion","sourceId":829972,"databundleVersionId":16650597,"modelInstanceId":631258,"modelId":643154},{"sourceType":"modelInstanceVersion","sourceId":838941,"databundleVersionId":16780923,"modelInstanceId":638158,"modelId":650143},{"sourceType":"modelInstanceVersion","sourceId":847000,"databundleVersionId":16901617,"modelInstanceId":644070,"modelId":656033},{"sourceType":"modelInstanceVersion","sourceId":847815,"databundleVersionId":16912868,"modelInstanceId":644669,"modelId":656621},{"sourceType":"modelInstanceVersion","sourceId":822944,"databundleVersionId":16521652,"modelInstanceId":625490,"modelId":637354},{"sourceType":"kernelVersion","sourceId":171884838},{"sourceType":"kernelVersion","sourceId":312874589},{"sourceType":"kernelVersion","sourceId":315000306}],"dockerImageVersionId":31328,"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":"2026-06-19T11:35:42.11472Z","iopub.execute_input":"2026-06-19T11:35:42.115027Z","iopub.status.idle":"2026-06-19T11:35:42.157007Z","shell.execute_reply.started":"2026-06-19T11:35:42.115005Z","shell.execute_reply":"2026-06-19T11:35:42.15608Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install transformers torch rdkit","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-19T11:35:42.158571Z","iopub.execute_input":"2026-06-19T11:35:42.158754Z","iopub.status.idle":"2026-06-19T11:35:46.001067Z","shell.execute_reply.started":"2026-06-19T11:35:42.158735Z","shell.execute_reply":"2026-06-19T11:35:45.99995Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport gc\nimport torch\nimport polars as pl\nimport pandas as pd\nfrom tqdm import tqdm\nfrom torch.utils.data import DataLoader, Dataset\nfrom torch.amp import autocast\nfrom transformers import DataCollatorWithPadding\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport numpy as np\nfrom tqdm import tqdm\nfrom transformers import AutoTokenizer, AutoModel\nfrom sklearn.metrics import average_precision_score\nfrom sklearn.model_selection import train_test_split\nimport pyarrow.parquet as pq\nfrom torch.amp import GradScaler","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-19T11:35:46.002172Z","iopub.execute_input":"2026-06-19T11:35:46.002386Z","iopub.status.idle":"2026-06-19T11:35:46.140999Z","shell.execute_reply.started":"2026-06-19T11:35:46.002361Z","shell.execute_reply":"2026-06-19T11:35:46.140085Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"'''\n# --- 1. 數據優化抽樣 (Optimized Sampling) ---\n# 參數設定\nINPUT_FILE = \"/kaggle/input/competitions/leash-BELKA/train.csv\"\nOUTPUT_FILE = \"optimized_train_data.csv\"\nCHUNK_SIZE = 3_000_000\n\ndef csv_to_parquet_optimized(input_csv, output_parquet):\n    print(\"正在將 CSV 轉換為高效能 Parquet 格式...\")\n    # 使用 Polars 讀取，並指定數據類型以節省空間\n    df = pl.read_csv(input_csv, dtypes={\n        \"molecule_smiles\": pl.Utf8,\n        \"protein_name\": pl.Categorical,\n        \"binds\": pl.Int8\n    })\n    df.write_parquet(output_parquet, compression=\"snappy\")\n    print(f\"✅ 轉換完成: {output_parquet}\")\n    del df\n    gc.collect()\n\ndef optimized_sampling_to_parquet(input_parquet, output_parquet, chunk_size=3_000_000, neg_to_pos_ratio=16):\n    print(\"Sampling...\")\n    \n    # 關鍵動作：建立 Parquet 檔案物件\n    parquet_file = pq.ParquetFile(input_parquet)\n    \n    # ==========================================\n    # 第一趟 (Pass 1)：母體普查\n    # ==========================================\n    print(\"\\n [Pass 1] 正在進行全資料庫極速普查...\")\n    total_pos_count = 0\n    total_neg_count = 0\n    \n    # 使用 iter_batches 來分批讀取，並轉成 Pandas DataFrame\n    for batch in tqdm(parquet_file.iter_batches(batch_size=chunk_size, columns=['binds']), desc=\"Counting\"):\n        chunk = batch.to_pandas()\n        total_pos_count += (chunk['binds'] == 1).sum()\n        total_neg_count += (chunk['binds'] == 0).sum()\n        \n    print(f\"母體總數 - 正樣本: {total_pos_count:,} | 負樣本: {total_neg_count:,}\")\n    \n    # 精算抽樣率\n    target_neg_count = total_pos_count * neg_to_pos_ratio\n    exact_sampling_rate = target_neg_count / total_neg_count\n    \n    print(f\" 目標負樣本數: {target_neg_count:,} (比例 15:1)\")\n    print(f\" 算出抽樣機率: {exact_sampling_rate:.8f}\")\n\n    # ==========================================\n    # 第二趟 (Pass 2)：精準資料提取\n    # ==========================================\n    print(\"\\n [Pass 2] 正在進行精準資料提取與抽樣...\")\n    use_cols = ['molecule_smiles', 'protein_name', 'binds', 'buildingblock1_smiles']\n    \n    pos_chunks = []\n    neg_chunks = []\n\n    for batch in tqdm(parquet_file.iter_batches(batch_size=chunk_size, columns=use_cols), desc=\"Extracting & Sampling\"):\n        chunk = batch.to_pandas()\n        \n        # 1. 抓出所有正樣本\n        pos_chunks.append(chunk[chunk['binds'] == 1])\n        \n        # 2. 抓出負樣本並使用機率抽樣\n        neg_subset = chunk[chunk['binds'] == 0]\n        neg_chunks.append(neg_subset.sample(frac=exact_sampling_rate, random_state=2))\n\n    print(\"\\n 讀取完畢！正在合併數據...\")\n    df_pos = pd.concat(pos_chunks)\n    df_neg = pd.concat(neg_chunks)\n    \n    # 合併並洗牌\n    df_final = pd.concat([df_pos, df_neg]).sample(frac=1, random_state=2).reset_index(drop=True)\n\n    print(f\"最終樣本數: {len(df_final):,} (正: {len(df_pos):,}, 負: {len(df_neg):,})\")\n    \n    # ==========================================\n    # 寫入與輸出\n    # ==========================================\n    print(\"正在轉換為 Polars 並寫入高效能 Parquet...\")\n    \n    pl_df = pl.from_pandas(df_final).with_columns([\n        pl.col(\"molecule_smiles\").cast(pl.Utf8),\n        pl.col(\"buildingblock1_smiles\").cast(pl.Utf8),  \n        pl.col(\"protein_name\").cast(pl.Categorical),    \n        pl.col(\"binds\").cast(pl.Int8)\n    ])\n    \n    pl_df.write_parquet(output_parquet, compression=\"snappy\")\n    \n    print(f\"完美轉換！檔案已儲存至 {output_parquet}\")\n    \n    del df_pos, df_neg, df_final, pl_df\n    gc.collect()\n\n\n\noptimized_sampling_to_parquet(\"/kaggle/input/competitions/leash-BELKA/train.parquet\", \"optimized_train_data.parquet\")\n\ngc.collect()\n\nif torch.cuda.is_available():\n    torch.cuda.empty_cache()\n'''","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-19T11:35:46.143126Z","iopub.execute_input":"2026-06-19T11:35:46.143379Z","iopub.status.idle":"2026-06-19T11:35:46.152938Z","shell.execute_reply.started":"2026-06-19T11:35:46.143356Z","shell.execute_reply":"2026-06-19T11:35:46.152073Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"'''\nimport polars as pl\nimport gc\n\ndef merge_and_shuffle_data(main_parquet_path, external_csv_path, output_path, read_count_threshold=15):\n    print(\"[Phase 1] 啟動外部資料融合...\")\n    \n    # 1. 讀取主資料集\n    print(f\"讀取主資料集: {main_parquet_path}\")\n    df_main = pl.read_parquet(main_parquet_path)\n    \n    # 2. 讀取外部資料並進行嚴格篩選\n    print(f\"📥 讀取外部資料: {external_csv_path}\")\n    df_ext_raw = pl.read_csv(external_csv_path)\n    \n    # 過濾低頻雜訊 \n    initial_count = len(df_ext_raw)\n    df_ext_pos_filtered = df_ext_raw.filter(pl.col(\"read_count\") >= read_count_threshold)\n    print(f\"pos_filtered -> {len(df_ext_pos_filtered)} 筆)\")\n    df_ext_neg_filtered = df_ext_raw.filter(pl.col(\"read_count\") == 0)\n    print(f\"neg_filtered -> {len(df_ext_neg_filtered)} 筆)\")\n    \n    # 保留 [Dy] 並使用 iso 立體結構\n    df_positive = df_ext_pos_filtered.select([\n        pl.col(\"new_structure\").alias(\"molecule_smiles\"),         # 完美保留 [Dy]\n        pl.col(\"bb1_iso\").alias(\"buildingblock1_smiles\"),         # 完美保留立體化學特徵\n        pl.lit(\"sEH\").cast(pl.Categorical).alias(\"protein_name\"), \n        pl.lit(1).cast(pl.Int8).alias(\"binds\")                    # 因為經過高 read_count 篩選，現在可以安心給 1\n    ])\n\n    df_negative = df_ext_neg_filtered.select([\n        pl.col(\"new_structure\").alias(\"molecule_smiles\"),         # 完美保留 [Dy]\n        pl.col(\"bb1_iso\").alias(\"buildingblock1_smiles\"),         # 完美保留立體化學特徵\n        pl.lit(\"sEH\").cast(pl.Categorical).alias(\"protein_name\"), \n        pl.lit(0).cast(pl.Int8).alias(\"binds\")                    # 0\n    ])\n\n    df_ext = pl.concat([df_positive, df_negative])\n    \n    # 3. 合併與洗牌\n    print(\"正在合併與全局洗牌...\")\n    df_main_subset = df_main.select(df_ext.columns) \n    df_combined = pl.concat([df_main_subset, df_ext])\n    df_combined = df_combined.sample(fraction=1.0, shuffle=True, seed=42)\n\n    # 5. 寫出 Parquet\n    print(f\"正在寫入最終訓練檔至: {output_path}\")\n    df_combined.write_parquet(output_path, compression=\"snappy\")\n    \n    print(f\"總樣本數：{len(df_combined):,} (新增 {len(df_ext_pos_filtered):,} 筆強效正樣本)(新增 {len(df_ext_neg_filtered):,} 筆嚴格負樣本)\")\n    \n    del df_main, df_ext_raw, df_ext_pos_filtered, df_ext_neg_filtered, df_positive, df_negative, df_ext, df_combined, df_main_subset\n    gc.collect()\n\n\n\n# 執行合併 \nmerge_and_shuffle_data(\"optimized_train_data.parquet\", \"/kaggle/input/notebooks/chemdatafarmer/additional-seh-data/DNA_Labeled_Data.csv\", \"final_merged_train_data.parquet\", read_count_threshold=15)\n'''","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-19T11:35:46.153809Z","iopub.execute_input":"2026-06-19T11:35:46.154066Z","iopub.status.idle":"2026-06-19T11:35:46.178539Z","shell.execute_reply.started":"2026-06-19T11:35:46.154043Z","shell.execute_reply":"2026-06-19T11:35:46.177773Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import polars as pl\nimport numpy as np\nfrom sklearn.model_selection import StratifiedGroupKFold\n\ndef load_and_split_data(parquet_path, n_splits=5):\n    print(\"正在透過 Polars 載入數據...\")\n    df = pl.read_parquet(parquet_path)\n    \n    print(\"啟動 StratifiedGroupKFold 嚴格切分 (BB1 隔離)...\")\n    sgkf = StratifiedGroupKFold(n_splits=n_splits, shuffle=True, random_state=42)\n    \n    y = df['binds'].to_numpy()\n    groups = df['buildingblock1_smiles'].to_numpy()\n    X = np.zeros(len(y)) \n    \n    for train_idx, val_idx in sgkf.split(X, y, groups):\n        train_df = df[train_idx].to_pandas()\n        val_df = df[val_idx].to_pandas()\n        break \n        \n    print(f\"分割完成 | 訓練集: {len(train_df)} | 驗證集: {len(val_df)}\")\n    \n    # 4. 嚴格洩漏檢查 (此時還需要用到 BB1)\n    train_groups = set(train_df['buildingblock1_smiles'].unique())\n    val_groups = set(val_df['buildingblock1_smiles'].unique())\n    leakage = train_groups.intersection(val_groups)\n    \n    if len(leakage) == 0:\n        print(\"洩漏檢查完美通過：訓練集與驗證集完全沒有重複的化學骨架 (BB1)！\")\n    else:\n        print(f\"警告：發現洩漏！重疊的 BB1 數量: {len(leakage)}\")\n    \n    # ==========================================\n    # 卸載 BB1 欄位\n    # ==========================================\n    print(\"🧹 正在清理暫存欄位，無縫對齊下游 Dataset...\")\n    train_df = train_df.drop(columns=['buildingblock1_smiles'])\n    val_df = val_df.drop(columns=['buildingblock1_smiles'])\n    \n    return train_df, val_df\n\n\nprint(\"開始執行Scaffold Split 數據載入...\")\ntrain_df, val_df = load_and_split_data(\"/kaggle/input/notebooks/t8101349/predict-new-medicines-with-belka-rusessemble-data4/final_merged_train_data.parquet\", n_splits=5) ###\n\nimport gc\ngc.collect()\n\nprint(f\"分割完成！訓練集大小: {len(train_df)}, 驗證集大小: {len(val_df)}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-19T11:35:46.179554Z","iopub.execute_input":"2026-06-19T11:35:46.179818Z","iopub.status.idle":"2026-06-19T11:38:51.609323Z","shell.execute_reply.started":"2026-06-19T11:35:46.179789Z","shell.execute_reply":"2026-06-19T11:38:51.608199Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torchvision\nimport torch.nn as nn\n\nclass FocalLoss(nn.Module):\n    def __init__(self, alpha=0.25, gamma=3.0, label_smoothing=0.005):\n        super().__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n        self.label_smoothing = label_smoothing\n        self.bce = nn.BCEWithLogitsLoss(reduction='none')\n\n    def forward(self, inputs, targets):\n        targets = targets.view(-1, 1).float()\n        inputs = inputs.view(-1, 1).float()\n        \n        if self.label_smoothing > 0:\n            smoothed_targets = targets * (1.0 - self.label_smoothing) + 0.5 * self.label_smoothing\n        else:\n            smoothed_targets = targets\n\n        bce_loss = self.bce(inputs, smoothed_targets)\n        probas = torch.sigmoid(inputs)\n        p_t = probas * targets + (1 - probas) * (1 - targets)\n        focal_weight = (1 - p_t) ** self.gamma\n        alpha_t = self.alpha * targets + (1 - self.alpha) * (1 - targets)\n        \n        loss = alpha_t * focal_weight * bce_loss\n        return loss.mean()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-19T11:38:51.610393Z","iopub.execute_input":"2026-06-19T11:38:51.61065Z","iopub.status.idle":"2026-06-19T11:38:52.427797Z","shell.execute_reply.started":"2026-06-19T11:38:51.610618Z","shell.execute_reply":"2026-06-19T11:38:52.42709Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport pandas as pd\nimport polars as pl\nimport numpy as np\nfrom tqdm import tqdm\nfrom transformers import AutoTokenizer, AutoModel\nfrom sklearn.metrics import average_precision_score\n\n# --- 1. 配置與蛋白質預快取 ---\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nPROT_MODEL_NAME = \"facebook/esm2_t12_35M_UR50D\"\n\n# 預先載入 Tokenizer\nprot_tokenizer = AutoTokenizer.from_pretrained(PROT_MODEL_NAME)\nprot_model = AutoModel.from_pretrained(PROT_MODEL_NAME).to(device).eval()\n\ndef precompute_prot_embeddings():\n    \"\"\"預先計算 BELKA 的三種蛋白質 Embedding\"\"\"\n    prot_sequences = {\n        'BRD4': \"MSAESGPGTRLRNLPVMGDGLETSQMSTTQAQAQPQPANAASTNPPPPETSNPNKPKRQTNQLQYLLRVVLKTLWKHQFAWPFQQPVDAVKLNLPDYYKIIKTPMDMGTIKKRLENNYYWNAQECIQDFNTMFTNCYIYNKPGDDIVLMAEALEKLFLQKINELPTEETEIMIVQAKGRGRGRKETGTAKPGVSTVPNTTQASTPPQTQTPQPNPPPVQATPHPFPAVTPDLIVQTPVMTVVPPQPLQTPPPVPPQPQPPPAPAPQPVQSHPPIIAATPQPVKTKKGVKRKADTTTPTTIDPIHEPPSLPPEPKTTKLGQRRESSRPVKPPKKDVPDSQQHPAPEKSSKVSEQLKCCSGILKEMFAKKHAAYAWPFYKPVDVEALGLHDYCDIIKHPMDMSTIKSKLEAREYRDAQEFGADVRLMFSNCYKYNPPDHEVVAMARKLQDVFEMRFAKMPDEPEEPVVAVSSPAVPPPTKVVAPPSSSDSSSDSSSDSDSSTDDSEEERAQRLAELQEQLKAVHEQLAALSQPQQNKPKKKKKKKKKKKK\", \n        'HSA': \"MKWVTFISLLFLFSSAYSRGVFRRDAHKSEVAHRFKDLGEENFKALVLIAFAQYLQQCPFEDHVKLVNEVTEFAKTCVADESAENCDKSLHTLFGDKLCTVATLRETYGEMADCCAKQEPERNECFLQHKDDNPNLPRLVRPEVDVMCTAFHDNEETFLKKYLYEIARRHPYFYAPELLFFAKRYKAAFTECCQAADKAACLLPKLDELRDEGKASSAKQRLKCASLQKFGERAFKAWAVARLSQRFPKAEFAEVSKLVTDLTKVHTECCHGDLLECADDRADLAKYICENQDSISSKLKECCEKPLLEKSHCIAEVENDEMPADLPSLAADFVESKDVCKNYAEAKDVFLGMFLYEYARRHPDYSVVLLLRLAKTYETTLEKCCAAADPHECYAKVFDEFKPLVEEPQNLIKQNCELFEQLGEYKFQNALLVRYTKKVPQVSTPTLVEVSRNLGKVGSKCCKHPEAKRMPCAEDYLSVVLNQLCVLHEKTPVSDRVTKCCTESLVNRRPCFSALEVDETYVPKEFNAETFTFHADICTLSEKERQIKKQTALVELVKHKPKATKEQLKAVMDDFAAFVEKCCKADDKETCFAEEGKKLVAASQAALGL\",\n        'sEH': \"MTLRAAVFDLDGVLALPAVFGVLGRTEEALALPRGLLNDAFQKGGPEGATTRLMKGEITLSQWIPLMEENCRKCSETAKVCLPKNFSIKEIFDKAISARKINRPMLQAALMLRKKGFTTAILTNTWLDDRAERDGLAQLMCELKMHFDFLIESCQVGMVKPEPQIYKFLLDTLKASPSEVVFLDDIGANLKPARDLGMVTILVQDTDTALKELEKVTGIQLLNTPAPLPTSCNPSDMSHGYVTVKPRVRLHFVELGSGPAVCLCHGFPESWYSWRYQIPALAQAGYRVLAMDMKGYGESSAPPEIEEYCMEVLCKEMVTFLDKLGLSQAVFIGHDWGGMLVWYMALFYPERVRAVASLNTPFIPANPNMSPLESIKANPVFDYQLYFQEPGVAEAELEQNLSRTFKSLFRASDESVLSMHKVCEAGGLFVNSPEEPSLSRMVTEEEIQFYVQQFKKSGFRGPLNWYRNMERNWKWACKSLGRKILIPALMVTAEKDFVLVPQMSQHMEDWIPHLKRGHIEDCGHWTQMDKPTEVNQILIKWLDSDARNPPVVSKM\"\n    }\n    prot_map = {}\n    for name, seq in prot_sequences.items():\n        inputs = prot_tokenizer(seq, return_tensors=\"pt\", truncation=True, max_length=1024).to(device)\n        with torch.no_grad():\n            outputs = prot_model(**inputs)\n            # 取 Mean Pooling\n            prot_map[name] = outputs.last_hidden_state.mean(dim=1).cpu().numpy().squeeze()\n    return prot_map\n\n# 預處理好的 map\nprot_map = precompute_prot_embeddings()\n\ndel prot_model, prot_tokenizer\ntorch.cuda.empty_cache()\n\n\n# --- 2. 數據集與模型定義 ---\n\nclass BELKADynamicDataset(Dataset):\n    def __init__(self, df, prot_map):\n        self.smiles = df['molecule_smiles'].values\n        self.prot_names = df['protein_name'].values\n        self.labels = df['binds'].values\n        self.prot_map = prot_map\n\n    def __len__(self):\n        return len(self.labels)\n\n    def __getitem__(self, idx):\n        # 1. 抓取蛋白質特徵\n        prot_vector = self.prot_map[self.prot_names[idx]]\n        if isinstance(prot_vector, torch.Tensor):\n            prot_vector = prot_vector.clone().detach().float()\n        else:\n            prot_vector = torch.tensor(prot_vector, dtype=torch.float32)\n\n        return {\n            \"smiles\": self.smiles[idx], \n            \"prot_emb\": prot_vector,\n            \"label\": torch.tensor([self.labels[idx]], dtype=torch.float32)\n        }\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-19T11:38:52.428845Z","iopub.execute_input":"2026-06-19T11:38:52.429178Z","iopub.status.idle":"2026-06-19T11:39:01.760101Z","shell.execute_reply.started":"2026-06-19T11:38:52.42915Z","shell.execute_reply":"2026-06-19T11:39:01.759118Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom transformers import AutoModel\n\n\nclass PeriLNGatedLayer(nn.Module):\n    \"\"\"Peri-LN + GEGLU + ReZero-style 殘差閘門\"\"\"\n    def __init__(self, d_model, nhead, dim_feedforward, dropout=0.1):\n        super().__init__()\n        self.self_attn = nn.MultiheadAttention(d_model, nhead, dropout=dropout, batch_first=True)\n        self.linear1 = nn.Linear(d_model, dim_feedforward * 2)\n        self.linear2 = nn.Linear(dim_feedforward, d_model)\n\n        self.norm_attn_in = nn.LayerNorm(d_model)\n        self.norm_attn_out = nn.LayerNorm(d_model)\n        self.norm_ffn_in = nn.LayerNorm(d_model)\n        self.norm_ffn_out = nn.LayerNorm(d_model)\n\n        self.dropout = nn.Dropout(dropout)\n        self.alpha_attn = nn.Parameter(torch.zeros(1))\n        self.alpha_ffn = nn.Parameter(torch.zeros(1))\n\n    def forward(self, src, src_key_padding_mask=None):\n        nx = self.norm_attn_in(src)\n        attn_out, _ = self.self_attn(nx, nx, nx, key_padding_mask=src_key_padding_mask, need_weights=False)\n        src = src + self.alpha_attn * self.dropout(self.norm_attn_out(attn_out))\n\n        nx = self.norm_ffn_in(src)\n        a, gate = self.linear1(nx).chunk(2, dim=-1)\n        ffn = self.linear2(a * F.gelu(gate))\n        src = src + self.alpha_ffn * self.dropout(self.norm_ffn_out(ffn))\n        return src\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-19T11:39:01.761229Z","iopub.execute_input":"2026-06-19T11:39:01.761474Z","iopub.status.idle":"2026-06-19T11:39:01.768428Z","shell.execute_reply.started":"2026-06-19T11:39:01.761453Z","shell.execute_reply":"2026-06-19T11:39:01.767838Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import copy\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom transformers import AutoModel\n\n\nclass CrossAttnFusionLayer(nn.Module):\n    \"\"\"\n    Protein queries 主動詢問 molecule tokens。\n    Q = protein queries , K=V = molecule tokens \n    Peri-LN + GEGLU + ReZero α gate, 沿用上一輪的穩定組合。\n    \"\"\"\n    def __init__(self, d_model, nhead, dim_feedforward, dropout=0.1):\n        super().__init__()\n        self.cross_attn = nn.MultiheadAttention(d_model, nhead, dropout=dropout, batch_first=True)\n        self.linear1 = nn.Linear(d_model, dim_feedforward * 2)\n        self.linear2 = nn.Linear(dim_feedforward, d_model)\n\n        self.norm_q_in   = nn.LayerNorm(d_model)\n        self.norm_kv_in  = nn.LayerNorm(d_model)\n        self.norm_attn_out = nn.LayerNorm(d_model)\n        self.norm_ffn_in   = nn.LayerNorm(d_model)\n        self.norm_ffn_out  = nn.LayerNorm(d_model)\n\n        self.dropout = nn.Dropout(dropout)\n        self.alpha_attn = nn.Parameter(torch.zeros(1))\n        self.alpha_ffn  = nn.Parameter(torch.zeros(1))\n\n    def forward(self, q, kv, kv_pad_mask=None):\n        # q: (B, Q, D), kv: (B, L, D), kv_pad_mask: (B, L)  True=PAD\n        nq, nkv = self.norm_q_in(q), self.norm_kv_in(kv)\n        attn_out, _ = self.cross_attn(nq, nkv, nkv,\n                                      key_padding_mask=kv_pad_mask,\n                                      need_weights=False)\n        q = q + self.alpha_attn * self.dropout(self.norm_attn_out(attn_out))\n\n        nq = self.norm_ffn_in(q)\n        a, gate = self.linear1(nq).chunk(2, dim=-1)\n        ffn = self.linear2(a * F.gelu(gate))\n        q = q + self.alpha_ffn * self.dropout(self.norm_ffn_out(ffn))\n        return q\n\n\nclass QuerySelfAttnLayer(nn.Module):\n    \"\"\"讓 protein queries 之間互通有無 (各自吸收的分子片段做整合)\"\"\"\n    def __init__(self, d_model, nhead, dim_feedforward, dropout=0.1):\n        super().__init__()\n        self.self_attn = nn.MultiheadAttention(d_model, nhead, dropout=dropout, batch_first=True)\n        self.linear1 = nn.Linear(d_model, dim_feedforward * 2)\n        self.linear2 = nn.Linear(dim_feedforward, d_model)\n        self.norm_attn_in = nn.LayerNorm(d_model)\n        self.norm_attn_out = nn.LayerNorm(d_model)\n        self.norm_ffn_in   = nn.LayerNorm(d_model)\n        self.norm_ffn_out  = nn.LayerNorm(d_model)\n        self.dropout = nn.Dropout(dropout)\n        self.alpha_attn = nn.Parameter(torch.zeros(1))\n        self.alpha_ffn  = nn.Parameter(torch.zeros(1))\n\n    def forward(self, x):\n        nx = self.norm_attn_in(x)\n        attn_out, _ = self.self_attn(nx, nx, nx, need_weights=False)\n        x = x + self.alpha_attn * self.dropout(self.norm_attn_out(attn_out))\n\n        nx = self.norm_ffn_in(x)\n        a, gate = self.linear1(nx).chunk(2, dim=-1)\n        ffn = self.linear2(a * F.gelu(gate))\n        x = x + self.alpha_ffn * self.dropout(self.norm_ffn_out(ffn))\n        return x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-19T11:39:01.770674Z","iopub.execute_input":"2026-06-19T11:39:01.770912Z","iopub.status.idle":"2026-06-19T11:39:01.797084Z","shell.execute_reply.started":"2026-06-19T11:39:01.770891Z","shell.execute_reply":"2026-06-19T11:39:01.796111Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ChemBERTaCrossAttnFusion(nn.Module):\n    \"\"\"\n    ChemBERTa (永久凍結) → mol features\n    Protein 條件化的 K 個 query tokens → cross-attend mol → 互相 self-attend\n    讀出: K 個 query 的 attention pooling → head\n    \"\"\"\n    def __init__(\n        self,\n        chem_model_name=\"DeepChem/ChemBERTa-77M-MLM\",\n        d_model=384,\n        prot_dim=480,\n        num_queries=8,           # 平行的蛋白質「探針」數量\n        num_fusion_layers=4,\n        nhead=6,\n        dropout=0.1,\n    ):\n        super().__init__()\n\n        # --- 1. ChemBERTa 永久凍結 (eval 模式 + no_grad 推論) ---------------\n        self.chemberta = AutoModel.from_pretrained(chem_model_name)\n        for p in self.chemberta.parameters():\n            p.requires_grad = False\n        self.chemberta.eval()\n        chem_dim = self.chemberta.config.hidden_size\n\n        # --- 2. 分子特徵對齊 (常駐的小投影 + LN) --------------------------\n        self.mol_proj = nn.Sequential(\n            nn.LayerNorm(chem_dim),\n            nn.Linear(chem_dim, d_model),\n        )\n\n        # --- 3. 蛋白質投影 -------------------------------------------------\n        self.prot_proj = nn.Sequential(\n            nn.LayerNorm(prot_dim),\n            nn.Linear(prot_dim, d_model),\n            nn.GELU(),\n            nn.Linear(d_model, d_model),\n        )\n\n        # --- 4. 可學習的 query seed (K 個探針) + 加上蛋白質條件 ------------\n        self.num_queries = num_queries\n        self.query_seeds = nn.Parameter(torch.randn(num_queries, d_model) * 0.02)\n        self.query_pos   = nn.Embedding(num_queries, d_model)\n\n        # --- 5. Fusion stack: Cross-Attn → Query-SelfAttn ------------------\n        self.cross_layers = nn.ModuleList([\n            CrossAttnFusionLayer(d_model, nhead, d_model * 4, dropout)\n            for _ in range(num_fusion_layers)\n        ])\n        self.query_layers = nn.ModuleList([\n            QuerySelfAttnLayer(d_model, nhead, d_model * 4, dropout)\n            for _ in range(num_fusion_layers)\n        ])\n        self.final_norm = nn.LayerNorm(d_model)\n\n        # --- 6. Attention pooling 聚合 K 個 query ---------------------------\n        self.pool_query = nn.Parameter(torch.randn(1, 1, d_model) * 0.02)\n        self.pool_attn  = nn.MultiheadAttention(d_model, nhead, dropout=dropout, batch_first=True)\n\n        # --- 7. 預測頭 ------------------------------------------------------\n        self.head = nn.Sequential(\n            nn.Linear(d_model, 256),\n            nn.LayerNorm(256),\n            nn.GELU(),\n            nn.Dropout(dropout * 3),\n            nn.Linear(256, 128),\n            nn.LayerNorm(128),\n            nn.GELU(),\n            nn.Dropout(dropout * 2),\n            nn.Linear(128, 1),\n        )\n\n    def train(self, mode=True):\n        \"\"\"強制 ChemBERTa 永遠在 eval 模式 (避免 dropout 污染特徵)\"\"\"\n        super().train(mode)\n        self.chemberta.eval()\n        return self\n\n    def forward(self, input_ids, attention_mask, prot_emb):\n        B = input_ids.size(0)\n\n        # 1) ChemBERTa 凍結前向 (no_grad 省顯存與算力)\n        with torch.no_grad():\n            chem_out = self.chemberta(\n                input_ids=input_ids, attention_mask=attention_mask\n            ).last_hidden_state                         # (B, L, chem_dim)\n        mol = self.mol_proj(chem_out)                   # (B, L, D)\n        mol_pad = (attention_mask == 0)                 # (B, L) True=PAD\n\n        # 2) 構造蛋白質條件化的 query: seed + position + prot vector\n        prot_vec = self.prot_proj(prot_emb).unsqueeze(1)                  # (B, 1, D)\n        seeds = self.query_seeds.unsqueeze(0).expand(B, -1, -1)           # (B, K, D)\n        qpos  = self.query_pos.weight.unsqueeze(0).expand(B, -1, -1)      # (B, K, D)\n        q = seeds + qpos + prot_vec                                       # 廣播加總 (B, K, D)\n\n        # 3) Cross-Attn (q 看 mol) → Query Self-Attn (q 之間互通)\n        for ca, sa in zip(self.cross_layers, self.query_layers):\n            q = ca(q, mol, kv_pad_mask=mol_pad)\n            q = sa(q)\n        q = self.final_norm(q)                                            # (B, K, D)\n\n        # 4) Attention pooling: 用 1 個 pool token 對 K 個 query 做 weighted sum\n        pq = self.pool_query.expand(B, -1, -1)                            # (B, 1, D)\n        fused, _ = self.pool_attn(pq, q, q, need_weights=False)           # (B, 1, D)\n        fused = fused.squeeze(1)                                          # (B, D)\n\n        return self.head(fused)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-19T11:39:01.798278Z","iopub.execute_input":"2026-06-19T11:39:01.798623Z","iopub.status.idle":"2026-06-19T11:39:01.82363Z","shell.execute_reply.started":"2026-06-19T11:39:01.798588Z","shell.execute_reply":"2026-06-19T11:39:01.82296Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ModelEMA:\n    \"\"\"\n    只追蹤可訓練參數的 shadow weights。\n    凍結的 ChemBERTa 不複製,省下 77M 參數的記憶體。\n    \"\"\"\n    def __init__(self, model, decay=0.999, warmup_steps=2000):\n        self.decay = decay\n        self.warmup_steps = warmup_steps\n        self.num_updates = 0\n\n        base = model.module if hasattr(model, \"module\") else model\n        self.shadow = {}\n        for n, p in base.named_parameters():\n            if p.requires_grad:\n                self.shadow[n] = p.data.detach().clone()\n\n    @torch.no_grad()\n    def update(self, model):\n        \"\"\"每個 optimizer.step() 後呼叫\"\"\"\n        self.num_updates += 1\n        # 早期用較低 decay (避免初始隨機權重污染太久),後期才到 0.999\n        d = min(self.decay, (1 + self.num_updates) / (10 + self.num_updates))\n\n        base = model.module if hasattr(model, \"module\") else model\n        for n, p in base.named_parameters():\n            if n in self.shadow:\n                self.shadow[n].mul_(d).add_(p.data, alpha=1 - d)\n\n    @torch.no_grad()\n    def swap_in(self, model):\n        \"\"\"驗證/推論前: 把 EMA 權重塞進 model,回傳原權重備份\"\"\"\n        base = model.module if hasattr(model, \"module\") else model\n        backup = {}\n        for n, p in base.named_parameters():\n            if n in self.shadow:\n                backup[n] = p.data.clone()\n                p.data.copy_(self.shadow[n])\n        return backup\n\n    @torch.no_grad()\n    def swap_out(self, model, backup):\n        \"\"\"驗證完畢: 還原訓練中的權重\"\"\"\n        base = model.module if hasattr(model, \"module\") else model\n        for n, p in base.named_parameters():\n            if n in backup:\n                p.data.copy_(backup[n])\n\n    def state_dict(self):\n        return {\"shadow\": self.shadow, \"num_updates\": self.num_updates, \"decay\": self.decay}\n\n    def load_state_dict(self, sd):\n        self.shadow = sd[\"shadow\"]\n        self.num_updates = sd[\"num_updates\"]\n        self.decay = sd[\"decay\"]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-19T11:39:01.824834Z","iopub.execute_input":"2026-06-19T11:39:01.825086Z","iopub.status.idle":"2026-06-19T11:39:01.841613Z","shell.execute_reply.started":"2026-06-19T11:39:01.825064Z","shell.execute_reply":"2026-06-19T11:39:01.840974Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import time\nimport torch\nfrom torch.amp import autocast, GradScaler\nfrom sklearn.metrics import average_precision_score\nfrom tqdm import tqdm\n\ndef train_and_validate_with_ema(\n    model, train_loader, val_loader, optimizer, criterion, device,\n    epochs=10, time_limit_hours=9.0, max_lr=5e-4,\n    save_name=\"best_belka_model.pth\", scheduler_type='onecycle',\n    ema_decay=0.999, ema_warmup_steps=2000,\n):\n    # ---- 動態切換 Scheduler ----\n    if scheduler_type == 'onecycle':\n        scheduler = torch.optim.lr_scheduler.OneCycleLR(\n            optimizer, max_lr=max_lr, steps_per_epoch=len(train_loader),\n            epochs=epochs, pct_start=0.1, anneal_strategy='cos'\n        )\n    else:\n        scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n            optimizer, T_max=epochs * len(train_loader), eta_min=1e-7\n        )\n\n    scaler = torch.amp.GradScaler('cuda')\n    ema = ModelEMA(model, decay=ema_decay, warmup_steps=ema_warmup_steps)  # ← 加上\n    best_ap, start_time = 0.0, time.time()\n\n    for epoch in range(epochs):\n        # ===== Train =====\n        model.train()\n        train_loss = 0\n        pbar = tqdm(train_loader, desc=f\"Epoch {epoch+1}/{epochs} [Train]\")\n        for batch in pbar:\n            ids   = batch[\"input_ids\"].to(device, non_blocking=True)\n            mask  = batch[\"attention_mask\"].to(device, non_blocking=True)\n            prot  = batch[\"prot_emb\"].to(device, non_blocking=True)\n            labels= batch[\"labels\"].to(device, non_blocking=True)\n\n            optimizer.zero_grad(set_to_none=True)\n            with torch.amp.autocast('cuda'):\n                logits = model(input_ids=ids, attention_mask=mask, prot_emb=prot)\n                loss = criterion(logits.float(), labels.view_as(logits).float())\n\n            scaler.scale(loss).backward()\n            scaler.unscale_(optimizer)\n            torch.nn.utils.clip_grad_norm_(\n                [p for p in model.parameters() if p.requires_grad], max_norm=1.0\n            )\n            scaler.step(optimizer)\n            scaler.update()\n            scheduler.step()\n\n            ema.update(model)                                # ← 每個 step 更新 EMA\n\n            train_loss += loss.item()\n            pbar.set_postfix(loss=f\"{loss.item():.4f}\",\n                             lr=f\"{scheduler.get_last_lr()[0]:.6e}\")\n\n        # ===== Val (用 EMA 權重) =====\n        backup = ema.swap_in(model)                          # ← 切換到 EMA\n        model.eval()\n        all_preds, all_labels, val_loss = [], [], 0\n        with torch.no_grad():\n            for batch in tqdm(val_loader, desc=f\"Epoch {epoch+1}/{epochs} [Val-EMA]\"):\n                ids  = batch[\"input_ids\"].to(device, non_blocking=True)\n                mask = batch[\"attention_mask\"].to(device, non_blocking=True)\n                prot = batch[\"prot_emb\"].to(device, non_blocking=True)\n                labels = batch[\"labels\"].to(device, non_blocking=True)\n                with torch.amp.autocast('cuda'):\n                    logits = model(input_ids=ids, attention_mask=mask, prot_emb=prot)\n                    v_loss = criterion(logits.float(), labels.view_as(logits).float())\n                val_loss += v_loss.item()\n                all_preds.extend(torch.sigmoid(logits.float()).cpu().numpy())\n                all_labels.extend(labels.cpu().numpy())\n        ema.swap_out(model, backup)                          # ← 還原訓練權重\n\n        current_ap = average_precision_score(all_labels, all_preds)\n        print(f\"\\nEpoch {epoch+1}: Train {train_loss/len(train_loader):.4f} | \"\n              f\"Val {val_loss/len(val_loader):.4f} | AP (EMA) {current_ap:.4f}\")\n\n        # ===== Save EMA weights as best =====\n        if current_ap > best_ap:\n            best_ap = current_ap\n            # 直接存 EMA shadow\n            base = model.module if hasattr(model, \"module\") else model\n            backup = ema.swap_in(model)\n            torch.save({k: v.detach().cpu() for k, v in base.state_dict().items()\n                        if k.split('.')[0] != 'chemberta'},   \n                       save_name)\n            ema.swap_out(model, backup)\n            print(f\"  ✓ Best EMA model saved: {save_name} (AP={best_ap:.4f})\")\n\n        if (time.time() - start_time) / 3600 > time_limit_hours:\n            print(\"時間到,停止訓練\")\n            break\n\n    return best_ap","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-19T11:39:01.842588Z","iopub.execute_input":"2026-06-19T11:39:01.843188Z","iopub.status.idle":"2026-06-19T11:39:01.866986Z","shell.execute_reply.started":"2026-06-19T11:39:01.843102Z","shell.execute_reply":"2026-06-19T11:39:01.866433Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport torch\nimport polars as pl\nfrom sklearn.model_selection import train_test_split # 應該從 sklearn 導入\nfrom torch.utils.data import DataLoader, Dataset      # torch 只負責 Data 加載\nimport torch.nn as nn\nfrom transformers import AutoTokenizer\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Using device: {device}\")\n\n# === ChemBERTa 專用 tokenizer (關鍵: 不是 Model 1 的字元 tokenizer) ===\nchem_tok = AutoTokenizer.from_pretrained(\"DeepChem/ChemBERTa-77M-MLM\")\n\ndef make_chemberta_collate(tokenizer, max_length=160):\n    def collate(batch):\n        smiles = [b[\"smiles\"] for b in batch]\n        enc = tokenizer(smiles, padding=True, truncation=True,\n                        max_length=max_length, return_tensors=\"pt\")\n        return {\n            \"input_ids\":      enc[\"input_ids\"],\n            \"attention_mask\": enc[\"attention_mask\"],\n            \"prot_emb\":       torch.stack([b[\"prot_emb\"] for b in batch]),\n            \"labels\":         torch.stack([b[\"label\"]    for b in batch]),\n        }\n    return collate\n\nBATCH_SIZE = 3072\ncustom_collate = make_chemberta_collate(chem_tok, max_length=160)\n\n# 1) 先把 DataFrame 包成 Dataset —— 這兩行才是你之前漏掉的\ntrain_dataset = BELKADynamicDataset(train_df, prot_map)\nval_dataset   = BELKADynamicDataset(val_df,   prot_map)\n\n# 2) 再用 DataLoader 包 Dataset \ntrain_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True,\n                          num_workers=4, pin_memory=True, drop_last=True, collate_fn=custom_collate)\nval_loader   = DataLoader(val_dataset,   batch_size=BATCH_SIZE, shuffle=False,\n                          num_workers=4, pin_memory=True, collate_fn=custom_collate)\n\n# === 單階段訓練 ===\nprint(\"\\n\" + \"=\"*40)\nprint(\"單階段訓練 (ChemBERTa 全程凍結)\")\nprint(\"=\"*40)\n\nmodel = ChemBERTaCrossAttnFusion(\n    chem_model_name=\"DeepChem/ChemBERTa-77M-MLM\",\n    d_model=384, prot_dim=480,\n    num_queries=8, num_fusion_layers=4, nhead=6, dropout=0.1,\n).to(device)\n\nif torch.cuda.device_count() > 1:\n    print(f\"偵測到 {torch.cuda.device_count()} 張 GPU,啟動 DataParallel\")\n    model = nn.DataParallel(model)\n\ntrainable = [p for p in model.parameters() if p.requires_grad]\nprint(f\"可訓練參數: {sum(p.numel() for p in trainable)/1e6:.2f}M \"\n      f\"(ChemBERTa 凍結,不計入)\")\n\noptimizer = torch.optim.AdamW(trainable, lr=3e-4, weight_decay=1e-2)\ncriterion = FocalLoss(alpha=0.25, gamma=2.0, label_smoothing=0.05)\n\nbest_ap = train_and_validate_with_ema(\n    model, train_loader, val_loader, optimizer, criterion, device,\n    epochs=12, max_lr=3e-4, scheduler_type='onecycle',\n    save_name=\"ema_crossattn_final.pth\",\n    ema_decay=0.999, ema_warmup_steps=2000,\n)\n\nprint(f\"\\n✅ 訓練完成,最佳 EMA AP: {best_ap:.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-19T11:39:01.86759Z","iopub.execute_input":"2026-06-19T11:39:01.867753Z","iopub.status.idle":"2026-06-19T11:39:02.773758Z","shell.execute_reply.started":"2026-06-19T11:39:01.867734Z","shell.execute_reply":"2026-06-19T11:39:02.772383Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport gc\nimport torch\nimport numpy as np\nimport polars as pl\nimport pandas as pd\nfrom tqdm import tqdm\nfrom torch.utils.data import DataLoader, Dataset\nfrom transformers import AutoTokenizer\n\n# ==========================================\n# 0. ChemBERTa tokenizer (與訓練同一個!) + 對齊 max_length\n# ==========================================\nchem_tok = AutoTokenizer.from_pretrained(\"DeepChem/ChemBERTa-77M-MLM\")\nMAX_LEN = 160 \n\n# ==========================================\n# 1. 測試集 Dataset\n# ==========================================\nclass BELKATestDataset(Dataset):\n    def __init__(self, df, prot_map):\n        self.ids        = df['id'].to_numpy()\n        self.smiles     = df['molecule_smiles'].to_numpy()\n        self.prot_names = df['protein_name'].to_numpy()\n        self.prot_map   = prot_map\n\n    def __len__(self):\n        return len(self.smiles)\n\n    def __getitem__(self, idx):\n        prot_vector = self.prot_map[self.prot_names[idx]]\n        if isinstance(prot_vector, torch.Tensor):\n            prot_vector = prot_vector.clone().detach().float()\n        else:\n            prot_vector = torch.tensor(prot_vector, dtype=torch.float32)\n        return {\n            \"id\":       self.ids[idx],\n            \"smiles\":   self.smiles[idx],\n            \"prot_emb\": prot_vector,\n        }\n\n# ==========================================\n# 2. Collate (與訓練完全對齊: 產生 input_ids + attention_mask + tensor)\n# ==========================================\ndef make_chemberta_test_collate(tokenizer, max_length=MAX_LEN):\n    def collate(batch):\n        enc = tokenizer([b[\"smiles\"] for b in batch],\n                        padding=True, truncation=True,\n                        max_length=max_length, return_tensors=\"pt\")   # ← 關鍵: return_tensors + attention_mask\n        return {\n            \"ids\":            [b[\"id\"] for b in batch],\n            \"input_ids\":      enc[\"input_ids\"],\n            \"attention_mask\": enc[\"attention_mask\"],\n            \"prot_emb\":       torch.stack([b[\"prot_emb\"] for b in batch]),\n        }\n    return collate\n\n# ==========================================\n# 3. 推論 + 產生提交\n# ==========================================\ndef generate_submission(model, test_file, tokenizer, prot_map, device,\n                        output_name=\"submission.csv\", batch_size=4096):\n    model.eval()\n\n    print(\"使用 Polars 讀取測試集中...\")\n    test_df = pl.read_parquet(test_file, columns=[\"id\", \"molecule_smiles\", \"protein_name\"])\n\n    test_loader = DataLoader(\n        BELKATestDataset(test_df, prot_map),\n        batch_size=batch_size, shuffle=False,        # 不洗牌,維持 id 順序\n        num_workers=4, pin_memory=True,\n        collate_fn=make_chemberta_test_collate(tokenizer),\n    )\n\n    all_ids, pred_chunks = [], []\n    print(\"開始推論...\")\n    with torch.no_grad():\n        for batch in tqdm(test_loader, desc=\"Predicting\"):\n            with torch.amp.autocast('cuda'):\n                outputs = model(                                       # ← 正確的 forward 呼叫\n                    input_ids=batch[\"input_ids\"].to(device, non_blocking=True),\n                    attention_mask=batch[\"attention_mask\"].to(device, non_blocking=True),\n                    prot_emb=batch[\"prot_emb\"].to(device, non_blocking=True),\n                )\n            probs = torch.sigmoid(outputs.float()).cpu().numpy().flatten()\n            all_ids.extend(batch[\"ids\"])\n            pred_chunks.append(probs)                                  # 收集 array,最後一次合併\n\n    all_preds = np.concatenate(pred_chunks)                            # 比逐元素 extend 快得多\n\n    print(f\"寫入 {output_name} ...\")\n    pd.DataFrame({\"id\": all_ids, \"binds\": all_preds}).to_csv(output_name, index=False)\n    print(f\"✅ 提交檔完成: {len(all_preds):,} 筆\")\n\n    del test_df, all_ids, all_preds, pred_chunks\n    gc.collect()\n\n# ==========================================\n# 4. 載入模型\n# ==========================================\nmodel = ChemBERTaCrossAttnFusion(\n    chem_model_name=\"DeepChem/ChemBERTa-77M-MLM\",\n    d_model=384, prot_dim=480,\n    num_queries=8, num_fusion_layers=4, nhead=6, dropout=0.1,\n).to(device)\n\nckpt = torch.load(\"ema_crossattn_final.pth\", map_location=device) \nmissing, unexpected = model.load_state_dict(ckpt, strict=False)\n\n# 檢查: 融合層一定要真的載到,否則就是隨機權重亂數提交\nassert len(unexpected) == 0, f\"有未預期的 key: {unexpected[:5]}\"\nnon_chem_missing = [k for k in missing if not k.startswith(\"chemberta\")]\nassert len(non_chem_missing) == 0, f\"融合層沒載到! {non_chem_missing[:5]}\"\nprint(f\"✅ 權重載入正確 (僅 chemberta {len(missing)} keys 從預訓練還原)\")\nmodel.eval()\n\n# (可選) 雙卡推論加速 —— 一定要在 load_state_dict「之後」才包,否則 key 前綴對不上\n# if torch.cuda.device_count() > 1:\n#     import torch.nn as nn\n#     model = nn.DataParallel(model)\n\n# ==========================================\n# 5. 執行\n# ==========================================\nTEST_FILE       = \"/kaggle/input/competitions/leash-BELKA/test.parquet\"\nSAMPLE_SUB_FILE = \"/kaggle/input/competitions/leash-BELKA/sample_submission.csv\"\n\nif os.path.exists(TEST_FILE):\n    generate_submission(model, TEST_FILE, chem_tok, prot_map, device)  # ← 傳 chem_tok,不是 mol_tokenizer\nelse:\n    print(f\"❌ 找不到測試檔: {TEST_FILE}\")\n\n# ==========================================\n# 6. 提交健檢\n# ==========================================\nif os.path.exists(\"submission.csv\") and os.path.exists(SAMPLE_SUB_FILE):\n    sub_count    = pl.scan_csv(\"submission.csv\").select(pl.len()).collect().item()\n    sample_count = pl.scan_csv(SAMPLE_SUB_FILE).select(pl.len()).collect().item()\n    print(f\"行數: 提交 {sub_count:,} / 範本 {sample_count:,} → \"\n          f\"{'✅ 一致' if sub_count == sample_count else '❌ 不一致'}\")\n\n    sub = pl.read_csv(\"submission.csv\")\n    print(\"id 嚴格遞增:\", sub['id'].is_sorted())\n    print(\"binds 範圍:\", round(sub['binds'].min(), 5), \"~\", round(sub['binds'].max(), 5))  # 應落在 0~1\n\n    # 各蛋白平均預測機率 —— 檢查模型有沒有退化成常數輸出\n    test_meta = pl.read_parquet(TEST_FILE, columns=[\"id\", \"protein_name\"])\n    print(sub.join(test_meta, on=\"id\").group_by(\"protein_name\").agg(pl.col(\"binds\").mean()))\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-19T11:39:02.774482Z","iopub.status.idle":"2026-06-19T11:39:02.774788Z","shell.execute_reply.started":"2026-06-19T11:39:02.774658Z","shell.execute_reply":"2026-06-19T11:39:02.774677Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import polars as pl\nsub = pl.read_csv(\"submission.csv\")\nprint(\"ID 是否嚴格遞增且無遺漏：\", sub['id'].is_sorted())\n\n# 把 submission 的結果跟 test.parquet 裡的 protein_name 對接起來看平均機率\ntest_df = pl.read_parquet(\"/kaggle/input/competitions/leash-BELKA/test.parquet\", columns=[\"id\", \"protein_name\"])\nsub_merged = sub.join(test_df, on=\"id\")\nprint(sub_merged.group_by(\"protein_name\").agg(pl.col(\"binds\").mean()))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-19T11:39:02.775943Z","iopub.status.idle":"2026-06-19T11:39:02.776216Z","shell.execute_reply.started":"2026-06-19T11:39:02.776066Z","shell.execute_reply":"2026-06-19T11:39:02.77608Z"}},"outputs":[],"execution_count":null}]}