lnp_ml/scripts/measure_cost.py

89 lines
3.5 KiB
Python

"""测量 MARLIN / Backbone 的参数量、峰值显存与单-LNP 延迟。"""
import json, time, argparse
from pathlib import Path
import numpy as np
import pandas as pd
import torch
from torch.utils.data import DataLoader, Subset
from lnp_ml.dataset import LNPDataset, collate_fn, process_dataframe
from lnp_ml.modeling.nested_cv_optuna import create_model, _build_rag_pool
def build_marlin(bp, use_llm):
"""use_llm=True -> MARLIN(s1_both); False -> Backbone(s1_baseline)。"""
llm_kwargs = dict(
reg_bypass="on",
use_rag=use_llm, rag_top_k=4, use_llm=use_llm,
llm_model_path="models/qwen2.5-7b-instruct",
llm_freeze=False, llm_use_lora=False, llm_use_qlora=use_llm,
use_soft_prompt=use_llm,
llm_lora_r=bp.get("llm_lora_r", 8), llm_lora_alpha=16, llm_lora_dropout=0.05,
)
return create_model(
d_model=bp["d_model"], num_heads=bp["num_heads"], n_attn_layers=bp["n_attn_layers"],
fusion_strategy=bp["fusion_strategy"], head_hidden_dim=bp["head_hidden_dim"],
dropout=bp["dropout"], use_mpnn=True, mpnn_device="cuda",
set_transformer_block=bp["set_transformer_block"],
use_moe=use_llm, moe_n_experts=bp["moe_n_experts"], moe_top_k=bp["moe_top_k"],
moe_expert_hidden_mult=bp["moe_expert_hidden_mult"],
llm_kwargs=llm_kwargs,
)
@torch.no_grad()
def measure(model, full, train_idx, test_idx, device, use_llm, warmup=3):
model.eval().to(device)
tot = sum(p.numel() for p in model.parameters())
tr = sum(p.numel() for p in model.parameters() if p.requires_grad)
if use_llm: # RAG 前向需要检索池(只用训练集,防泄漏)
s, d, ex = _build_rag_pool(full, train_idx)
model.llm_prompt.set_retrieval_pool(s, d, pool_id="cost", extra_labels=ex)
loader = DataLoader(Subset(full, test_idx.tolist()), batch_size=1,
shuffle=False, collate_fn=collate_fn)
it = iter(loader)
for _ in range(warmup):
b = next(it)
model(b["smiles"], {k: v.to(device) for k, v in b["tabular"].items()})
torch.cuda.synchronize()
torch.cuda.reset_peak_memory_stats()
t0 = time.time(); n = 0
for b in loader:
model(b["smiles"], {k: v.to(device) for k, v in b["tabular"].items()}); n += 1
torch.cuda.synchronize()
dt_ms = (time.time() - t0) / n * 1000
peak_gb = torch.cuda.max_memory_allocated() / 1e9
return tot, tr, peak_gb, dt_ms
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--input", default="data/interim/internal.csv")
ap.add_argument("--run", default="models/abl_full/s1_both/seed42")
ap.add_argument("--fold", type=int, default=0)
args = ap.parse_args()
device = torch.device("cuda")
df = process_dataframe(pd.read_csv(args.input))
full = LNPDataset(df)
run = Path(args.run)
bp = json.load(open(run / "summary.json"))["fold_results"][args.fold]["best_params"]
sp = json.load(open(run / f"outer_fold_{args.fold}" / "splits.json"))
tr_idx = np.array(sp["outer_train_idx"]); te_idx = np.array(sp["outer_test_idx"])
for name, use_llm in [("Backbone", False), ("MARLIN", True)]:
model = build_marlin(bp, use_llm)
tot, trn, mem, lat = measure(model, full, tr_idx, te_idx, device, use_llm)
print(f"{name:9s} Total={tot/1e6:8.2f}M Trainable={trn/1e6:7.3f}M "
f"({100*trn/tot:.3f}%) PeakMem={mem:5.2f}GB Latency/LNP={lat:6.1f}ms")
del model; torch.cuda.empty_cache()
if __name__ == "__main__":
main()