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