"""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}