File size: 11,272 Bytes
0f8ae8e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
"""Low-latency inference engine: shape-bucketed CUDA graphs over the one-pass decision model.

Right padding + causal layers => pad positions never influence earlier slots, so no attention
mask is needed and every (B, T) bucket can be captured once and replayed.  The graph outputs
option-letter logits for all positions [B, T, K]; slots are gathered outside.
"""
import time, torch, torch._dynamo, torch.nn.functional as F
from decider.model import DecisionModel, collate
from decider.prompt import build, MAX_OPTIONS

T_BUCKETS = [64, 128, 192, 256, 320, 384, 512, 640, 768, 1024, 1280, 1536, 2048]
B_BUCKETS = [1, 2, 4, 8, 16, 32, 64]
GRAPH_MAX_T = 2048          # longer inputs (up to the 32k request budget) run eagerly: compute dominates there, and one graph
LONG_STEP = 1024            # per (B, T) shape would cost a compile + capture for every new length


def _bucket(x, buckets):
    for b in buckets:
        if x <= b:
            return b
    return None


def fused_causal_conv1d_fn(hidden_states, weight, bias=None, activation=None, **kwargs):
    """Depthwise causal conv (kernel k) as k shifted multiply-adds: fuses under torch.compile,
    unlike the cuDNN grouped conv fallback (which was ~11% of batched GPU time)."""
    B, C, T = hidden_states.shape; k = weight.shape[-1]
    x = F.pad(hidden_states.to(weight.dtype), (k - 1, 0))
    out = x[:, :, k - 1:k - 1 + T] * weight[:, k - 1][None, :, None]
    for j in range(k - 1):
        out = out + x[:, :, j:j + T] * weight[:, j][None, :, None]
    if bias is not None:
        out = out + bias[None, :, None]
    if activation == "silu":
        out = F.silu(out)
    elif activation is not None:
        from transformers.activations import ACT2FN
        out = ACT2FN[activation](out)
    return out.to(hidden_states.dtype)


def patch_conv():
    from transformers.models.qwen3_5 import modeling_qwen3_5 as mq
    mq.causal_conv1d_fn = fused_causal_conv1d_fn


def read_slots(out, rows, slots, nopts, temperature, n_per_item):
    """One gather + one softmax + one device-to-host copy for the whole batch (was: three small kernels and a sync per item).
    out [B, T, K] logits; rows/slots/nopts: flat python lists, one entry per question; n_per_item: questions per item."""
    dev = out.device; idx = torch.tensor([rows, slots, nopts], dtype=torch.long).to(dev, non_blocking=True)
    lg = out[idx[0], idx[1]]                                                                  # [N, K]
    lg = lg.masked_fill(torch.arange(lg.shape[1], device=dev)[None, :] >= idx[2][:, None], float("-inf"))
    p = torch.softmax(lg / temperature, -1).cpu()
    return list(torch.split(p, n_per_item))


def fill_ids(items_ids, B, T, pad):
    import numpy as np
    a = np.full((B, T), pad, dtype=np.int64)
    for b, x in enumerate(items_ids): a[b, :len(x)] = x
    return torch.from_numpy(a)


