"""端到端:原版 (HF eager/sdpa + autocast) vs TileLang fast path. 数值一致性 + 延迟对比.""" import os, sys, time, torch, numpy as np sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "apps")); sys.path.insert(0, os.path.dirname(__file__)) from common import get_agent from fast_laya import accelerate, restore from laya.common import build_sequence, collate_items, QTYPES agent = get_agent(os.environ.get("LAYA_VARIANT", "multilingual")) Q = {"department": {"type": "choice", "instructions": "Which team should handle this?", "criteria": {"billing": "invoices, refunds", "technical": "bugs, outages", "sales": "pricing", "shipping": "delivery"}}, "urgency": {"type": "score", "instructions": "How urgent is this?", "criteria": ["not urgent", "soon", "blocking"]}, "churn": {"type": "noul", "instructions": "Does the user threaten to cancel?"}} def qs(n): return {f"{k}{i}": v for i in range(n) for k, v in Q.items()} short = {"subject": "Duplicate charge on invoice 4411", "body": "We were billed twice for March. Please refund the duplicate or we're moving to a competitor."} long_ = {"subject": "Outage report", "body": ("Since yesterday our whole team cannot log in, the dashboard returns 502 errors and our release is blocked. " * 40)} def batch(state, q): items = [] for qid in q: qq = agent._to_internal(q[qid]); seq, m = build_sequence(agent.tok, state, qq, agent.cfg["max_len"], agent.cfg["head_max_len"]) items.append({"ids": seq, "markers": m, "qtype": QTYPES[qq["t"]]}) return {k: v.cuda() for k, v in collate_items([items], agent.tok.pad_token_id).items() if torch.is_tensor(v)} def run(b): with torch.no_grad(), torch.autocast("cuda", dtype=agent.dtype): return agent.model(b["input_ids"], b["attention_mask"], b["marker_pos"], b["marker_mask"], b["qtype"]) def timeit(fn, iters=30): for _ in range(3): fn() torch.cuda.synchronize(); t = time.perf_counter() for _ in range(iters): fn() torch.cuda.synchronize(); return (time.perf_counter() - t) / iters * 1000 # fp32 ground truth agent.dtype_bak = agent.dtype def run_fp32(b): with torch.no_grad(): return agent.model(b["input_ids"], b["attention_mask"], b["marker_pos"], b["marker_mask"], b["qtype"]) cases = [("short x1", short, {"department": Q["department"]}), ("short x3", short, Q), ("short x30", short, qs(10)), ("long x3", long_, Q), ("long x30", long_, qs(10))] print("=== numerics (probabilities), fp32 vs original-bf16 vs fast-bf16") rows = [] for name, st, q in cases: b = batch(st, q) restore(agent) l32, a32 = run_fp32(b); lo, ao = run(b) t_orig = timeit(lambda: run(b)) accelerate(agent) lf, af = run(b) t_fast = timeit(lambda: run(b)) accelerate(agent, use_graphs=False); t_nog = timeit(lambda: run(b)) P = lambda l: torch.softmax(l.float(), -1) d_orig = (P(lo) - P(l32)).abs().max().item(); d_fast = (P(lf) - P(l32)).abs().max().item(); d_of = (P(lf) - P(lo)).abs().max().item() agree = (lf.argmax(-1) == l32.argmax(-1)).float().mean().item() print(f"{name:10s} L={b['input_ids'].shape[1]:4d} |p_orig-p_fp32|={d_orig:.4f} |p_fast-p_fp32|={d_fast:.4f} |p_fast-p_orig|={d_of:.4f} argmax agree={agree:.2f}") rows.append((name, b["input_ids"].shape[1], t_orig, t_nog, t_fast)) print("\n=== model forward latency (ms)") print(f"{'case':10s} {'L':>5s} {'original':>10s} {'tilelang':>10s} {'tl+graph':>10s} {'speedup':>8s}") for name, L, to, tn, tf in rows: print(f"{name:10s} {L:5d} {to:10.2f} {tn:10.2f} {tf:10.2f} {to/tf:7.1f}x") # end-to-end predict() print("\n=== end-to-end agent.predict() (ms, incl. tokenization)") for name, st, q in cases: restore(agent); to = timeit(lambda: agent.predict(st, q)); accelerate(agent); tf = timeit(lambda: agent.predict(st, q)) print(f"{name:10s} original={to:8.2f} fast={tf:8.2f} speedup={to/tf:5.1f}x")