CoolFace
Modelpublic

FlagRelease/materials.smi-ted-nvidia-FlagOS

sourceHugging Faceapache-2.0updated 4mo agoView on Hugging Face
0likes
validate_accuracy_flaggems.py125 linesDownload Raw Back to inference
1#!/usr/bin/env python32"""3SMI-TED 推理正确性验证 — 启用 FlagGems 算子。4使用 MOSES 测试集 (1000 个分子), encode→decode 后计算 Morgan 指纹 Tanimoto 相似度。5 6用法:7  python scripts/inference/validate_accuracy_flaggems.py8 9  # 排除特定算子10  python scripts/inference/validate_accuracy_flaggems.py --flaggems-exclude mm gelu11"""12 13import sys14import os15import time16import argparse17import warnings18warnings.filterwarnings("ignore")19 20_SCRIPT_DIR = os.path.dirname(os.path.abspath(__file__))21_PROJECT_ROOT = os.path.dirname(os.path.dirname(_SCRIPT_DIR))22sys.path.insert(0, _PROJECT_ROOT)23 24import torch25import pandas as pd26import numpy as np27from rdkit import Chem28from rdkit.Chem import AllChem29from rdkit.DataStructs import TanimotoSimilarity30 31import flag_gems32 33from models.smi_ted.smi_ted_light.load import load_smi_ted, normalize_smiles34 35 36def main():37    parser = argparse.ArgumentParser(description="SMI-TED 推理正确性验证 (FlagGems)")38    parser.add_argument("--flaggems-exclude", type=str, nargs="*", default=["mm", "gelu"],39                        help="排除导致海光崩溃的算子 (默认: mm gelu)")40    parser.add_argument("--save-flaggems-ops", type=str,41                        default=os.path.join(_PROJECT_ROOT, "flaggems_enabled_ops.txt"),42                        help="保存启用的算子列表")43    args = parser.parse_args()44 45    model_dir = os.path.join(_PROJECT_ROOT, "models", "smi_ted", "smi_ted_light")46    data_path = os.path.join(_PROJECT_ROOT, "models", "smi_ted", "notebooks", "data", "moses_test.csv")47 48    if not os.path.isfile(data_path):49        print(f"[ERROR] 数据集不存在: {data_path}")50        sys.exit(1)51 52    device = "cuda" if torch.cuda.is_available() else "cpu"53 54    # ── 启用 FlagGems ──55    exclude = set(args.flaggems_exclude)56    log_path = os.path.join(_PROJECT_ROOT, "flag_gems_enable.txt")57    flag_gems.enable(unused=exclude if exclude else None, record=True, once=True, path=log_path)58    print(f"FlagGems 已启用 (version: {flag_gems.__version__})")59    if exclude:60        print(f"  排除: {exclude}")61    print()62 63    # ── 加载数据 ──64    df = pd.read_csv(data_path, nrows=1000)65    df["norm_smiles"] = df["SMILES"].apply(normalize_smiles)66    df = df.dropna()67    smiles_list = df["norm_smiles"].tolist()68    print(f"数据集: {len(smiles_list)} 个分子 (MOSES test)")69 70    # ── 加载模型 ──71    print("加载模型...")72    model = load_smi_ted(folder=model_dir, ckpt_filename="smi-ted-Light_40.pt")73    model.eval()74    if device == "cuda":75        model = model.cuda()76 77    # ── Encode + Decode ──78    print("推理中...")79    t0 = time.time()80 81    with torch.no_grad():82        embeddings = model.encode(smiles_list, batch_size=128, return_torch=True)83        reconstructed = model.decode(embeddings)84 85    elapsed = time.time() - t086    print(f"  耗时: {elapsed:.1f}s ({len(smiles_list)/elapsed:.0f} mol/s)")87 88    # ── 计算 Tanimoto 相似度 ──89    similarities = []90    failed = 091    for orig, rec in zip(smiles_list, reconstructed):92        mol_orig = Chem.MolFromSmiles(orig)93        mol_rec = Chem.MolFromSmiles(rec)94        if mol_orig is None or mol_rec is None:95            failed += 196            continue97        fp_orig = AllChem.GetMorganFingerprintAsBitVect(mol_orig, 2)98        fp_rec = AllChem.GetMorganFingerprintAsBitVect(mol_rec, 2)99        similarities.append(TanimotoSimilarity(fp_orig, fp_rec))100 101    # ── 输出结果 ──102    mean_sim = np.mean(similarities)103    min_sim = np.min(similarities)104    perfect = sum(1 for s in similarities if s >= 0.999)105 106    print(f"\n{'=' * 50}")107    print(f"验证结果 (FlagGems: 已启用)")108    print(f"{'=' * 50}")109    print(f"  总样本数:    {len(smiles_list)}")110    print(f"  有效样本:    {len(similarities)}")111    print(f"  RDKit 失败:  {failed}")112    print(f"  均值 Tanimoto: {mean_sim:.4f}")113    print(f"  最小 Tanimoto: {min_sim:.4f}")114    print(f"  完美重建 (≥0.999): {perfect}/{len(similarities)} ({perfect/len(similarities)*100:.1f}%)")115 116    if mean_sim >= 0.99:117        print(f"\n  ✓ 推理正确性验证通过 (均值 ≥ 0.99)")118    else:119        print(f"\n  ✗ 推理正确性验证未通过!")120    print(f"{'=' * 50}")121 122 123if __name__ == "__main__":124    main()125