lnp_ml/scripts/precompute_unimol.py

103 lines
3.7 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.

"""在 unimol 环境pip install unimol-tools中运行把数据集所有唯一 SMILES 编码成 UniMol 表征缓存。
用法(工作目录 = lnp_ml 根):
conda activate unimol
python scripts/precompute_unimol.py --out data/processed/unimol_embeddings.npz
产物 .npz 含 "smiles" (N,) 与 "embeddings" (N, D),供主项目查表使用。
首次运行会自动下载 UniMol 预训练权重HF 不通时先 export HF_ENDPOINT=https://hf-mirror.com
"""
import argparse
from pathlib import Path
import numpy as np
import pandas as pd
DATA_FILES = [
"data/interim/internal.csv",
"data/external/all_data_LiON.csv",
]
GLOB_FILES = [
"data/external/all_amine_split_for_LiON/cv_*/train.csv",
"data/external/all_amine_split_for_LiON/cv_*/test.csv",
]
def collect_smiles(root: Path) -> list:
seen: set = set()
files = [root / f for f in DATA_FILES]
for pattern in GLOB_FILES:
files += sorted(root.glob(pattern))
for f in files:
if not f.exists():
print(f"跳过不存在的文件: {f}")
continue
df = pd.read_csv(f, low_memory=False)
if "smiles" in df.columns:
seen.update(df["smiles"].dropna().astype(str).tolist())
return sorted(seen)
def main() -> None:
ap = argparse.ArgumentParser()
ap.add_argument("--out", default="data/processed/unimol_embeddings.npz")
ap.add_argument("--root", default=".", help="lnp_ml 仓库根目录")
ap.add_argument("--model-name", default="unimolv1", help="unimolv1 / unimolv2")
ap.add_argument("--model-size", default="84m", help="仅 unimolv2 生效")
ap.add_argument("--batch-size", type=int, default=64)
args = ap.parse_args()
from unimol_tools import UniMolRepr
clf = UniMolRepr(
data_type="molecule", remove_hs=False,
model_name=args.model_name, model_size=args.model_size,
)
smiles = collect_smiles(Path(args.root))
print(f"收集到 {len(smiles)} 个唯一 SMILES开始编码...")
def encode_one(s):
try:
rep = clf.get_repr([s], return_atomic_reprs=True)
return np.asarray(rep["cls_repr"], dtype=np.float32).reshape(-1)
except Exception as e:
print(f" 跳过(编码失败): {s[:40]}... ({e})")
return None
vecs: dict = {}
dim = None
bs = args.batch_size
for i in range(0, len(smiles), bs):
batch = smiles[i:i + bs]
try:
rep = clf.get_repr(batch, return_atomic_reprs=True)
cls = np.asarray(rep["cls_repr"], dtype=np.float32)
if cls.ndim != 2 or cls.shape[0] != len(batch):
raise ValueError(f"返回形状 {cls.shape} 与输入 {len(batch)} 不对齐")
for s, v in zip(batch, cls):
vecs[s] = v
dim = v.shape[0]
except Exception as e:
print(f" 批量失败,改逐条: {e}")
for s in batch:
v = encode_one(s)
vecs[s] = v
if v is not None:
dim = v.shape[0]
print(f" {min(i + bs, len(smiles))}/{len(smiles)}")
if dim is None:
raise RuntimeError("没有任何 SMILES 编码成功,请检查 unimol-tools 安装与权重。")
n_fail = sum(1 for s in smiles if vecs.get(s) is None)
embeddings = np.vstack([
vecs[s] if vecs.get(s) is not None else np.zeros(dim, np.float32) for s in smiles
])
out = Path(args.out)
out.parent.mkdir(parents=True, exist_ok=True)
np.savez_compressed(out, smiles=np.array(smiles), embeddings=embeddings)
print(f"已保存 {embeddings.shape}{out}(失败 {n_fail} 个已置零)")
if __name__ == "__main__":
main()