DicongLi 7eff2b3d02 feat: 支持 MolFormer/GraphMVP/GROVER 离线 embedding 作为化学 token
三个预训练分子编码器作为额外的 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 与预训练权重未入库(可由编码脚本复现)
2026-07-23 20:18:53 +08:00

58 lines
2.5 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""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)