化學資訊學入門與實作:Multi-Task Learning 讓模型同時預測多個分子毒性
讓 QSAR 模型同時預測 Tox21 十二個毒性端點,大幅提升學習效率。
化學資訊學入門與實作:Multi-Task Learning 讓模型同時預測多個分子毒性

讓 QSAR 模型同時預測 Tox21 十二個毒性端點,大幅提升學習效率。
在上一篇文章「化學資訊學入門與實作:Active Learning 策略在 QSAR 建模中的應用」中,我們介紹了 Active Learning 如何讓模型主動挑選最有價值的訓練樣本,以最少的實驗量達到最大的模型改進。這次,我們把視角從「如何選資料」轉向「如何讓模型同時學更多」 — — 也就是 Multi-Task Learning(多任務學習) 的概念與實作。
為什麼需要同時預測多個性質?
在藥物探索與毒性評估中,一個分子往往需要通過十幾個不同的篩選關卡。以 Tox21 資料集為例,它包含了 12 個毒性端點,涵蓋核受體(Nuclear Receptor, NR)路徑的 7 個標的和壓力反應(Stress Response, SR)路徑的 5 個標的。
傳統做法是為每個端點分別訓練一個 Single-Task 模型,彼此獨立。但這樣做有兩個明顯缺點:
-
每個模型只能用到自己的標記資料,當某個端點的正樣本極少時(毒性篩選中常見),模型很容易過擬合
-
不同端點之間往往存在生物學上的相關性,單任務模型完全忽略了這層資訊
Multi-Task Learning 的核心概念
Multi-Task Learning(MTL) 的基本思想是:讓多個相關任務共享一個底層的特徵表示(shared representation),各自再連接到獨立的輸出頭(task-specific head)。在分子性質預測中,「共享層」負責學習通用的分子化學特徵,「任務頭」則各自對每個端點給出預測。
這種架構有一個直觀的優勢:訓練 NR-AR 端點時學到的「具親核性的芳香族結構傾向有核受體活性」,可以透過共享層,隱性地幫助 NR-ER 的訓練,反之亦然。換句話說,任務之間相互提供了隱式的正規化,有助於緩解資料稀缺問題。
探索任務間的相關性:先做相關矩陣
在進入模型訓練前,先了解 12 個端點之間的相關性,有助於判斷 MTL 的潛在收益。相關性高的任務群(如 NR-AR 與 NR-AR-LBD)共享的生物機制較多,MTL 的正向遷移效果預計較顯著。
import deepchem as dc
import numpy as np
import matplotlib.pyplot as plt
# ── 載入 Tox21(ECFP 特徵,scaffold split)────────────────────────────────
# 資料集:Tox21;分子數:7,831;目標值:12 個二元毒性端點(0/1)
# 原始論文:https://doi.org/10.3389/fenvs.2020.00085
tasks_dc, datasets, transformers = dc.molnet.load_tox21(
featurizer='ECFP', # 1024-bit ECFP4
splitter='scaffold',
reload=True,
)
train, valid, test = datasets
y_train = train.y # shape: (N, 12),含 NaN(未測試的端點)
# 計算有效標記的 Pearson 相關係數(忽略 NaN)
n_tasks = y_train.shape[1]
corr = np.full((n_tasks, n_tasks), np.nan)
for i in range(n_tasks):
for j in range(n_tasks):
mask = ~np.isnan(y_train[:, i]) & ~np.isnan(y_train[:, j])
if mask.sum() > 50:
corr[i, j] = np.corrcoef(y_train[mask, i], y_train[mask, j])[0, 1]
# 視覺化
fig, ax = plt.subplots(figsize=(9, 7))
im = ax.imshow(corr, cmap='RdYlBu_r', vmin=-0.1, vmax=1.0)
ax.set_xticks(range(n_tasks)); ax.set_xticklabels(tasks_dc, rotation=45, ha='right')
ax.set_yticks(range(n_tasks)); ax.set_yticklabels(tasks_dc)
fig.colorbar(im, ax=ax)
ax.set_title('Tox21: Task Correlation Matrix')
plt.tight_layout(); plt.savefig('fig1_tox21_task_correlation.png', dpi=150, bbox_inches='tight')

