File size: 2,150 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
import os, sys, torch, time
sys.path.insert(0, "apps"); sys.path.insert(0, "kernels")
from common import get_agent; from fast_laya import accelerate
from laya.common import build_sequence, collate_items, QTYPES
agent = get_agent("multilingual"); fast = accelerate(agent, use_graphs=False)
Q = {"department": {"type": "choice", "instructions": "Which team should handle this?", "criteria": {"billing": "invoices, refunds", "technical": "bugs, outages", "sales": "pricing"}},
     "urgency": {"type": "score", "instructions": "How urgent is this?", "criteria": ["not urgent", "soon", "blocking"]},
     "churn": {"type": "noul", "instructions": "Does the user threaten to cancel?"}}
qs = {f"{k}{i}": v for i in range(10) for k, v in Q.items()}
for name, st in [("short", {"body": "We were billed twice for March. Please refund the duplicate."}), ("long", {"body": "Since yesterday our whole team cannot log in, the dashboard returns 502 errors. " * 50})]:
    items = []
    for qid in qs:
        qq = agent._to_internal(qs[qid]); seq, m = build_sequence(agent.tok, st, qq, agent.cfg["max_len"], agent.cfg["head_max_len"]); items.append({"ids": seq, "markers": m, "qtype": QTYPES[qq["t"]]})
    b = {k: v.cuda() for k, v in collate_items([items], agent.tok.pad_token_id).items() if torch.is_tensor(v)}
    def fwd():
        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"])
    for _ in range(3): fwd()
    from torch.profiler import profile, ProfilerActivity
    with profile(activities=[ProfilerActivity.CUDA]) as prof:
        for _ in range(5): fwd()
        torch.cuda.synchronize()
    print(f"\n##### {name}: N={b['input_ids'].shape[0]} L0={b['input_ids'].shape[1]}")
    ka = prof.key_averages()
    tot = sum(e.self_device_time_total for e in ka)
    for e in sorted(ka, key=lambda e: -e.self_device_time_total)[:14]:
        if e.self_device_time_total > 0: print(f"{e.self_device_time_total/5/1000:8.3f} ms {100*e.self_device_time_total/tot:5.1f}%  {e.key[:90]}")
    print(f"{tot/5/1000:8.3f} ms total GPU per forward")