mirror of
https://github.com/RYDE-WORK/lnp_ml.git
synced 2026-09-18 14:23:20 +08:00
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)
|