mirror of
https://github.com/RYDE-WORK/lnp_ml.git
synced 2026-09-18 14:23:20 +08:00
三个预训练分子编码器作为额外的 chemical token 接入(三选一,互斥):
- MolFormer-XL 768d(SMILES 序列,11亿分子预训练)
- GraphMVP 300d(2D 图自监督,ICLR'22)
- GROVER 3200d(图 Transformer,NeurIPS'20)
实现:embedding 离线编码存 npy,模型内按 {smiles: vector} 查表注入,
不改动前向逻辑,默认全关、不影响既有实验。
- models.py: 新增三个 CHEM_KEYS_WITH_*、统一的 _load_offline_emb 查表、
proj_input_dims 按开关裁剪、forward 注入;子类同步转发参数
- nested_cv_optuna.py / pretrain.py: 新增 --{molformer,graphmvp,grover}-emb/-csv
- scripts_run/encode_embeddings/: 三个离线编码脚本 + README(含权重来源与
新旧环境兼容补丁说明)
- results/pretrained_encoders/: 五组 15 折结果 summary
(trials15/inner3/repeats3/seed42, 均带 external 预训练)
注:npy embedding 与预训练权重未入库(可由编码脚本复现)
123 lines
4.6 KiB
Python
123 lines
4.6 KiB
Python
"""GraphMVP 离线 embedding (300d)
|
||
|
||
依赖:torch-geometric, ogb
|
||
权重:https://github.com/chao1224/GraphMVP → Drive 文件夹内
|
||
GraphMVP_complate_features_for_regression/GraphMVP/pretraining_model.pth
|
||
|
||
注意:新版 OGB 的 chirality 类别数为 5,而 checkpoint 为 4,
|
||
加载时按形状逐层对齐(前 4 行用预训练,新增类别保持随机初始化)。
|
||
|
||
用法:
|
||
python encode_graphmvp.py --ckpt pretraining_model.pth \
|
||
--csv data/interim/internal.csv --out data/interim/graphmvp_emb.npy
|
||
"""
|
||
import argparse
|
||
import numpy as np
|
||
import pandas as pd
|
||
import torch
|
||
import torch.nn as nn
|
||
import torch.nn.functional as F
|
||
from ogb.graphproppred.mol_encoder import AtomEncoder, BondEncoder
|
||
from ogb.utils.mol import smiles2graph
|
||
from torch_geometric.data import Batch, Data
|
||
from torch_geometric.nn import MessagePassing, global_mean_pool
|
||
|
||
|
||
class GINConv(MessagePassing):
|
||
def __init__(self, emb_dim: int) -> None:
|
||
super().__init__(aggr="add")
|
||
self.mlp = nn.Sequential(
|
||
nn.Linear(emb_dim, 2 * emb_dim), nn.BatchNorm1d(2 * emb_dim),
|
||
nn.ReLU(), nn.Linear(2 * emb_dim, emb_dim),
|
||
)
|
||
self.eps = nn.Parameter(torch.Tensor([0]))
|
||
self.bond_encoder = BondEncoder(emb_dim=emb_dim)
|
||
|
||
def forward(self, x, edge_index, edge_attr):
|
||
e = self.bond_encoder(edge_attr)
|
||
return self.mlp((1 + self.eps) * x + self.propagate(edge_index, x=x, edge_attr=e))
|
||
|
||
def message(self, x_j, edge_attr):
|
||
return F.relu(x_j + edge_attr)
|
||
|
||
def update(self, aggr_out):
|
||
return aggr_out
|
||
|
||
|
||
class GNNComplete(nn.Module):
|
||
def __init__(self, num_layer: int = 5, emb_dim: int = 300, drop_ratio: float = 0.0) -> None:
|
||
super().__init__()
|
||
self.num_layer, self.drop_ratio = num_layer, drop_ratio
|
||
self.atom_encoder = AtomEncoder(emb_dim)
|
||
self.gnns = nn.ModuleList([GINConv(emb_dim) for _ in range(num_layer)])
|
||
self.batch_norms = nn.ModuleList([nn.BatchNorm1d(emb_dim) for _ in range(num_layer)])
|
||
|
||
def forward(self, x, edge_index, edge_attr):
|
||
h = self.atom_encoder(x)
|
||
for layer in range(self.num_layer):
|
||
h = self.gnns[layer](h, edge_index, edge_attr)
|
||
h = self.batch_norms[layer](h)
|
||
h = F.dropout(h if layer == self.num_layer - 1 else F.relu(h),
|
||
self.drop_ratio, training=self.training)
|
||
return h
|
||
|
||
|
||
def load_weights(model: nn.Module, ckpt_path: str) -> None:
|
||
sd = torch.load(ckpt_path, map_location="cpu")
|
||
msd = model.state_dict()
|
||
loaded, skipped, partial = 0, [], []
|
||
for k, v in sd.items():
|
||
if k not in msd:
|
||
continue
|
||
if msd[k].shape == v.shape:
|
||
msd[k] = v
|
||
loaded += 1
|
||
elif v.dim() == 2 and msd[k].shape[1] == v.shape[1] and msd[k].shape[0] > v.shape[0]:
|
||
msd[k][: v.shape[0]] = v # OGB 类别数变化,前 N 行用预训练
|
||
partial.append(k)
|
||
loaded += 1
|
||
else:
|
||
skipped.append((k, tuple(v.shape), tuple(msd[k].shape)))
|
||
model.load_state_dict(msd)
|
||
print(f"已加载 {loaded}/{len(sd)} 个张量 | 部分加载: {partial} | 跳过: {skipped}")
|
||
assert not skipped, "有张量形状不兼容"
|
||
|
||
|
||
def main(ckpt: str, csv_path: str, out_npy: str, col: str = "smiles", batch: int = 64) -> None:
|
||
model = GNNComplete().cuda().eval()
|
||
load_weights(model, ckpt)
|
||
|
||
smi = pd.read_csv(csv_path, low_memory=False)[col].astype(str).tolist()
|
||
uniq = list(dict.fromkeys(smi))
|
||
print(f"{len(smi)} 行, {len(uniq)} 唯一分子")
|
||
|
||
lut = {}
|
||
with torch.no_grad():
|
||
for i in range(0, len(uniq), batch):
|
||
datas = []
|
||
for s in uniq[i:i + batch]:
|
||
g = smiles2graph(s)
|
||
datas.append(Data(
|
||
x=torch.tensor(g["node_feat"], dtype=torch.long),
|
||
edge_index=torch.tensor(g["edge_index"], dtype=torch.long),
|
||
edge_attr=torch.tensor(g["edge_feat"], dtype=torch.long),
|
||
))
|
||
b = Batch.from_data_list(datas).cuda()
|
||
hg = global_mean_pool(model(b.x, b.edge_index, b.edge_attr), b.batch).cpu().numpy()
|
||
for j, s in enumerate(uniq[i:i + batch]):
|
||
lut[s] = hg[j]
|
||
|
||
E = np.stack([lut[s] for s in smi])
|
||
np.save(out_npy, E)
|
||
print(f"{out_npy} {E.shape}")
|
||
|
||
|
||
if __name__ == "__main__":
|
||
p = argparse.ArgumentParser()
|
||
p.add_argument("--ckpt", required=True)
|
||
p.add_argument("--csv", required=True)
|
||
p.add_argument("--out", required=True)
|
||
p.add_argument("--col", default="smiles")
|
||
a = p.parse_args()
|
||
main(a.ckpt, a.csv, a.out, a.col)
|