lnp_ml/scripts_run/encode_embeddings/encode_molformer.py

52 lines
1.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.

"""MolFormer-XL 离线 embedding (768d)
依赖transformers>=4.57(需 transformers.masking_utils。若与项目 pin 的
4.45 冲突,装到独立目录后用 sys.path 隔离:
pip install --target=/tmp/mf_libs "transformers>=4.57" --no-deps
pip install --target=/tmp/mf_libs tokenizers>=0.22,<=0.23 --no-deps
# 本脚本开头 sys.path.insert(0, '/tmp/mf_libs')
用法:
python encode_molformer.py --csv data/interim/internal.csv \
--out data/interim/molformer_emb.npy
"""
import argparse
import numpy as np
import pandas as pd
import torch
MODEL = "ibm/MoLFormer-XL-both-10pct"
def main(csv_path: str, out_npy: str, col: str = "smiles", batch: int = 32) -> None:
from transformers import AutoModel, AutoTokenizer
smiles = pd.read_csv(csv_path, low_memory=False)[col].astype(str).tolist()
print(f"{len(smiles)} 行, {len(set(smiles))} 唯一分子")
tok = AutoTokenizer.from_pretrained(MODEL, trust_remote_code=True)
model = AutoModel.from_pretrained(
MODEL, trust_remote_code=True, deterministic_eval=True
).cuda().eval()
embs = []
with torch.no_grad():
for i in range(0, len(smiles), batch):
b = tok(smiles[i:i + batch], padding=True, truncation=True,
max_length=512, return_tensors="pt")
b = {k: v.cuda() for k, v in b.items()}
embs.append(model(**b).pooler_output.cpu().numpy())
E = np.vstack(embs)
np.save(out_npy, E)
print(f"{out_npy} {E.shape}")
if __name__ == "__main__":
p = argparse.ArgumentParser()
p.add_argument("--csv", required=True)
p.add_argument("--out", required=True)
p.add_argument("--col", default="smiles")
a = p.parse_args()
main(a.csv, a.out, a.col)