laya-browser / code /kernels /prof_fast.py
cklxx's picture
v19s: WebChain real-site trajectories, format v5, webgym x7 + DAgger, harness fixes; replaces v17s
454b3e6 verified
Raw History Blame Contribute Delete
2.15 kB
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")