一覧へ

グラフニューラルネットワークによる薬物-標的相互作用予測 — 実践的なSARS-CoV-2 Mpro標的薬再利用スクリーニング

RDKitを使用してSMILESを分子グラフに変換し、PyTorch Geometricグラフニューラルネットワーク(GAT)を使用して薬物-標的相互作用(DTI)を学習/予測します。SARS-CoV-2 Mpro標的薬再利用シナリオを含む完全なパイプラインで、BindingDB · Davis · KIBAベンチマークを再現し、新規標的/薬物の一般化を検証し、GATアテンションを使用して結合原子を視覚化します。

中級
|
40
|
検証済み (2026-07)
進捗0/15 (0%)

グラフニューラルネットワークによる薬物-標的相互作用予測:SARS-CoV-2 Mpro薬物の再利用スクリーニングの実践的なガイド

薬物候補分子とタンパク質標的との結合親和性と結合強度を予測することは、創薬における重要なタスクです。ドッキング(パート08)は3次元座標を扱いますが、このパートでは化学構造グラフの取り扱いに焦点を当てます。SMILES文字列は、ノード(原子)とエッジ(結合)で構成されるグラフとして表現され、グラフニューラルネットワーク(GNN)はこのトポロジー情報から学習して結合を予測します。小型のGPU1つで、数万の薬物候補を数秒でスクリーニングし、どの原子が結合に寄与するかを解釈するためにアテンションを可視化することができます。このパートでは、SARS-CoV-2の主要プロテアーゼ(Mpro)を標的とした薬物の再利用のための完全なパイプラインを構築します。

📚 推奨される事前知識(強く推奨)

これは、AI×生物学の高度で詳細なモジュールです。始める前に、以下のDryBenchモジュールを最初に確認することをお勧めします。

事前知識なしに開始しようとすると、グラフ畳み込み、メッセージパッシング、アテンションヘッド、PyTorchの自動微分などの原理を再説明することなく進むため、このパートの実践的なコードを理解するのが難しくなるでしょう。


DryBenchで学んだこと

DryBench ai-native #2では、ニューラルネットワークが、入力のローカルパターンをフィルタリングする層を積み重ねることで、階層的な表現を学習することを学びました。#12では、PyTorchがテンソル、自動微分、GPU演算の基盤であることを学びました。#13では、HuggingFaceが、事前学習済みのモデルをロードするための標準的なAPIを提供することを学びました。

グラフニューラルネットワーク(GNN)は、この原理を不規則なグラフ構造に拡張します。画像はピクセルのグリッドであり、テキストは線形シーケンスであるため、CNNとRNNは自然な選択肢です。しかし、分子はグラフであり、各ノードは異なる数の隣接ノードを持っています。GNNにおけるメッセージパッシングには、各ノードがその隣接ノードと情報を交換し、自身の表現を更新することが含まれます。これにより、GNNは分子、タンパク質、ソーシャルネットワークなどのグラフデータに対する標準的な選択肢となります。このパートでは、この原理を、薬物-標的相互作用(DTI)予測に集中的に適用します。

問題の定義

実用的なシナリオ:SARS-CoV-2 Mpro薬物の再利用

2020-2023年のパンデミック中、いくつかの研究チームが、承認された薬(FDA承認薬、DrugBank、Repurposing Hub)のライブラリをスクリーニングし、SARS-CoV-2の主要プロテアーゼ(Mpro、3CLpro)に結合する可能性のある薬物を探索しました。in vitro実験は、1つの薬あたり数万ドルかかり、数日かかるため、in silicoによる事前フィルタリングが不可欠です。当社のパイプラインは、以下を目指します。

  • 入力: 3000個の承認済み薬のSMILES + Mpro配列(UniProt P0DTD1、306残基)+ PDB構造(例:6LU7、7BQY)。
  • 出力: 各薬物とMproとの結合確率と結合強度スコア、および上位20個の候補の推奨。
  • 検証: PDBbindとBindingDBに測定されたMproリガンド、および既知のプロテアーゼ阻害剤(例:ニルマトレルビル)との一貫性。
  • コスト: ローカルの小型GPUで、30分以内にパイプライン全体を完了させる。

