FlagRelease/materials.smi-ted-nvidia-FlagOS
0
1#!/usr/bin/env python32"""3SMI-TED BBBP 分类验证 — 启用 FlagGems 算子。4使用 encode() 提取嵌入 + XGBoost 分类,对标官方 notebook ROC-AUC = 0.9194。5 6用法:7 python scripts/inference/validate_bbbp_flaggems.py8 9 # 排除特定算子10 python scripts/inference/validate_bbbp_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 xgboost import XGBClassifier28from sklearn.metrics import roc_auc_score29 30import flag_gems31 32from models.smi_ted.smi_ted_light.load import load_smi_ted, normalize_smiles33 34 35def main():36 parser = argparse.ArgumentParser(description="SMI-TED BBBP 分类验证 (FlagGems)")37 parser.add_argument("--flaggems-exclude", type=str, nargs="*", default=["mm", "gelu"],38 help="排除导致海光崩溃的算子 (默认: mm gelu)")39 args = parser.parse_args()40 41 model_dir = os.path.join(_PROJECT_ROOT, "models", "smi_ted", "smi_ted_light")42 train_path = os.path.join(_PROJECT_ROOT, "models", "smi_ted", "finetune", "moleculenet", "bbbp", "train.csv")43 test_path = os.path.join(_PROJECT_ROOT, "models", "smi_ted", "finetune", "moleculenet", "bbbp", "test.csv")44 45 device = "cuda" if torch.cuda.is_available() else "cpu"46 47 # ── 启用 FlagGems ──48 exclude = set(args.flaggems_exclude)49 log_path = os.path.join(_PROJECT_ROOT, "flag_gems_enable.txt")50 flag_gems.enable(unused=exclude if exclude else None, record=True, once=True, path=log_path)51 print(f"FlagGems 已启用 (version: {flag_gems.__version__})")52 if exclude:53 print(f" 排除: {exclude}")54 print()55 56 # ── 加载数据 ──57 train_df = pd.read_csv(train_path)58 test_df = pd.read_csv(test_path)59 train_df["norm_smiles"] = train_df["smiles"].apply(normalize_smiles)60 test_df["norm_smiles"] = test_df["smiles"].apply(normalize_smiles)61 train_df = train_df.dropna()62 test_df = test_df.dropna()63 64 print(f"数据集: BBBP (血脑屏障渗透)")65 print(f" 训练集: {len(train_df)} 分子")66 print(f" 测试集: {len(test_df)} 分子")67 68 # ── 加载模型 ──69 print("加载模型...")70 model = load_smi_ted(folder=model_dir, ckpt_filename="smi-ted-Light_40.pt")71 model.eval()72 if device == "cuda":73 model = model.cuda()74 75 # ── 提取嵌入 ──76 print("提取训练集嵌入...")77 t0 = time.time()78 with torch.no_grad():79 train_emb = model.encode(train_df["norm_smiles"].tolist())80 print(f" 训练集: {train_emb.shape} 耗时: {time.time() - t0:.1f}s")81 82 print("提取测试集嵌入...")83 t0 = time.time()84 with torch.no_grad():85 test_emb = model.encode(test_df["norm_smiles"].tolist())86 print(f" 测试集: {test_emb.shape} 耗时: {time.time() - t0:.1f}s")87 88 # ── XGBoost 分类 ──89 print("训练 XGBoost 分类器...")90 xgb = XGBClassifier(n_estimators=2000, learning_rate=0.04, max_depth=8)91 xgb.fit(train_emb, train_df["p_np"])92 y_prob = xgb.predict_proba(test_emb)[:, 1]93 roc_auc = roc_auc_score(test_df["p_np"], y_prob)94 95 # ── 输出结果 ──96 print(f"\n{'=' * 50}")97 print(f"BBBP 分类验证结果 (FlagGems: 已启用)")98 print(f"{'=' * 50}")99 print(f" ROC-AUC: {roc_auc:.4f}")100 print(f" 官方参考值: 0.9194")101 print(f" 差异: {abs(roc_auc - 0.9194):.4f}")102 103 if abs(roc_auc - 0.9194) < 0.02:104 print(f"\n ✓ 验证通过 (差异 < 0.02)")105 else:106 print(f"\n ✗ 验证未通过!")107 print(f"{'=' * 50}")108 109 110if __name__ == "__main__":111 main()112 