58 lines
2.5 KiB
Python
Raw 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)