DTIにおける根本的な課題

  • 分類タスク: 結合と非結合の二値予測(BindingDB、DrugBank、BioSNAP)。
  • 回帰タスク: pIC50、Kd、またはKiなどの結合親和性の定量的な予測(Davis、KIBA、Metzベンチマーク)。
  • コールドターゲット汎化: 訓練中に見たことのない新しい標的タンパク質に対する予測性能。
  • コールド薬物汎化: 訓練中に見たことのない新しい薬物候補に対する予測性能。
  • アッセイの多様性: 訓練データは、さまざまなアッセイ(生化学的、細胞、in vivo)の混合物です。アッセイ固有のバイアスが存在します。

既存のアプローチと、このパートの位置づけ

  • ランダムフォレスト + ECFP特徴: 高速で、優れたベースラインです。R² 0.4-0.5。
  • 1D CNN + SMILES文字(DeepDTA 2018): シーケンスとして処理し、構造情報を失います。R² 0.5-0.6 [1]。
  • GraphDTA(Nguyen 2020): 最初のグラフベースのDTI SOTA。R² 0.6-0.65 [2]。
  • MGraphDTA(Yang 2022): マルチスケールグラフとアテンション。R² 0.65-0.70 [3]。
  • GeNNius(2024): 超軽量で超高速。非常に高速な訓練と推論 [4]。
  • DrugBAN(2023): 二重線形アテンションと拡張された解釈可能性。
  • ファウンデーションモデル(ChemBERTa、MolFormer、DeepChem): 大規模な事前学習の後にファインチューニング。
  • AlphaFold3、Boltz-2(パート11): ドッキングを統合します。最も高い精度ですが、GPU要件が高い。

このパートでは、GATベースのベースライン(メインパート)+ ChemBERTa埋め込みの統合(拡張アイデア)+ アテンションの可視化を扱います。

このパートの目標指標

  • BindingDBの二値分類AUCが0.90以上(ランダム分割)。
  • Davisのピアソンのrが0.80以上。
  • コールドターゲット設定2のAUCが0.75以上(汎化の検証)。
  • Mproのシナリオにおいて、既知の阻害剤のうち少なくとも10個を上位20位以内に取得する(つまり、上位20位以内に既知の20個の阻害剤のうち少なくとも10個を取得する)。
  • 訓練にかかる時間が30分未満(小型GPU)。
  • 推論スループットが1秒あたり1000個以上のSMILES。

必要なツールとインフラストラクチャ

ツール役割ライセンス
RDKitSMILES → 分子グラフ、フィンガープリントBSD-3-Clause
PyTorch Geometric (PyG)グラフニューラルネットワークフレームワークMIT
DGL (オプション)PyG の代替Apache 2.0
GeNNius (GitHub)超軽量 GNN ベ이스ラインMIT (予定)
ChemBERTa (HuggingFace)事前学習済みの分子埋め込み (オプション)Apache 2.0
MolFormer (IBM)代替の分子基盤Apache 2.0
BindingDB、Davis、KIBA、BioSNAPトレーニング/検証ベンチマーク学術利用無料
DrugBank、DrugRepurposing Hub薬剤ライブラリ学術利用無料、登録が必要

インフラストラクチャ要件:

  • 小さなコンシューマーGPU (RTX 4060 8GB 以上を推奨。CPU でもトレーニングと推論が可能ですが、5〜10倍遅くなります)。
  • 16GB 以上の RAM。
  • ディスク: BindingDB のデータセットは約 500MB、Davis と KIBA はそれぞれ 100MB 未満、DrugBank は約 200MB です。

学習者向けの概算コスト: API コストは 0 です (完全にローカル)。GPU 時間は、トレーニングに 30 分未満、推論は数秒です。

実用的なパイプラインの実装

全体の流れ:

mermaid

ステップ1: SMILES → 分子グラフ

