mirror of
https://github.com/RYDE-WORK/lnp_ml.git
synced 2026-09-18 19:13:21 +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 与预训练权重未入库(可由编码脚本复现)
58 lines
2.5 KiB
Python
58 lines
2.5 KiB
Python
"""GROVER fingerprint 离线编码 (3200d = atom 1600 + bond 1600)
|
||
|
||
代码:https://github.com/tencent-ailab/grover ,权重 grover_base
|
||
原项目基于 Python 3.6.8 / PyTorch 1.1,在现代环境需对 grover 源码打两处补丁:
|
||
1. grover/util/utils.py 的 torch.load 加 weights_only=False
|
||
2. build_model 前用 checkpoint 自带 args 补齐 current_args 缺失字段,
|
||
并为 dropout / features_only / ffn_num_layers 等推理期字段提供默认值
|
||
|
||
三步流程:
|
||
# 1) 生成唯一 SMILES 输入
|
||
python encode_grover.py prepare --csv data/interim/internal.csv --out /tmp/grover_in.csv
|
||
# 2) 跑 GROVER fingerprint(在 grover 仓库目录下)
|
||
python main.py fingerprint --data_path /tmp/grover_in.csv \
|
||
--checkpoint_path grover_base.pt --fingerprint_source both --output /tmp/grover_fp.npz
|
||
# 3) 按 csv 行顺序对齐存 npy
|
||
python encode_grover.py align --fp /tmp/grover_fp.npz --input /tmp/grover_in.csv \
|
||
--target data/interim/internal.csv --out data/interim/grover_emb.npy
|
||
"""
|
||
import argparse
|
||
import numpy as np
|
||
import pandas as pd
|
||
|
||
|
||
def prepare(csv_path: str, out_path: str, col: str = "smiles") -> None:
|
||
smi = pd.read_csv(csv_path, low_memory=False)[col].astype(str).tolist()
|
||
uniq = list(dict.fromkeys(smi))
|
||
pd.DataFrame({"smiles": uniq, "label": [0] * len(uniq)}).to_csv(out_path, index=False)
|
||
print(f"{len(smi)} 行 -> {len(uniq)} 唯一分子 -> {out_path}")
|
||
|
||
|
||
def align(fp_npz: str, input_csv: str, target_csv: str, out_npy: str, col: str = "smiles") -> None:
|
||
fps = np.load(fp_npz)["fps"]
|
||
uniq = pd.read_csv(input_csv)["smiles"].astype(str).tolist()
|
||
assert len(uniq) == fps.shape[0], f"{len(uniq)} vs {fps.shape[0]}"
|
||
lut = {s: fps[i] for i, s in enumerate(uniq)}
|
||
rows = pd.read_csv(target_csv, low_memory=False)[col].astype(str).tolist()
|
||
E = np.stack([lut[s] for s in rows])
|
||
np.save(out_npy, E)
|
||
print(f"{out_npy} {E.shape}")
|
||
|
||
|
||
if __name__ == "__main__":
|
||
p = argparse.ArgumentParser()
|
||
sub = p.add_subparsers(dest="cmd", required=True)
|
||
a = sub.add_parser("prepare")
|
||
a.add_argument("--csv", required=True)
|
||
a.add_argument("--out", required=True)
|
||
b = sub.add_parser("align")
|
||
b.add_argument("--fp", required=True)
|
||
b.add_argument("--input", required=True)
|
||
b.add_argument("--target", required=True)
|
||
b.add_argument("--out", required=True)
|
||
args = p.parse_args()
|
||
if args.cmd == "prepare":
|
||
prepare(args.csv, args.out)
|
||
else:
|
||
align(args.fp, args.input, args.target, args.out)
|