Tox21 十二個毒性端點的 Pearson 相關矩陣。虛線分隔 NR(核受體)與 SR(壓力反應)兩大群。
用 DeepChem 建立 Multi-Task DNN
DeepChem 的 MultitaskClassifier 是針對多任務分類問題設計的神經網路模型,它的架構由若干全連接共享層加上 12 個獨立的輸出節點組成,對每個端點輸出一個 sigmoid 機率。
需注意Tox21 資料存在大量 missing label(並非每個分子都有 12 個端點的實驗數據)。DeepChem 在計算損失時會自動遮蔽 NaN 標記,只對有效的(分子, 任務)對回傳梯度,確保遺失值不會影響訓練。
from deepchem.models import MultitaskClassifier
from deepchem.metrics import Metric
import deepchem as dc
# ── 模型定義 ─────────────────────────────────────────────────────────────────
model = MultitaskClassifier(
n_tasks=12,
n_features=1024, # ECFP4 特徵維度
layer_sizes=[1024, 512, 256], # 三層共享隱藏層
dropouts=0.25,
learning_rate=1e-4,
batch_size=128,
n_epochs=50,
model_dir='./tox21_multitask_model'
)
# ── 訓練 ─────────────────────────────────────────────────────────────────────
model.fit(train, nb_epoch=50)
# ── 評估(AUROC,各任務平均)───────────────────────────────────────────────
roc_auc = Metric(dc.metrics.roc_auc_score, np.mean)
train_scores = model.evaluate(train, [roc_auc], transformers)
valid_scores = model.evaluate(valid, [roc_auc], transformers)
print(f"Train AUROC (mean): {train_scores['mean-roc_auc_score']:.4f}")
print(f"Valid AUROC (mean): {valid_scores['mean-roc_auc_score']:.4f}")

Multi-Task DNN 與 Single-Task Random Forest 在各端點上的 AUROC 比較。MTL 在大多數任務有 2–6% 的提升。
訓練過程監控:Loss Curve 告訴我們什麼
在訓練 Multi-Task DNN 時,觀察 Loss Curve 尤其重要。由於不同任務的正負樣本比例(class imbalance)差異極大(某些端點陽性率不到 5%),整體損失的下降速率不一定能反映所有任務的學習狀況。
從下圖可以看到,訓練損失與驗證損失的差距在第 30 個 epoch 後開始拉大,這是過擬合的早期信號。使用 Early Stopping(以驗證集 AUROC 為監控指標)可以有效避免這個問題,最佳模型通常落在 30–40 個 epoch 之間。

Multi-Task DNN 的平均 Binary Cross-Entropy Loss 訓練曲線。綠色虛線標示驗證集最佳 epoch。
MTL 的適用條件與限制
MTL 並非對所有情境都有益。具體來說,當任務之間的相關性很低甚至存在負遷移(negative transfer)時,強制共享底層資訊表示反而會拉低單個任務的性能。實驗建議:先計算任務相關矩陣(如上方圖 1),若多數任務對的相關係數 < 0.1,建議改用分組 MTL(將相關任務歸為一組分別訓練)。
此外,MTL 對超參數選擇較敏感:共享層的寬度與深度、各任務損失的加權比例(task weighting),都可能顯著影響結果。一個常見的起點是對所有任務使用相等的損失權重,再依據各任務在驗證集的表現做加權調整。
結語
Multi-Task Learning 在毒性預測與藥物探索中展現出相當的潛力,透過共享分子表示來讓稀缺的毒性標記互相補充,是提升 QSAR 模型泛化能力的有效手段。配合 DeepChem 的標準化 API 與 Tox21 等公開資料集,我們可以在相對短的時間內建立一個覆蓋多個生物端點的預測模型,為後續的分子篩選提供更全面的參考。
之後我們將介紹 Conformal Prediction — — 一種為 QSAR 預測附上統計保證區間的方法,讓模型不只告訴你「這個分子預測活性為 0.82」,還能說「有 90% 的信心,真實值落在 [0.75, 0.89]」。
메타데이터
- post_id
- 9dfd3bbecbee
- slug
- 化學資訊學入門與實作-multi-task-learning-讓模型同時預測多個分子毒性-9dfd3bbecbee
- url
- https://medium.com/@cheninformatics/%E5%8C%96%E5%AD%B8%E8%B3%87%E8%A8%8A%E5%AD%B8%E5%85%A5%E9%96%80%E8%88%87%E5%AF%A6%E4%BD%9C-multi-task-learning-%E8%AE%93%E6%A8%A1%E5%9E%8B%E5%90%8C%E6%99%82%E9%A0%90%E6%B8%AC%E5%A4%9A%E5%80%8B%E5%88%86%E5%AD%90%E6%AF%92%E6%80%A7-9dfd3bbecbee
- canonical_url
- https://medium.com/@cheninformatics/%E5%8C%96%E5%AD%B8%E8%B3%87%E8%A8%8A%E5%AD%B8%E5%85%A5%E9%96%80%E8%88%87%E5%AF%A6%E4%BD%9C-multi-task-learning-%E8%AE%93%E6%A8%A1%E5%9E%8B%E5%90%8C%E6%99%82%E9%A0%90%E6%B8%AC%E5%A4%9A%E5%80%8B%E5%88%86%E5%AD%90%E6%AF%92%E6%80%A7-9dfd3bbecbee
- author_url
- https://medium.com/@cheninformatics
- status
- ok
- fetched_at
- 2026-06-21 07:44:09