lnp_ml/scripts/precompute_chemeleon.py

74 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.

"""在 chemeleon 环境chemprop>=2.2)中运行:把数据集里所有唯一 SMILES 编码成 CheMeleon 指纹缓存。
用法(工作目录 = lnp_ml 根,且同目录有 chemeleon_fingerprint.py:
conda activate chemeleon
python scripts/precompute_chemeleon.py --out data/processed/chemeleon_embeddings.npz
产物 .npz 含 "smiles" (N,) 与 "embeddings" (N, D),供主项目查表使用。
首次运行会自动把权重下载到 ~/.chemprop/chemeleon_mp.pt。
"""
import argparse
import sys
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/chemeleon_embeddings.npz")
ap.add_argument("--root", default=".", help="lnp_ml 仓库根目录")
ap.add_argument("--batch-size", type=int, default=1024)
ap.add_argument("--device", default=None, help="cpu / cuda默认自动")
args = ap.parse_args()
# chemeleon_fingerprint.py 在仓库根,确保可 import
sys.path.insert(0, str(Path(args.root).resolve()))
from chemeleon_fingerprint import CheMeleonFingerprint
fp = CheMeleonFingerprint(device=args.device)
smiles = collect_smiles(Path(args.root))
print(f"收集到 {len(smiles)} 个唯一 SMILES开始编码...")
chunks = []
for i in range(0, len(smiles), args.batch_size):
batch = smiles[i : i + args.batch_size]
chunks.append(np.asarray(fp(batch), dtype=np.float32))
print(f" {min(i + args.batch_size, len(smiles))}/{len(smiles)}")
embeddings = np.vstack(chunks)
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}")
if __name__ == "__main__":
main()