FlagRelease/materials.smi-ted-nvidia-FlagOS
0
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 