mirror of
https://github.com/RYDE-WORK/lnp_ml.git
synced 2026-07-25 15:35:25 +08:00
103 lines
3.7 KiB
Python
103 lines
3.7 KiB
Python
"""在 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() |