python
from dataclasses import dataclass
from typing import Any
import torch
from torch_geometric.data import Data
from rdkit import Chem, RDLogger
RDLogger.DisableLog("rdApp.*")
ATOM_FEATURES = {
"atomic_num": list(range(1, 119)),
"degree": [0, 1, 2, 3, 4, 5, 6],
"formal_charge": [-3, -2, -1, 0, 1, 2, 3],
"hybridization": [
Chem.rdchem.HybridizationType.SP,
Chem.rdchem.HybridizationType.SP2,
Chem.rdchem.HybridizationType.SP3,
Chem.rdchem.HybridizationType.SP3D,
Chem.rdchem.HybridizationType.SP3D2,
],
"num_h": [0, 1, 2, 3, 4],
"is_aromatic": [False, True],
"is_in_ring": [False, True],
}
def one_hot(value: Any, choices: list) -> list[int]:
"""不一致の場合、最後のスロットに1を配置(OOV処理)。"""
if value in choices:
idx = choices.index(value)
else:
idx = len(choices) - 1
result = [0] * len(choices)
result[idx] = 1
return result
def atom_features(atom: Chem.Atom) -> list[float]:
features = []
features += one_hot(atom.GetAtomicNum(), ATOM_FEATURES["atomic_num"])
features += one_hot(atom.GetDegree(), ATOM_FEATURES["degree"])
features += one_hot(atom.GetFormalCharge(), ATOM_FEATURES["formal_charge"])
features += one_hot(atom.GetHybridization(), ATOM_FEATURES["hybridization"])
features += one_hot(atom.GetTotalNumHs(), ATOM_FEATURES["num_h"])
features += one_hot(atom.GetIsAromatic(), ATOM_FEATURES["is_aromatic"])
features += one_hot(atom.IsInRing(), ATOM_FEATURES["is_in_ring"])
return features
def bond_features(bond: Chem.Bond) -> list[float]:
bt = bond.GetBondType()
return [
int(bt == Chem.rdchem.BondType.SINGLE),
int(bt == Chem.rdchem.BondType.DOUBLE),
int(bt == Chem.rdchem.BondType.TRIPLE),
int(bt == Chem.rdchem.BondType.AROMATIC),
int(bond.GetIsConjugated()),
int(bond.IsInRing()),
]
def smiles_to_graph(smiles: str, drug_id: str = "") -> Data | None:
"""SMILES → PyG Data。解析に失敗した場合はNoneを返します。"""
mol = Chem.MolFromSmiles(smiles)
if mol is None or mol.GetNumAtoms() == 0:
return None
canonical = Chem.MolToSmiles(mol, canonical=True)
node_features = torch.tensor(
[atom_features(atom) for atom in mol.GetAtoms()],
dtype=torch.float,
)
edge_indices = []
edge_attrs = []
for bond in mol.GetBonds():
i, j = bond.GetBeginAtomIdx(), bond.GetEndAtomIdx()
feat = bond_features(bond)
edge_indices += [[i, j], [j, i]] # 双方向
edge_attrs += [feat, feat]
if not edge_indices:
edge_index = torch.empty((2, 0), dtype=torch.long)
edge_attr = torch.empty((0, 6), dtype=torch.float)
else:
edge_index = torch.tensor(edge_indices, dtype=torch.long).t().contiguous()
edge_attr = torch.tensor(edge_attrs, dtype=torch.float)
data = Data(x=node_features, edge_index=edge_index, edge_attr=edge_attr)
data.smiles = canonical
data.drug_id = drug_id
return data

ステップ2: ターゲットタンパク質のエンコード(ESM2埋め込みまたはCNNベースライン)

python
def protein_to_onehot(sequence: str, max_length: int = 1200) -> torch.Tensor:
"""タンパク質配列 → (max_length, 21) one-hot。"""
aa_to_idx = {aa: i for i, aa in enumerate("ACDEFGHIKLMNPQRSTVWY")}
seq = sequence[:max_length].upper()
encoded = torch.zeros(max_length, 21)
for i, aa in enumerate(seq):
encoded[i, aa_to_idx.get(aa, 20)] = 1.0
return encoded
class ProteinCNN(torch.nn.Module):
"""タンパク質配列 → 埋め込みベクトル(ベースライン)。"""
def __init__(self, output_dim: int = 128):
super().__init__()
self.conv1 = torch.nn.Conv1d(21, 32, kernel_size=8, padding=3)
self.conv2 = torch.nn.Conv1d(32, 64, kernel_size=8, padding=3)
self.conv3 = torch.nn.Conv1d(64, output_dim, kernel_size=8, padding=3)
self.pool = torch.nn.AdaptiveAvgPool1d(1)
def forward(self, x: torch.Tensor) -> torch.Tensor:
# x: (B, max_length, 21)
x = x.transpose(1, 2) # (B, 21, L)
x = torch.relu(self.conv1(x))
x = torch.relu(self.conv2(x))
x = torch.relu(self.conv3(x))
return self.pool(x).squeeze(-1)
class ProteinESM2(torch.nn.Module):
"""ESM2埋め込みラッパー(02と同様)。"""
MODEL_ID = "facebook/esm2_t6_8M_UR50D" # 8M 小さい、非常に軽量
def __init__(self, device: str = "cuda"):
super().__init__()
from transformers import AutoTokenizer, AutoModel
self.tokenizer = AutoTokenizer.from_pretrained(self.MODEL_ID)
self.model = AutoModel.from_pretrained(self.MODEL_ID).to(device).eval()
self.device = device
self.output_dim = self.model.config.hidden_size # 320 for 8M
@torch.no_grad()
def forward(self, sequences: list[str]) -> torch.Tensor:
embs = []
for seq in sequences:
inputs = self.tokenizer(seq, return_tensors="pt", truncation=True, max_length=1024).to(self.device)
out = self.model(**inputs)
embs.append(out.last_hidden_state.mean(dim=1).squeeze(0))
return torch.stack(embs).to(self.device)

