File size: 3,896 Bytes
454b3e6
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
"""端到端:原版 (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")