class Engine:
    """compile: torch.compile the forward (needs use_cache=False; ~1.4x batched, fuses elementwise work).
    fp8: e4m3 weights + per-token activation scaling on the big linears (Hopper tensor cores).
    conv_patch: fusable depthwise causal conv instead of the cuDNN fallback."""
    def __init__(self, path, device="cuda", dtype=torch.bfloat16, use_graphs=True, max_ctx_tokens=1536,
                 compile=True, fp8=False, conv_patch=True):
        if conv_patch:
            patch_conv()
        self.m = DecisionModel(path, dtype=dtype, grad_ckpt=False).to(device).eval()
        self.tok = self.m.tok; self.dev = device; self.use_graphs = use_graphs; self.max_ctx = max_ctx_tokens
        self.core, self.W = self.m.lm.model, self.m.lm.lm_head.weight[self.m.letters].detach().clone()
        self.cfg = dict(compile=compile, fp8=fp8, conv_patch=conv_patch, graphs=use_graphs)
        if fp8:
            from decider.fp8 import convert_to_fp8
            self.cfg["fp8_layers"] = convert_to_fp8(self.core)
        if compile:
            torch._dynamo.config.cache_size_limit = 128
            self._fwd_impl = torch.compile(self._fwd_eager, dynamic=False)
        else:
            self._fwd_impl = self._fwd_eager
        self.graphs = {}                       # (B, T) -> (static_ids, static_out, graph)
        self.pool = torch.cuda.graph_pool_handle() if use_graphs else None
        self.stats = dict(graph_captures=0, forwards=0)

    def _fwd_eager(self, ids):
        h = self.core(input_ids=ids, use_cache=False).last_hidden_state
        return F.linear(h, self.W).float()                          # [B, T, K]

    @torch.no_grad()
    def _fwd(self, ids):
        return self._fwd_impl(ids)

    def _capture(self, B, T):
        s_ids = torch.full((B, T), self.tok.pad_token_id, dtype=torch.long, device=self.dev)
        st = torch.cuda.Stream(); st.wait_stream(torch.cuda.current_stream())
        with torch.cuda.stream(st):
            for _ in range(3): self._fwd(s_ids)                     # warm-up: compile / triton autotune
        torch.cuda.current_stream().wait_stream(st)
        g = torch.cuda.CUDAGraph()
        with torch.cuda.graph(g, pool=self.pool):
            s_out = self._fwd(s_ids)
        self.stats["graph_captures"] += 1
        return s_ids, s_out, g

    @torch.no_grad()
    def logits_all(self, ids):
        """ids: [B, T] long on device (already right-padded to a bucket). Returns [B, T, K] float."""
        B, T = ids.shape; self.stats["forwards"] += 1
        if T > GRAPH_MAX_T:
            self.stats["long_forwards"] = self.stats.get("long_forwards", 0) + 1
            return self._fwd_eager(ids)
        if not self.use_graphs:
            return self._fwd(ids)
        key = (B, T)
        if key not in self.graphs:
            self.graphs[key] = self._capture(B, T)
        s_ids, s_out, g = self.graphs[key]
        s_ids.copy_(ids); g.replay()
        return s_out

    @torch.no_grad()
    def score_items(self, items, temperature=1.0):
        """items: list of dicts from prompt.build. Returns list of [n_q, MAX_OPTIONS] prob tensors (cpu)."""
        Tmax = max(len(it["ids"]) for it in items)
        T = _bucket(Tmax, T_BUCKETS) or -(-Tmax // LONG_STEP) * LONG_STEP
        B = (_bucket(len(items), B_BUCKETS) or len(items)) if T <= GRAPH_MAX_T else len(items)
        ids = fill_ids([it["ids"] for it in items], B, T, self.tok.pad_token_id)
        out = self.logits_all(ids.to(self.dev, non_blocking=True))
        return read_slots(out, [b for b, it in enumerate(items) for _ in it["slots"]], [s for it in items for s in it["slots"]],
                          [n for it in items for n in it["nopts"]], temperature, [len(it["slots"]) for it in items])

    @torch.no_grad()
    def score_shared(self, items, temperature=1.0, min_prefix=192):
        """Rows that start with the same tokens (one state, one question per row): run the shared prefix once, fork its
        cache (attention KV + delta-net conv/recurrent states) to every row, and run only the question suffixes.
        Same answers as score_items up to kernel round-off; cost ~ state + sum(questions) instead of n * state."""
        ids = [it["ids"] for it in items]; n = len(ids)
        lcp = 0; short = min(len(x) for x in ids) - 1
        while lcp < short and all(x[lcp] == ids[0][lcp] for x in ids): lcp += 1
        if n < 2 or lcp < min_prefix:
            return self.score_items(items, temperature)
        self.stats["shared_prefix_calls"] = self.stats.get("shared_prefix_calls", 0) + 1
        pre = torch.tensor(ids[0][:lcp], device=self.dev)[None]
        cache = self.core(input_ids=pre, use_cache=True).past_key_values
        cache.reorder_cache(torch.zeros(n, dtype=torch.long, device=self.dev))            # fork: every row gets a copy of row 0
        Ts = max(len(x) for x in ids) - lcp
        suf = fill_ids([x[lcp:] for x in ids], n, Ts, self.tok.pad_token_id)
        h = self.core(input_ids=suf.to(self.dev), past_key_values=cache, use_cache=True).last_hidden_state
        rows = [b for b, it in enumerate(items) for _ in it["slots"]]; sl = [s - lcp for it in items for s in it["slots"]]
        idx = torch.tensor([rows, sl], device=self.dev)
        return read_slots(F.linear(h[idx[0], idx[1]], self.W).float()[:, None, :], list(range(len(rows))), [0] * len(rows),
                          [n for it in items for n in it["nopts"]], temperature, [len(it["slots"]) for it in items])

    def warmup(self, shapes=((1, 128), (1, 256), (1, 384), (1, 512), (8, 256), (8, 512), (32, 256), (32, 512))):
        t = time.time()
        for B, T in shapes:
            self.logits_all(torch.full((B, T), self.tok.pad_token_id, dtype=torch.long, device=self.dev))
        torch.cuda.synchronize(); return time.time() - t


if __name__ == "__main__":
    import sys, random, numpy as np
    from decider import data as D
    from decider.infer import Decider
    path = sys.argv[1] if len(sys.argv) > 1 else "runs/r3_v2/model"
    cfg = dict(compile="nocompile" not in sys.argv[2:], fp8="fp8" in sys.argv[2:], conv_patch="noconv" not in sys.argv[2:])
    _, evals = D.load_cache("data/tasks.pkl")
    eng = Engine(path, **cfg); print("engine cfg", eng.cfg)
    rng = random.Random(0)
    exs = evals["support_tickets"][:64] + evals["clinc_oos"][:64] + evals["race"][:32]
    items = [build(e, eng.tok, rng, max_ctx_tokens=1536) for e in exs]
    # correctness vs eager masked forward (DecisionModel.slot_logits)
    ref = []
    with torch.no_grad():
        for i in range(0, len(items), 16):
            b = collate(items[i:i + 16], eng.tok.pad_token_id)
            lg = eng.m.slot_logits(b["input_ids"].cuda(), b["attention_mask"].cuda(), b["slot_idx"].cuda(), b["slot_batch"].cuda(), b["nopts"].cuda())
            ref.append(torch.softmax(lg, -1).cpu())
    ref = torch.cat(ref)
    got = torch.cat(eng.score_items(items))
    print(f"max |p_graph - p_eager| = {(ref - got).abs().max():.4f} over {len(ref)} questions; argmax agreement {(ref.argmax(1) == got.argmax(1)).float().mean():.4f}")
    print(f"warmup capture of 8 buckets: {eng.warmup():.1f}s; captures so far {eng.stats['graph_captures']}")
    # latency: single real requests
    for name, pool in [("support_tickets", exs[:64]), ("clinc_oos", exs[64:128]), ("race", exs[128:])]:
        its = [build(e, eng.tok, rng) for e in pool]
        ts = []
        for it in its[:40]:
            torch.cuda.synchronize(); t = time.time(); eng.score_items([it]); torch.cuda.synchronize(); ts.append(time.time() - t)
        ts = np.array(ts[5:]) * 1000
        print(f"single request {name:16s}: p50 {np.median(ts):5.1f} ms  p90 {np.percentile(ts, 90):5.1f} ms  (avg {np.mean([len(i['ids']) for i in its]):.0f} tok, {len(its[0]['slots'])} q)")
        for bs in (8, 32):
            ts = []
            for i in range(0, min(len(its), bs * 6), bs):
                chunk = its[i:i + bs]
                if len(chunk) < bs: break
                torch.cuda.synchronize(); t = time.time(); eng.score_items(chunk); torch.cuda.synchronize(); ts.append(time.time() - t)
            ts = np.array(ts[1:]) * 1000
            print(f"   batch {bs:2d}: p50 {np.median(ts):6.1f} ms -> {bs/np.median(ts)*1000:6.0f} ctx/s, {bs*len(its[0]['slots'])/np.median(ts)*1000:6.0f} decisions/s")
    print("stats", eng.stats, "graphs", len(eng.graphs), f"mem {torch.cuda.memory_reserved()/1e9:.1f} GB")