laya-browser / code /apps /fast_batch.py
cklxx's picture
laya-browser v10 / v10s: laya fine-tuned as a browser-agent decision head + code + results
adf912b verified
Raw History Blame Contribute Delete
5.29 kB
"""One-shot batched predict for laya: tokenize the shared state once, build every question's sequence from cached ids,
single forward, vectorised post-processing. Same outputs as agent.predict() (up to float rounding).
from fast_batch import predict_fast, profile_step
"""
import json, time
import numpy as np, torch
from laya.common import QTYPES, collate_items, confidence_from_probs, render_options, serialize_state, temp_bucket
def build_items(agent, state, questions):
tok = agent.tok
max_len, head_max_len = agent.cfg.get("max_len", 512), agent.cfg.get("head_max_len", 192)
mask_tok, mask_id = tok.mask_token, tok.mask_token_id
st_ids = None # shared state tokens, computed lazily once
items, meta = [], []
for qid, qdef in questions.items():
q = agent._to_internal(qdef)
opts = render_options(q)
ins = str(q["ins"]).replace(mask_tok, " ")
head_ids = tok("%s question: %s" % (q["t"], ins), add_special_tokens=False)["input_ids"]
opt_txt = [" " + o.replace(mask_tok, " ") for o in opts]
opt_enc = tok(opt_txt, add_special_tokens=False)["input_ids"] # one batched tokenizer call for all options
opt_ids = [[mask_id] + o[:48] for o in opt_enc]
opt_budget = head_max_len - sum(len(o) for o in opt_ids)
if opt_budget < 16:
per = max(4, (head_max_len - 16) // max(1, len(opt_ids)))
opt_ids = [o[:per] for o in opt_ids]
opt_budget = head_max_len - sum(len(o) for o in opt_ids)
head_ids = head_ids[: max(8, opt_budget)]
ids = [tok.cls_token_id] + head_ids + [tok.sep_token_id]
markers = []
for o in opt_ids:
markers.append(len(ids)); ids.extend(o)
ids.append(tok.sep_token_id)
room = max(0, max_len - len(ids) - 1)
if st_ids is None:
st_ids = tok(serialize_state(state).replace(mask_tok, " "), add_special_tokens=False)["input_ids"]
ids = (ids + st_ids[:room] + [tok.sep_token_id])[:max_len]
markers = [m for m in markers if m < max_len]
if len(markers) != len(opts):
raise ValueError("question %r options exceed head_max_len=%d" % (qid, head_max_len))
items.append({"ids": ids, "markers": markers, "qtype": QTYPES[q["t"]]})
meta.append((qid, q, len(markers)))
return items, meta
@torch.no_grad()
def predict_fast(agent, state, questions, timing=None):
t0 = time.perf_counter()
items, meta = build_items(agent, state, questions)
b = collate_items([items], agent.tok.pad_token_id)
t1 = time.perf_counter()
dev = agent.device
with torch.autocast(device_type=dev.type, dtype=agent.dtype, enabled=dev.type == "cuda"):
logits, act = agent.model(b["input_ids"].to(dev, non_blocking=True), b["attention_mask"].to(dev, non_blocking=True),
b["marker_pos"].to(dev, non_blocking=True), b["marker_mask"].to(dev, non_blocking=True), b["qtype"].to(dev, non_blocking=True))
logits = logits.float().cpu().numpy(); act = torch.softmax(act.float(), -1).cpu().numpy()
t2 = time.perf_counter()
answers = {}
for r, (qid, q, k) in enumerate(meta):
qt = QTYPES[q["t"]]
t_scale = agent.temperature_by_options.get(temp_bucket(qt, k), agent.temperature[qt])
z = logits[r, :k] / max(1e-3, float(t_scale)); p = np.exp(z - z.max()); p /= p.sum()
conf = round(confidence_from_probs(p, k), 4); ext = {"act_probability": round(float(act[r, 0]), 4)}
if q["t"] == "choice":
keys = list(q["crit"].keys())
answers[qid] = {"type": "choice", "choice": keys[int(p.argmax())], "probabilities": {kk: round(float(v), 4) for kk, v in zip(keys, p)}, "confidence": conf, "action": ext}
elif q["t"] == "score":
answers[qid] = {"type": "score", "score": round(float((np.arange(k) * p).sum()), 4), "legend": {str(i): c for i, c in enumerate(q["crit"])},
"probabilities": {str(i): round(float(v), 4) for i, v in enumerate(p)}, "confidence": conf, "action": ext}
else:
answers[qid] = {"type": "noul", "noul": round(float(p[1]), 4), "confidence": round(max(float(p[1]), 1 - float(p[1])), 4), "action": ext}
t3 = time.perf_counter()
if timing is not None:
timing.update(tokenize_ms=(t1 - t0) * 1e3, forward_ms=(t2 - t1) * 1e3, post_ms=(t3 - t2) * 1e3, tokens=int(b["attention_mask"].sum()))
return {"model": "laya-rl-agent", "answers": answers, "usage": {"input_tokens": int(b["attention_mask"].sum()), "output_tokens": 0}}
def profile_step(agent, state, questions, n=20):
"""Compare agent.predict vs predict_fast on one recorded browser step."""
for _ in range(3): agent.predict(state, questions); predict_fast(agent, state, questions)
torch.cuda.synchronize(); t = time.perf_counter()
for _ in range(n): agent.predict(state, questions)
torch.cuda.synchronize(); slow = (time.perf_counter() - t) / n * 1e3
tm = {}; torch.cuda.synchronize(); t = time.perf_counter()
for _ in range(n): predict_fast(agent, state, questions, tm)
torch.cuda.synchronize(); fast = (time.perf_counter() - t) / n * 1e3
return {"predict_ms": slow, "predict_fast_ms": fast, **tm}