File size: 5,294 Bytes
adf912b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
"""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}