ステップ3: GATベースのDTIモデル

python
from torch_geometric.nn import GATConv, global_mean_pool
class GAT_DTI(torch.nn.Module):
"""分子GAT + タンパク質CNN/ESM + concat MLP → DTI予測。"""
def __init__(
self,
atom_feat_dim: int = 133, # atom_features の出力次元
gat_hidden: int = 128,
gat_heads: int = 4,
protein_dim: int = 128,
fusion_dim: int = 256,
task: str = "classification", # "regression" または "classification"
):
super().__init__()
self.gat1 = GATConv(atom_feat_dim, gat_hidden, heads=gat_heads, dropout=0.1)
self.gat2 = GATConv(gat_hidden * gat_heads, gat_hidden, heads=1, dropout=0.1)
self.protein_encoder = ProteinCNN(output_dim=protein_dim)
self.fusion = torch.nn.Sequential(
torch.nn.Linear(gat_hidden + protein_dim, fusion_dim),
torch.nn.ReLU(),
torch.nn.Dropout(0.2),
torch.nn.Linear(fusion_dim, fusion_dim // 2),
torch.nn.ReLU(),
)
self.head = torch.nn.Linear(fusion_dim // 2, 1)
self.task = task
def forward(self, mol_data: Data, protein_onehot: torch.Tensor, return_attention: bool = False) -> torch.Tensor:
x, edge_index = mol_data.x, mol_data.edge_index
if return_attention:
x, (_, alpha1) = self.gat1(x, edge_index, return_attention_weights=True)
else:
x = self.gat1(x, edge_index)
x = torch.relu(x)
x = self.gat2(x, edge_index)
mol_emb = global_mean_pool(x, mol_data.batch) # (B, gat_hidden)
prot_emb = self.protein_encoder(protein_onehot) # (B, protein_dim)
combined = torch.cat([mol_emb, prot_emb], dim=-1)
hidden = self.fusion(combined)
out = self.head(hidden).squeeze(-1)
if self.task == "classification":
out = torch.sigmoid(out)
if return_attention:
return out, alpha1
return out

ステップ4: DataLoaderと学習ループ

python
from torch_geometric.loader import DataLoader as PyGLoader
class DTIDataset(torch.utils.data.Dataset):
"""DTI学習データセット。"""
def __init__(self, records: list[dict]):
self.records = []
for r in records:
g = smiles_to_graph(r["smiles"], drug_id=r.get("drug_id", ""))
if g is None:
continue
g.protein = protein_to_onehot(r["protein_sequence"])
g.y = torch.tensor(r["label"], dtype=torch.float)
g.target_id = r.get("target_id", "")
self.records.append(g)
def __len__(self):
return len(self.records)
def __getitem__(self, idx):
return self.records[idx]
def train_dti(
train_records: list[dict],
val_records: list[dict],
epochs: int = 50,
batch_size: int = 64,
lr: float = 1e-3,
device: str = "cuda",
task: str = "classification",
) -> torch.nn.Module:
train_ds = DTIDataset(train_records)
val_ds = DTIDataset(val_records)
train_loader = PyGLoader(train_ds, batch_size=batch_size, shuffle=True)
val_loader = PyGLoader(val_ds, batch_size=batch_size)
model = GAT_DTI(task=task).to(device)
optimizer = torch.optim.Adam(model.parameters(), lr=lr, weight_decay=1e-5)
criterion = torch.nn.BCELoss() if task == "classification" else torch.nn.MSELoss()
best_val_metric = 0.0 if task == "classification" else float("inf")
for epoch in range(epochs):
# 学習
model.train()
total_loss = 0.0
for batch in train_loader:
batch = batch.to(device)
optimizer.zero_grad()
pred = model(batch, batch.protein.view(batch.num_graphs, -1, 21))
loss = criterion(pred, batch.y)
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
optimizer.step()
total_loss += loss.item()
# 検証
val_metric = evaluate_model(model, val_loader, device, task)
print(f"Epoch {epoch+1}: train_loss={total_loss/len(train_loader):.4f}, val_metric={val_metric:.4f}")
# ベストモデルの保存
if (task == "classification" and val_metric > best_val_metric) or \
(task == "regression" and val_metric < best_val_metric):
best_val_metric = val_metric
torch.save(model.state_dict(), "best_dti.pt")
return model
def evaluate_model(model, loader, device: str, task: str) -> float:
from sklearn.metrics import roc_auc_score, mean_squared_error
model.eval()
all_preds, all_labels = [], []
with torch.no_grad():
for batch in loader:
batch = batch.to(device)
pred = model(batch, batch.protein.view(batch.num_graphs, -1, 21))
all_preds.extend(pred.cpu().tolist())
all_labels.extend(batch.y.cpu().tolist())
if task == "classification":
return roc_auc_score(all_labels, all_preds)
return mean_squared_error(all_labels, all_preds, squared=False)

ステップ5: コールドターゲット、コールドドラッグ分割(汎化の検証)

DTIベンチマークの真の課題。コールド分割は、ランダム分割よりも実際の創薬シナリオにより近い [5]。

python
import random
def cold_target_split(
records: list[dict], val_target_ratio: float = 0.2, seed: int = 42,
) -> tuple[list, list]:
"""コールドターゲット分割:検証ターゲットはトレーニングセットに存在しない。"""
random.seed(seed)
all_targets = list({r["target_id"] for r in records})
random.shuffle(all_targets)
n_val = int(len(all_targets) * val_target_ratio)
val_targets = set(all_targets[:n_val])
train = [r for r in records if r["target_id"] not in val_targets]
val = [r for r in records if r["target_id"] in val_targets]
return train, val
def cold_drug_split(
records: list[dict], val_drug_ratio: float = 0.2, seed: int = 42,
) -> tuple[list, list]:
"""コールドドラッグ分割:検証ドラッグはトレーニングセットに存在しない。"""
random.seed(seed)
all_drugs = list({r["drug_id"] for r in records})
random.shuffle(all_drugs)
n_val = int(len(all_drugs) * val_drug_ratio)
val_drugs = set(all_drugs[:n_val])
train = [r for r in records if r["drug_id"] not in val_drugs]
val = [r for r in records if r["drug_id"] in val_drugs]
return train, val

ステップ6: SARS-CoV-2 Mproのドラッグリポジショニングスクリーニング

python
import pandas as pd
MPRO_SEQUENCE = ( # SARS-CoV-2 Mpro (UniProt P0DTD1, 306残基)
"SGFRKMAFPSGKVEGCMVQVTCGTTTLNGLWLDDVVYCPRHVICTSEDMLNPNYEDLLIRKSNHNFLVQAGNVQLRVIGHSMQNCVLKLKVDTANPKTPKYKFVR"
"IQPGQTFSVLACYNGSPSGVYQCAMRPNFTIKGSFLNGSCGSVGFNIDYDCVSFCYMHHMELPTGVHAGTDLEGNFYGPFVDRQTAQAAGTDTTITVNVLAWLYA"
"AVINGDRWFLNRFTTTLNDFNLVAMKYNYEPLTQDHVDILGPLSAQTGIAVLDMCASLKELLQNGMNGRTILGSALLEDEFTPFDVVRQCSGVTFQ"
)
KNOWN_MPRO_INHIBITORS = [
# nirmatrelvir (Paxlovid成分)
"CC1(C)C2CC1C(NC(=O)C1CCCN1C(=O)C(NC(=O)OC(C)(C)C)C(C)(C)C)C(=O)NC(C#N)CC2=O",
# ensitrelvir (Xocova) など。
# 実際には、ChEMBL、DrugBankなどから取得する。
]
def screen_mpro_repurposing(
model: torch.nn.Module,
drug_smiles_list: list[tuple[str, str]], # (drug_id, smiles)
protein_sequence: str = MPRO_SEQUENCE,
top_k: int = 20,
device: str = "cuda",
) -> pd.DataFrame:
"""承認済みのドラッグライブラリからMproに対する上位K候補を返す。"""
model.eval()
prot_encoded = protein_to_onehot(protein_sequence).unsqueeze(0).to(device)
results = []
for drug_id, smiles in drug_smiles_list:
g = smiles_to_graph(smiles, drug_id)
if g is None:
continue
g.batch = torch.zeros(g.x.size(0), dtype=torch.long)
g = g.to(device)
with torch.no_grad():
pred = model(g, prot_encoded).item()
results.append({
"drug_id": drug_id,
"smiles": smiles,
"mpro_score": pred,
})
df = pd.DataFrame(results).sort_values("mpro_score", ascending=False)
return df.head(top_k)
def validate_against_known_inhibitors(
model: torch.nn.Module,
all_screened: pd.DataFrame,
known_inhibitors_smiles: list[str],
top_k: int = 20,
) -> dict:
"""既知の阻害剤のうち、上位K個の中にいくつ含まれているか?"""
known_set = set(known_inhibitors_smiles)
top_smiles = set(all_screened.head(top_k)["smiles"].tolist())
recall = len(top_smiles & known_set) / max(len(known_set), 1)
return {
"top_k": top_k,
"recall_of_known": recall,
"known_in_top_k": list(top_smiles & known_set),
}

ステップ7: アテンションの可視化(解釈可能性)

GATのアテンションウェイトを使用して、どの原子が予測に寄与しているかを可視化する。

python
def visualize_attention_on_molecule(
model: GAT_DTI,
smiles: str,
protein_sequence: str,
output_path: str = "attention.png",
device: str = "cuda",
) -> None:
"""分子上にアテンションを重ねて表示する。"""
from rdkit.Chem import Draw
from rdkit.Chem.Draw import rdMolDraw2D
from matplotlib.cm import get_cmap
g = smiles_to_graph(smiles)
g.batch = torch.zeros(g.x.size(0), dtype=torch.long)
g = g.to(device)
prot = protein_to_onehot(protein_sequence).unsqueeze(0).to(device)
model.eval()
with torch.no_grad():
pred, alpha = model(g, prot, return_attention=True)
# alpha: (num_edges, num_heads)。アトムレベルのアテンションを集約
edge_index = g.edge_index.cpu().numpy()
alpha_np = alpha.mean(dim=1).cpu().numpy() # (num_edges,)
n_atoms = g.x.size(0)
atom_attn = np.zeros(n_atoms)
for e_idx, dest in enumerate(edge_index[1]):
atom_attn[dest] += float(alpha_np[e_idx])
atom_attn = atom_attn / max(atom_attn.max(), 1e-8)
# RDKitで色を強調表示
mol = Chem.MolFromSmiles(smiles)
cmap = get_cmap("Reds")
highlight_colors = {i: cmap(float(atom_attn[i]))[:3] for i in range(n_atoms)}
drawer = rdMolDraw2D.MolDraw2DCairo(500, 500)
drawer.drawOptions().addAtomIndices = False
rdMolDraw2D.PrepareAndDrawMolecule(
drawer, mol,
highlightAtoms=list(range(n_atoms)),
highlightAtomColors=highlight_colors,
)
drawer.FinishDrawing()
with open(output_path, "wb") as f:
f.write(drawer.GetDrawingText())
print(f"アテンションの可視化を保存: {output_path} (予測スコア={pred.item():.3f})")

統合パイプライン

python
def full_pipeline(
training_records: list[dict],
mpro_repurposing_library: list[tuple[str, str]],
output_dir: Path,
device: str = "cuda",
) -> None:
"""DTIモデルのトレーニング → Mproスクリーニング → アテンションの可視化。"""
output_dir.mkdir(parents=True, exist_ok=True)
print("[1/4] コールドターゲット分割")
train, val = cold_target_split(training_records, val_target_ratio=0.2)
print(f" train {len(train)} / val {len(val)}")
print("[2/4] GAT_DTIのトレーニング")
model = train_dti(train, val, epochs=50, task="classification", device=device)
print("[3/4] Mproリポジショニングスクリーニング")
top20 = screen_mpro_repurposing(model, mpro_repurposing_library, top_k=20, device=device)
top20.to_csv(output_dir / "mpro_top20.csv", index=False)
print("[4/4] 上位3つのアテンションを可視化")
for _, row in top20.head(3).iterrows():
visualize_attention_on_molecule(
model, row["smiles"], MPRO_SEQUENCE,
output_path=str(output_dir / f"attn_{row['drug_id']}.png"),
device=device,
)
# 既知の阻害剤に対して検証
validation = validate_against_known_inhibitors(
model, top20, KNOWN_MPRO_INHIBITORS,
)
print(f"上位20個に含まれる既知の阻害剤のリコール: {validation['recall_of_known']:.2f}")
## パフォーマンス、コスト、および既知の失敗事例
### パフォーマンスの参照(公開ベンチマークを使用)
| モデル | ベンチマーク(Davis、KIBA、BindingDB) | AUC / ピアソンの相関係数 | 参照 |
|------|------------------------------|:----------------:|------|
| ランダムフォレスト + ECFP | Davis回帰 | r = 0.55 | 既存のデータ |
| DeepDTA (1D CNN) | Davis | r = 0.62 | Öztürk et al., Bioinformatics 2018 [1] |
| GraphDTA (GAT) | Davis | r = 0.68 | Nguyen et al., Bioinformatics 2020 [2] |
| MGraphDTA | Davis、KIBA | r = 0.72 | Yang et al., Chem Sci 2022 [3] |
| GeNNius | BindingDB | AUC 0.92 | ML4BM Lab 2024 [4] |
| DrugBAN | BioSNAP | AUC 0.89 | Bai et al., Nat Mach Intell 2023 |
| コールドターゲット設定 | Davisコールドターゲット | r = 0.350.50(大幅な低下) | Pahikkala 2015 [5] |
| SARS-CoV-2 Mpro ベンチマーク(COVID Moonshot) | 候補ライブラリ | さまざまな論文のベンチマーク | Moonshot consortium [6] |
### 学習者が再現するための推定コスト
- APIコスト:0(完全にローカル)。
- Davisの学習(3万ペア)を小型の消費者向けGPUで行う場合:約3060分。
- 推論:1秒あたり1000以上のSMILESを処理可能。
- 3000個の薬物のMproスクリーニング:数秒で完了。
### 5つの既知の失敗事例(コミュニティ/論文の収集)
1. **コールドターゲットシナリオにおけるパフォーマンスの著しい低下(一般化の失敗)**
症状:ランダム分割ではAUCが0.90であるが、コールドターゲット設定では0.55に低下する。
原因:トレーニングデータ内の一般的なターゲット(キナーゼなど)への過剰適合。新しいターゲットは異なる配列表現を持つ。
軽減策:(a) ESM2/ESM3の事前学習済みモデルでターゲットエンコーダーを強化する、(b) コントラスト学習によってターゲット埋め込み空間を正規化する、(c) ベンチマークにコールド分割を使用させる(このセクションのステップ5)、(d) ChemBERTa/MolFormerの分子埋め込みの初期化を使用する。
参照:Pahikkala et al. "Toward more realistic drug-target interaction predictions." Brief Bioinform 2015 [5].
2. **SMILESの標準化がないことによる重複学習**
症状:同じ分子の複数のSMILES表現が、トレーニング中に異なるサンプルとして扱われ、データの漏洩が発生する。
原因:RDKitによる標準化が行われていない場合、`CCO`と`C(C)O`のような表現が異なって扱われる。
軽減策:トレーニングデータのプリプロセス中に`Chem.MolToSmiles(mol, canonical=True)`を強制的に実行し、重複を削除し、タウトマーの標準化(RDKit `MolStandardize`)を実行する。
参照:RDKit Discussions [7].
3. **PyG DataLoaderのコラートエラー(分子とタンパク質のバッチ間の不一致)**
症状:バッチ内の分子とタンパク質が誤って整列されるか、トレーニング中に形状の不一致が発生する。
原因:PyGのDataオブジェクトは自動的にバッチ化されるが、タンパク質テンソルは別途処理する必要がある。
軽減策:カスタムコラート関数を使用して、(a) PyGで分子をバッチ化し、(b) PyTorchでタンパク質テンソルをスタックし、(c) ラベルをスタックする。あるいは、タンパク質のフィーチャーをPyG Dataオブジェクトにアタッチする。
参照:PyG GitHub Discussions [8].
4. **アッセイの多様性によるラベルのバイアス**
症状:トレーニングデータ内のラベルは、さまざまなアッセイ(生化学的Kd、細胞性IC50、生体内)から取得される。値のスケールが異なる。
原因:異なるアッセイは、異なる動的範囲と検出限界を持つ。
軽減策:(a) アッセイタイプを特徴量として追加する、(b) 同じアッセイからのデータのみでトレーニングを行う、(c) 対数変換とzスコア正規化を行う、(d) マルチタスク学習(各アッセイ用の個別のヘッド)を使用する。
参照:Landrum et al. RDKit chemistry blog [7].
5. **アテンション視覚化の解釈における落とし穴**
症状:GATにおける高アテンション原子は、必ずしも結合に最も寄与する原子であるとは限らない。誤解を招く解釈。
原因:アテンション重みは、学習プロセスの副産物であり、因果関係または生物学的な意味を保証するものではない。
軽減策:(a) 統合勾配やGNNExplainerなどの複数の説明可能性手法を組み合わせる、(b) アテンションを仮説の生成にのみ使用し、実験的な検証と組み合わせる、(c) 異なるシードで実行された複数のトレーニングで一貫して出現する原子のみを信頼する。
参照:Jain & Wallace "Attention is not Explanation." NAACL 2019 [9].
## 拡張のアイデア
- **ファウンデーション埋め込みの組み合わせ:** ChemBERTa、MolFormer、およびSELFormerの埋め込みを、GNNの初期ノード特徴量として使用する。
- **マルチタスク学習:** DTI、溶解度(水性)、hERG毒性、およびBBB透過性を同時にトレーニングする。
- **拡張された説明可能性:** アテンション、GNNExplainer、および統合勾配を組み合わせる。
- **アクティブラーニング:** 不確実性の高い新しいペアに実験を優先的に行い、ラベルを取得し、再トレーニングする。
- **3Dを考慮したGNN:** SchNet、DimeNet、およびEquiformerのようなモデルからの3D座標を使用し、ドッキング(セクション08)と組み合わせる。
- **構造ベースのリランキング:** GNNから予測された上位の候補を、3D結合親和性についてBoltz-2(セクション11)を使用して再評価する。
## 次のセクション
- セクション08:`docking-hybrid-diffusion`: GNNによって予測された上位の候補のポーズをドッキングを使用して検証する。
- セクション11:`structure-affinity-boltz`: Boltz-2を使用して、最終的な親和性ランキングを実行する。
- セクション12:`single-cell-perturbation`: DTIの結果を、単一細胞応答予測に接続する。
- セクション14:`bio-mcp-agent`: DTIの予測をMCPプラットフォームのツールとして公開し、エージェントによる自律的なスクリーニングを可能にする。
## 参考文献
1. Öztürk H, Özgür A, Ozkirimli E. "DeepDTA: deep drug-target binding affinity prediction." Bioinformatics 2018. `https://academic.oup.com/bioinformatics/article/34/17/i821/5093245`
2. Nguyen T, Le H, Quinn TP, et al. "GraphDTA: Predicting drug-target binding affinity with graph neural networks." Bioinformatics 2020. `https://academic.oup.com/bioinformatics/article/37/8/1140/5942970`
3. Yang Z, Zhong W, Zhao L, Chen CY-C. "MGraphDTA: deep multiscale graph neural network for explainable drug-target binding affinity prediction." Chemical Science 2022. `https://pubs.rsc.org/en/content/articlelanding/2022/sc/d1sc05180f`
4. Muñoz-Gil G, et al. "GeNNius: An ultrafast drug-target interaction inference method based on graph neural networks." 2024. `https://github.com/ML4BM-Lab/GeNNius`
5. Pahikkala T, Airola A, Pietilä S, et al. "Toward more realistic drug-target interaction predictions." Briefings in Bioinformatics 2015.
6. COVID Moonshot consortium: `https://postera.ai/moonshot/`
7. RDKit Discussions: `https://github.com/rdkit/rdkit/discussions`
8. PyG GitHub Discussions: `https://github.com/pyg-team/pytorch_geometric/discussions`
9. Jain S, Wallace BC. "Attention is not Explanation." NAACL 2019. `https://arxiv.org/abs/1902.10186`
10. PyTorch Geometric documentation: `https://pytorch-geometric.readthedocs.io/`
11. DGL (Deep Graph Library): `https://www.dgl.ai/`
12. BindingDB: `https://www.bindingdb.org/`
13. Davis benchmark: Davis MI et al. Nat Biotechnol 2011.
14. KIBA benchmark: Tang J et al. J Chem Inf Model 2014.
15. BioSNAP dataset: `https://snap.stanford.edu/biodata/`
16. ChemBERTa: `https://huggingface.co/DeepChem/ChemBERTa-77M-MTR`
17. MolFormer (IBM): `https://github.com/IBM/molformer`
18. DrugBank: `https://go.drugbank.com/`
19. Drug Repurposing Hub (Broad Institute): `https://clue.io/repurposing`
20. UniProt SARS-CoV-2 Mpro (P0DTD1): `https://www.uniprot.org/uniprotkb/P0DTD1/entry`

💬 質問・コメント

0件のコメント

ログインせずに投稿できます。ゲスト投稿は投稿者自身で編集・削除できません。

0/2000

読み込み中...