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