CoolFace
Modelpublic

FlagRelease/materials.smi-ted-nvidia-FlagOS

sourceHugging Faceapache-2.0updated 4mo agoView on Hugging Face
0likes
validate_bbbp_flaggems.py112 linesDownload Raw Back to inference
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