diff --git a/python/sglang/multimodal_gen/runtime/layers/attention/backends/subblock_sparse/README.md b/python/sglang/multimodal_gen/runtime/layers/attention/backends/subblock_sparse/README.md new file mode 100644 index 000000000..3a5dcc4a4 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/layers/attention/backends/subblock_sparse/README.md @@ -0,0 +1,169 @@ +# SubBlock sparse attention — training-free block sparsity for the MiniMax-H3 DiT + +Routes FlashInfer's 64-token block-sparse kernel (`bsa_attn_blk64_fwd`) with a +sub-block score. Nothing is trained and no weights change: a cheap estimator +runs before attention and hands the kernel a `q2k_block_index`. + +Spelled out in full, with every key at its default — which is the recommended +configuration and what every number below was measured at: + +```bash +sglang serve --model-path MiniMaxAI/MiniMax-H3 --model-variant fl2va \ + --num-gpus 8 --ulysses-degree 8 --performance-mode speed \ + --attention-backend subblock_sparse_attn \ + --component-attention-backends text_encoder=fa \ + --attention-backend-config '{"sparsity": 0.75, "n_k": 4, "n_q": 4, + "skip_first_steps": 10, "skip_first_layers": 0, + "min_seq_len": 4096}' +``` + +**`text_encoder=fa` is not optional.** `--attention-backend` applies to every +component, and the Qwen3-VL text encoder admits only `fa` / `torch_sdpa` / +`sage_attn_3`; without the override it raises and the server never starts. Put +the override on the *encoder*, not the DiT — `transformer=subblock_sparse_attn` +appears to work and silently does nothing, because H3 resolves the DiT backend +lazily on the first forward, outside the component-loading context that the +override applies to. + +`--attention-backend-config` is optional and overrides only the keys it names, +so `'{"sparsity": 0.85}'` alone trades quality for another 6%. Inline JSON gets +mangled by `shlex.split`; pass a **file path** instead if the shell eats the +quotes. + +## What it runs on + +Everything below comes from `bsa_attn_blk64_fwd`, not from this backend. + +| | | +| --- | --- | +| GPU | **compute capability 10.0 only** — B200 / GB200 class. The kernel is built `-gencode=arch=compute_100a,code=sm_100a`, which is arch-specific and does not forward-run on 10.3 (B300 / GB300) or 12.x (RTX PRO 6000, RTX 50xx). | +| dtype | bfloat16 | +| head_dim | 128 | +| attention | non-causal, one contiguous sequence per call | + +Inside the DiT, anything the kernel cannot serve — cross attention, the token +refiner, sequences under `min_seq_len`, non-bf16 activations, head_dim != 128 — +falls back to dense for that call, so no layer has to be excluded by hand. + +**On an unsupported GPU it is not a fallback, it is an error at startup.** The +resolver checks the compute capability before anything loads and refuses +anything but 10.0, so an H100 or a B300 fails at launch rather than after ten +dense denoise steps. Do not rely on the kernel's own guard for this: it compares +only the major version, so it would accept 10.3 and then fail with no cubin. + +## How the score works + +The usual proxy for a 64x64 block is `mean(Q_block) · mean(K_block)`. Averaging +64 keys into one vector destroys exactly the variation that decides which keys a +query wants. So each block is cut into sub-blocks — `n_k` on the key side, `n_q` +on the query side — and every sub-block pair is scored and combined with a +log-sum-exp: + +``` +score(i, j) = log Σ_{a,b} exp( mean(Q_{i,a}) · mean(K_{j,b}) · softmax_scale ) +``` + +which estimates the block's un-normalised softmax mass directly — the quantity +that says how much is lost by skipping the block. + +Splitting the query side *alone* is worse than not splitting: a block's mass sums +over its query rows, so with one key vector to score against the query detail +averages out. Splitting both together is a different proposition, and the only +estimator change in this family that has held up end to end. `router.py` carries +the recall table behind `n_q = n_k = 4` and the record of what was tried and +rejected. + +## Configuration + +| key | default | meaning | +| --- | ---: | --- | +| `sparsity` | 0.75 | key blocks dropped per query block, as an upper bound | +| `n_k` | 4 | key sub-blocks per 64-token block (1, 2, 4, 8) | +| `n_q` | 4 | query sub-blocks per 64-token block (1, 2, 4, 8) | +| `skip_first_steps` | 10 | leading denoise forwards kept dense | +| `skip_first_layers` | 0 | leading DiT blocks kept dense | +| `min_seq_len` | 4096 | shorter sequences run dense | + +**`sparsity` is an upper bound, not an exact figure.** The kernel pads each query +row's block count up to a multiple of 8 with phantom slots it then masks out, so +148 blocks costs exactly what 152 costs; the router takes the 152. At 590 blocks, +0.75 requested delivers 0.7424, and the startup log reports what was kept. It is +the speed lever — see below — and the only knob most users should touch. + +**`n_k` and `n_q` buy score accuracy, not speed.** They set how finely a block is +cut before scoring: `n_k=4` means four 16-token key sub-blocks, and the block's +score is the log-sum-exp over all `n_q * n_k` sub-block pairs. Raising them +sharpens the estimate of which blocks carry mass, at `n_q * n_k` times the score +matrix — 0.5% of denoise time at the default, so cost is not the constraint. +Raise `n_q` and `n_k` **together**: splitting the query side alone is worse than +not splitting at all. + +**The two schedule cutoffs are asymmetric on purpose.** `skip_first_steps` keeps +the leading denoise forwards dense; those steps settle the layout of the sample +and visibly re-frame the shot when approximated — lowering it from 10 to 5 halves +cosine against the dense render. Depth does not behave that way, so +`skip_first_layers` defaults to 0 and every DiT layer is sparse. Do not lower +`skip_first_steps` without looking at the output. + +**`min_seq_len` is a floor, not a tuning knob.** Below it the whole call runs +dense, and in packed varlen batches the test is per document, so H3's padding +tail goes dense while the 37.7k-token media document is routed. Two things break +down on short sequences: the router is four fixed Triton launches against an +attention cost that falls as S², so the overhead stops paying for itself; and the +budget goes coarse — 4096 keys is only 64 blocks, and at 1024 keys the +multiple-of-8 floor already keeps half of them. 4096 sits well below any real +video sequence and well above where either effect bites; it was chosen on that +reasoning rather than from a measured threshold sweep. + +## Measured + +MiniMax-H3 t2va, 1344x768 / 5 s / 50 steps, 8x B200, Ulysses-8, bf16, at the +shipped defaults (152 of 590 key blocks per query block). All arms in one session +on one node, cold sample dropped; spread within an arm is under 0.07 s. + +| | DiT denoise | vs dense | +| --- | ---: | ---: | +| dense (FlashAttention) | 18.270 s | 1.000x | +| SubBlock sparse | 16.061 s | **1.138x** | +| SubBlock sparse + [flashinfer#4397][fi] | 15.012 s | **1.217x** | + +[flashinfer-ai/flashinfer#4397][fi] rebuilds the kernel's internal Q/K/V tile +layout in one pass instead of three. It is bit-identical and **not required**: +worth 1.070x on its own. + +[fi]: https://github.com/flashinfer-ai/flashinfer/pull/4397 + +Sparsity is the speed lever and it saturates — 0.75 gives 1.136x, 0.80 gives +1.178x, 0.85 gives 1.211x. Cutting the budget 40% past 0.75 buys 6%, because at +37.7k tokens attention is no longer the bulk of the step, and 0.85 rendered worst +of the three on cosine against dense. `n_k` moves the denoise time by 0.3% across +its whole range: it is a quality knob, not a speed one. + +**The speedup is bounded by sequence length, not by the method.** The same config +measured 1.13x at 37.7k tokens, 1.20x at 52k and 1.47x at 96k — the backend only +touches attention, and attention's share of the DiT grows with S. Treat 1.2x as +the 768p/5 s number, not the ceiling. + +The same effect shows up in the sequence-parallel degree, since that sets how +much of the sequence each GPU holds: on 4x B200 at Ulysses-4 the identical +config gives **1.168x** on denoise and **1.138x** end to end, against 1.138x on +denoise at Ulysses-8. + +Peak memory is unchanged (99,356 vs 99,358 MiB/GPU): block sparsity saves +compute, not activations, and the `[B,H,Gq,Gk]` score matrix is ~20 MB at +S=37.7k. Absolute times are node-specific; only ratios measured in one session +are comparable. + +## Files + +| | | +| --- | --- | +| `router.py` | `SubBlockRouter` — pooling, scoring, selection, `RoutingPlan` | +| `kernels.py` | Triton pooling / segmented-LSE / fused top-k | +| `../subblock_sparse_attn.py` | the `AttentionBackend`: schedule, gating, dense fallback | + +Tests: `test/unit/test_subblock_sparse_attention.py`. The trick that makes the sparse +kernel checkable against dense is running it at a sparsity just above 0 — every +block is then inside the budget, so the result must reproduce dense attention up +to bf16 rounding, which pins the routing indices, the ragged tail block sizes +and the softmax scale in one assertion. diff --git a/python/sglang/multimodal_gen/runtime/layers/attention/backends/subblock_sparse/__init__.py b/python/sglang/multimodal_gen/runtime/layers/attention/backends/subblock_sparse/__init__.py new file mode 100644 index 000000000..312eddeff --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/layers/attention/backends/subblock_sparse/__init__.py @@ -0,0 +1,15 @@ +# SPDX-License-Identifier: Apache-2.0 +"""SubBlock -- training-free block-sparse attention routing for video DiTs. + +Originally vendored from the standalone SubBlock repository; ``router.py`` and +``kernels.py`` have since diverged from it. + +``router.py`` scores every (query block, key block) pair from sub-block-pooled +Q/K and turns the scores into the ``q2k_block_index`` that FlashInfer's +``bsa_attn_blk64_fwd`` consumes (SM100, bf16, head_dim 128). The estimator and +the measurements behind its defaults are documented there. +""" + +from .router import SubBlockRouter, load_bsa_attn_blk64_fwd + +__all__ = ["SubBlockRouter", "load_bsa_attn_blk64_fwd"] diff --git a/python/sglang/multimodal_gen/runtime/layers/attention/backends/subblock_sparse/kernels.py b/python/sglang/multimodal_gen/runtime/layers/attention/backends/subblock_sparse/kernels.py new file mode 100644 index 000000000..17a71ed4e --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/layers/attention/backends/subblock_sparse/kernels.py @@ -0,0 +1,300 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Fused sub-block score kernel: bf16 GEMM + segmented log-sum-exp in one pass. + + score[i, j] = log2 sum_{b < n_k} 2 ** ( qbar_i . kbar_{j,b} ) + +The PyTorch router materialises ``[B, H, Gq, Gk*n_k]`` fp32 (254 MB at S=96k, n_k=4), +reduces it, and throws it away -- the GEMM is 4% of its time, the rest is that tensor's +round trip to HBM. Here the reduction happens in registers before anything is written, so +only ``[B, H, Gq, Gk]`` ever reaches memory: n_k times less traffic. + +exp2/log2 are used internally (they are the hardware instructions; the caller folds +``softmax_scale * log2(e)`` into Q), and the result is converted back to natural-log units +so it matches the reference implementation exactly, not just up to ranking. + +Padding: sub-cells are ordered, and validity is monotone, so a single ``n_valid`` scalar +(the number of key sub-cells holding at least one real token) is enough -- everything at or +past it is forced to -inf so it can never win a slot. +""" + +import math + +import torch +import triton +import triton.language as tl + +_NEG = tl.constexpr(-1.0e30) # Triton only lets @jit read constexpr globals +_LN2 = tl.constexpr(0.6931471805599453) + + +@triton.jit +def _score_kernel( + Q, + K, + O, + stride_qm, + stride_ql, + stride_kn, + stride_kl, + stride_om, + stride_on, + stride_ol, + M, + M_VALID, + N_VALID, + NOUT, + MOUT, + BLK_M: tl.constexpr, + BLK_N: tl.constexpr, + NK: tl.constexpr, + NQR: tl.constexpr, + D: tl.constexpr, +): + pid_m = tl.program_id(0) + pid_n = tl.program_id(1) + pid_l = tl.program_id(2) + + offs_m = pid_m * BLK_M + tl.arange(0, BLK_M) + offs_n = pid_n * BLK_N + tl.arange(0, BLK_N) + offs_d = tl.arange(0, D) + + q = tl.load( + Q + pid_l * stride_ql + offs_m[:, None] * stride_qm + offs_d[None, :], + mask=offs_m[:, None] < M, + other=0.0, + ) + k = tl.load( + K + pid_l * stride_kl + offs_n[:, None] * stride_kn + offs_d[None, :], + mask=offs_n[:, None] < N_VALID, + other=0.0, + ) + acc = tl.dot(q, tl.trans(k), out_dtype=tl.float32) # [BLK_M, BLK_N] + + # a sub-cell past the last real one must not contribute to its group's log-sum-exp + acc = tl.where(offs_n[None, :] < N_VALID, acc, _NEG) + # same on the query side: with NQ > 1 the last query block can own sub-cells that + # are entirely padding, and those pool to zero -- an exp2(0) = 1 term that would + # otherwise be folded into the block's score. + acc = tl.where(offs_m[:, None] < M_VALID, acc, _NEG) + + acc = tl.reshape(acc, (BLK_M, BLK_N // NK, NK)) + m = tl.max(acc, axis=2) + s = tl.sum(tl.exp2(acc - m[:, :, None]), axis=2) + lse = m + tl.log2(s) + lse = tl.where(m > _NEG / 2, lse, _NEG) # whole group was padding + + # Fold the NQR query sub-cells of a query block together. Log-sum-exp is + # associative, so reducing NK then NQ is the same one log-sum-exp over all + # NQ*NK sub-block pairs -- and two stages keeps both reductions on an axis + # that is already contiguous in registers. + if NQR > 1: + lse = tl.reshape(lse, (BLK_M // NQR, NQR, BLK_N // NK)) + m2 = tl.max(lse, axis=1) + s2 = tl.sum(tl.exp2(lse - m2[:, None, :]), axis=1) + lse = tl.where(m2 > _NEG / 2, m2 + tl.log2(s2), _NEG) + + # exp2/log2 internally (they map to the hardware instructions), then back to natural + # log units so the fused and reference backends return the same numbers, not just the + # same ranking. One multiply in registers. + out = lse * _LN2 + out = out.to(O.dtype.element_ty) # bf16 halves what the selection step has to read + + offs_o = pid_n * (BLK_N // NK) + tl.arange(0, BLK_N // NK) + offs_q = pid_m * (BLK_M // NQR) + tl.arange(0, BLK_M // NQR) + tl.store( + O + + pid_l * stride_ol + + offs_q[:, None] * stride_om + + offs_o[None, :] * stride_on, + out, + mask=(offs_q[:, None] < MOUT) & (offs_o[None, :] < NOUT), + ) + + +BLK_M = BLK_N = 128 # score tile; must hold whole blocks, so a multiple of n_q and n_k + + +def fused_scores(qp, kp, out, *, n_k, n_valid, n_q, m_valid): + """qp: [L, Gq*n_q, D] bf16 (already carrying softmax_scale*log2e), + kp: [L, Gk*n_k, D] bf16 -> out: [L, Gq, Gk], natural-log scores. + + The n_q query sub-cells of a query block are folded together by log-sum-exp. + + ``n_valid`` / ``m_valid`` are the counts of key / query sub-cells holding at + least one real token; the rest pooled to zero and must not contribute. + """ + L, M, D = qp.shape + N = kp.shape[1] + Mout, Nout = out.shape[1], out.shape[2] + grid = (triton.cdiv(M, BLK_M), triton.cdiv(N, BLK_N), L) + _score_kernel[grid]( + qp, + kp, + out, + qp.stride(1), + qp.stride(0), + kp.stride(1), + kp.stride(0), + out.stride(1), + out.stride(2), + out.stride(0), + M, + m_valid, + n_valid, + Nout, + Mout, + BLK_M=BLK_M, + BLK_N=BLK_N, + NK=n_k, + NQR=n_q, + D=D, + num_warps=4, + num_stages=3, + ) + return out + + +@triton.jit +def _pool_kernel( + X, + Y, + stride_xb, + stride_xt, + stride_xh, + stride_yl, + stride_yn, + S, + H, + SUB: tl.constexpr, + D: tl.constexpr, + SCALE, +): + """[B, S, H, D] -> [B*H, n_cells, D]: masked mean of every SUB consecutive tokens. + + Fused so the activation is read once and the pooled result is written straight in + bf16; the PyTorch version needs an fp32 temporary plus a transpose. + """ + cell = tl.program_id(0) + l = tl.program_id(1) + b = l // H + h = l % H + offs_t = cell * SUB + tl.arange(0, SUB) + offs_d = tl.arange(0, D) + mask = offs_t < S + x = tl.load( + X + + b * stride_xb + + offs_t[:, None] * stride_xt + + h * stride_xh + + offs_d[None, :], + mask=mask[:, None], + other=0.0, + ).to(tl.float32) + cnt = tl.sum(mask.to(tl.float32), axis=0) + acc = tl.sum(x, axis=0) / tl.maximum(cnt, 1.0) * SCALE + tl.store(Y + l * stride_yl + cell * stride_yn + offs_d, acc.to(tl.bfloat16)) + + +def fused_pool(x, n_cells, sub, out, scale=1.0): + """x: [B, S, H, D] bf16 -> out: [B*H, n_cells, D] bf16""" + B, S, H, D = x.shape + _pool_kernel[(n_cells, B * H)]( + x, + out, + x.stride(0), + x.stride(1), + x.stride(2), + out.stride(0), + out.stride(1), + S, + H, + SUB=sub, + D=D, + SCALE=scale, + # one warp, not four: the tile is only SUB x 128, so extra warps buy no parallelism + # and cost scheduling. Measured 0.272 -> 0.074 ms at S=96k, 1.27 -> 4.7 TB/s. + num_warps=1, + ) + return out + + +@triton.jit +def _topk_kernel(S, OUT, G, K, BLK: tl.constexpr, ITERS: tl.constexpr): + """Exact-enough per-row top-K in a single pass over global memory. + + A score row is only Gk values -- 3 KB in bf16 at S=96k -- so it fits in registers. Load + it once, then do everything on chip. torch.topk instead makes several passes over the + whole matrix, which is why it costs half the router. + + The threshold search interpolates on the count rather than halving the interval: the + count-vs-threshold curve is the row's empirical CDF, so a secant step lands far closer + than a bisection step. Full-row reductions are what this kernel pays for, so fewer + steps is the whole game. + + The invariant is count(s >= lo) >= K, so the compaction can only over-fill, never + under-fill; `pos < K` truncates the boundary group. At ITERS=12 about 1.6% of rows end + up with a different set than an exact top-K, but only among blocks that tie at the + threshold -- the total selected score differs by ~1e-6 relative, which is nothing. + """ + row = tl.program_id(0) + offs = tl.arange(0, BLK) + m = offs < G + s = tl.load(S + row * G + offs, mask=m, other=-float("inf")).to(tl.float32) + lo = tl.min(tl.where(m, s, float("inf"))) + hi = tl.max(tl.where(m, s, -float("inf"))) + 1.0 + clo = tl.sum(m.to(tl.int32), axis=0).to(tl.float32) + chi = 0.0 + for _ in tl.static_range(ITERS): + den = clo - chi + t = (clo - K) / tl.where(den > 0.5, den, 1.0) + t = tl.minimum(tl.maximum(t, 0.05), 0.95) # keep the step inside the bracket + mid = lo + (hi - lo) * t + cnt = tl.sum(((s >= mid) & m).to(tl.int32), axis=0).to(tl.float32) + take = cnt >= K + lo = tl.where(take, mid, lo) + clo = tl.where(take, cnt, clo) + hi = tl.where(take, hi, mid) + chi = tl.where(take, chi, cnt) + sel = (s >= lo) & m + pos = tl.cumsum(sel.to(tl.int32), axis=0) - 1 + tl.store(OUT + row * K + pos, offs.to(tl.int32), mask=sel & (pos < K)) + + +def topk_iters(G, k): + """Threshold-search steps needed to match ``torch.topk``. + + The search interpolates on the *count* above a trial threshold, which assumes the count + is linear in the threshold. That holds near the median but not in the tail, where the + score density decays roughly exponentially and the secant undershoots -- so the further + into the tail k sits, the more steps are needed. Measured on 21 real files as the + fraction of query rows selecting a different block set than exact top-k: + + log2(G/k) 1.0 2.1 3.3 4.3 5.6 (sparsity .50 .76 .90 .95 .98) + 16 steps 0.09% 0.14% 1.62% 5.17% 1.88% + 24 steps 0.02% 0.01% 0.14% 0.61% 0.21% + 32 steps 0.02% 0.01% 0.02% 0.19% 0.05% + + A flat 16 is fine at the usual operating points and silently wrong past sparsity 0.9 + (+3.2% relative L2 at 0.90, +7.1% at 0.95). These cutoffs hold every regime under 0.2%. + Interpolating on log(count) instead linearises the tail and does fix sparsity >= 0.95, + but it is far worse where the tail model does not apply (27% differing rows at sparsity + 0.5), so the step count is the robust knob, not the model. + """ + L = math.log2(max(G, 1) / max(k, 1)) + return 16 if L <= 2.5 else 24 if L <= 3.5 else 32 + + +def fused_topk(scores2d, k): + """scores2d: [rows, G] contiguous -> [rows, k] int32 column ids (unsorted).""" + rows, G = scores2d.shape + out = torch.empty(rows, k, dtype=torch.int32, device=scores2d.device) + _topk_kernel[(rows,)]( + scores2d, + out, + G, + k, + BLK=triton.next_power_of_2(G), + ITERS=topk_iters(G, k), + num_warps=4 if G >= 1024 else 2, # short rows do not fill four warps + ) + return out diff --git a/python/sglang/multimodal_gen/runtime/layers/attention/backends/subblock_sparse/router.py b/python/sglang/multimodal_gen/runtime/layers/attention/backends/subblock_sparse/router.py new file mode 100644 index 000000000..f701b2640 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/layers/attention/backends/subblock_sparse/router.py @@ -0,0 +1,256 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Sub-block block-sparse routing for FlashInfer's ``bsa_attn_blk64_fwd``. + +Training-free. Runs *before* attention, produces the ``q2k_block_index`` tensor the +64-token block-sparse kernel consumes. + +Why sub-blocks +-------------- +The usual proxy score for a 64x64 block is ``mean(Q_block) . mean(K_block)``. Averaging +64 keys into one vector throws away exactly the variation that decides which keys a query +wants. Splitting each 64-token block into ``n`` sub-blocks of ``64/n`` tokens, scoring all +sub-block pairs and combining them with a log-sum-exp recovers most of that: + + score(i, j) = log sum_{a +.6779 and relative L2 of the rebuilt attention output .2043 -> .1982, paired t = +8.0 and +-5.7. + +It is worse in the pixels, on **0 of 15** prompts, by 0.107 cosine against the dense render +(paired t = -6.4). Single-layer output error, even measured directly, does not order these +estimators the way 40 denoise steps through 50 layers do. Nothing short of an end-to-end +render has predicted this correctly yet -- neither block mass recall nor single-step output +L2. + +Worth trying, not yet exposed +----------------------------- +A per-head budget beats any estimator upgrade measured here: at a fixed mean sparsity, +spending more blocks on diffuse heads and fewer on peaked ones lifts 5th-percentile mass +recall from .52 to .90. It needs a rule for setting the per-head split, which nothing in +the pipeline currently produces. + +Usage +----- + router = SubBlockRouter(n_k=4, n_q=4) + plan = router.route(q, k, sparsity=0.8, softmax_scale=d**-0.5) # q, k: [B, S, H, D] + out, _ = bsa_attn_blk64_fwd(q, k, v, plan.index, plan.topk, + block_sizes=SubBlockRouter.block_sizes(S, q.device), + q2k_block_nums=None) +""" + +from __future__ import annotations + +import functools +import importlib.util +import math +import sys +import types +from pathlib import Path + +import msgspec +import torch + +from .kernels import fused_pool, fused_scores, fused_topk + + +@functools.lru_cache(maxsize=1) +def load_bsa_attn_blk64_fwd(): + """FlashInfer's 64-block sparse attention entry point. + + ``flashinfer.cute_dsl.sparse.__init__`` also pulls in the blk128 CuTe-DSL backend, + which breaks in ways blk64 does not care about: it hard-requires the ``quack`` + package, and it tracks a moving ``cutlass.cute`` API (0.6.15.post1 raises + ``AttributeError: module 'cutlass.cute.core' has no attribute 'ThrMma'``). blk64 is + plain CUDA and needs none of it, so whatever the package import trips over we load + ``bsa_attn_blk64.py`` under a synthetic parent package instead -- same file, same + kernel. If blk64 itself is broken or absent, that load raises and the caller sees it. + """ + try: + from flashinfer.cute_dsl.sparse import bsa_attn_blk64_fwd + + return bsa_attn_blk64_fwd + except Exception: + pass + import flashinfer + + base = Path(flashinfer.__file__).resolve().parent / "cute_dsl" / "sparse" + pkg = "_flashinfer_sparse_blk64_only" + + def _load(name: str, path: Path, is_pkg: bool): + spec = importlib.util.spec_from_file_location( + name, + path, + submodule_search_locations=[str(path.parent)] if is_pkg else None, + ) + mod = importlib.util.module_from_spec(spec) + sys.modules[name] = mod + spec.loader.exec_module(mod) + return mod + + if pkg not in sys.modules: + parent = types.ModuleType(pkg) + parent.__path__ = [str(base)] + sys.modules[pkg] = parent + _load(f"{pkg}.blk64", base / "blk64" / "__init__.py", True) + mod = _load(f"{pkg}.bsa_attn_blk64", base / "bsa_attn_blk64.py", False) + return mod.bsa_attn_blk64_fwd + + +LOG2E = 1.4426950408889634 +BLOCK = 64 # the kernel's block granularity (kSparseBlockSize=64) +BUDGET_GRANULARITY = 8 # blocks per query row the kernel bills in, padding to fit +VALID_N = (1, 2, 4, 8) # sub-blocks per 64-token block -> 64 / 32 / 16 / 8 tokens + + +def _snap_up_to_8(topk: int, num_blocks: int) -> int: + """Round a block budget up to what the kernel is going to charge for anyway. + + ``bsa_attn_blk64_fwd`` pads each query row's block count up to a multiple of + ``BUDGET_GRANULARITY`` with phantom slots that repeat the last real block and + are then masked out of the softmax. Asking for 148 blocks therefore costs + exactly what 152 costs, with four of the slots computed and thrown away. + Measured at S=37.7k on B200, 2 prompts in one session: 152 blocks take + 16.055 s against 16.061 s for 148, and 120 take 15.490 s against 15.496 s + for 118 -- free, inside the noise. So take the blocks. + + The consequence for the caller is that ``sparsity`` is an upper bound rather + than an exact figure: 0.75 of 590 blocks becomes 152 kept, 0.7424 dropped. + """ + return min(num_blocks, max(1, -(-topk // BUDGET_GRANULARITY)) * BUDGET_GRANULARITY) + + +class RoutingPlan(msgspec.Struct, frozen=True): + """What the kernel needs, plus the budget that produced it.""" + + index: torch.Tensor # [B, H, Gq, topk] int32 + topk: int # key blocks kept per query block + num_blocks: int # key blocks available + + @property + def density(self) -> float: + return self.topk / self.num_blocks + + +class SubBlockRouter: + """Builds ``q2k_block_index`` from sub-block-pooled Q/K. + + Args: + n_k: key sub-blocks per 64-token block (1, 2, 4 or 8). 1 reproduces plain avg + pooling; 4 is the quality/cost point the recall table above lands on. + n_q: query sub-blocks, same values. Splitting Q *alone* (n_q>1 with n_k=1) is + worse than not splitting; splitting both sides together is what the default + does. Costs n_q times the score matrix, 0.5% of the denoise time. + + Structural block reservation (an attention sink, or forcing the diagonal j == i) was + measured on 200 real H3 attention cells and is deliberately absent: at a fixed budget + the diagonal changed relative L2 by 0.2% and the sink only helped in DiT layers 2-32, + which did not survive to the pixels. + """ + + def __init__(self, n_k: int = 4, n_q: int = 4) -> None: + if n_k not in VALID_N or n_q not in VALID_N: + raise ValueError( + f"n_q/n_k must be one of {VALID_N}, got n_q={n_q}, n_k={n_k}" + ) + self.n_k, self.n_q = n_k, n_q + + @torch.no_grad() + def scores( + self, q: torch.Tensor, k: torch.Tensor, softmax_scale: float + ) -> torch.Tensor: + """``[B, S, H, D] -> [B, H, Gq, Gk]`` block scores (log-space, higher = keep). + + Two Triton kernels: pool, then GEMM + segmented log-sum-exp in registers, so the + ``[B, H, Gq*n_q, Gk*n_k]`` intermediate never reaches memory. + + ``softmax_scale * log2(e)`` is folded into Q so the kernel can use the exp2/log2 + hardware instructions; it multiplies by ln 2 on the way out, so scores come back + in natural-log units. Selection is a top-k and any monotone rescale leaves that + alone, so the units only matter to a reader of the magnitudes. + + The scores stay **float32**. bf16 would halve what selection reads, but with 8 + mantissa bits many blocks tie exactly at the threshold and the fused selector + breaks ties by column index -- which systematically prefers early key blocks, one + region of the video. Measured +3.9% relative L2 at S=96k. + """ + b, s, h, d = q.shape + sk = k.shape[1] + gq, gk = -(-s // BLOCK), -(-sk // BLOCK) + nq, nk = self.n_q, self.n_k + sub_q, sub_k = BLOCK // nq, BLOCK // nk + + # Pooling handles the ragged tail on the *pooled* tensor: padding q/k up to + # G*BLOCK first would copy the whole 300+ MB activation to add a few rows. + # Sub-cells past the last real token pool to zero, and `*_valid` tells the score + # kernel to drop them -- left in, each would contribute an exp(0)=1 term that + # both inflates the score and flattens the differences the ranking depends on. + pooled_q = torch.empty(b * h, gq * nq, d, device=q.device, dtype=torch.bfloat16) + pooled_k = torch.empty(b * h, gk * nk, d, device=k.device, dtype=torch.bfloat16) + fused_pool(q, gq * nq, sub_q, pooled_q, scale=softmax_scale * LOG2E) + fused_pool(k, gk * nk, sub_k, pooled_k) + + out = torch.empty(b * h, gq, gk, device=q.device, dtype=torch.float32) + fused_scores( + pooled_q, + pooled_k, + out, + n_k=nk, + n_valid=-(-sk // sub_k), + n_q=nq, + m_valid=-(-s // sub_q), + ) + return out.view(b, h, gq, gk) + + @torch.no_grad() + def route( + self, q: torch.Tensor, k: torch.Tensor, sparsity: float, softmax_scale: float + ) -> RoutingPlan: + """Select the top ``(1 - sparsity)`` fraction of key blocks per query block.""" + b, s, h, d = q.shape + gk = -(-k.shape[1] // BLOCK) + scores = self.scores(q, k, softmax_scale) # [B, H, Gq, Gk] + gq = scores.shape[2] + topk = _snap_up_to_8(math.ceil((1.0 - sparsity) * gk), gk) + # One pass over the score matrix instead of torch.topk's several; the kernel + # accepts the blocks in any order, so nothing sorts them. + index = fused_topk(scores.reshape(-1, gk), topk).view(b, h, gq, topk) + return RoutingPlan(index=index, topk=topk, num_blocks=gk) + + @staticmethod + def block_sizes(seq_len: int, device) -> torch.Tensor: + """Real token count per 64-block, for the kernel's tail masking.""" + g = -(-seq_len // BLOCK) + start = torch.arange(g, device=device, dtype=torch.int32) * BLOCK + return (seq_len - start).clamp(0, BLOCK).to(torch.int32) diff --git a/python/sglang/multimodal_gen/runtime/layers/attention/backends/subblock_sparse_attn.py b/python/sglang/multimodal_gen/runtime/layers/attention/backends/subblock_sparse_attn.py new file mode 100644 index 000000000..f79328fa3 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/layers/attention/backends/subblock_sparse_attn.py @@ -0,0 +1,407 @@ +# SPDX-License-Identifier: Apache-2.0 +"""SubBlock block-sparse attention backend. + +Routes FlashInfer's 64-token block-sparse kernel with a K-side sub-block +log-sum-exp score (see ``backends/subblock_sparse/``). Everything is training-free: +the router runs before attention and produces the ``q2k_block_index`` the +kernel consumes. + +Sparsity is not applied everywhere. The early denoise steps settle the layout +of the sample and tolerate approximation badly, so the backend falls back to +dense attention for them. Depth turns out not to matter the same way, which is +why the layer cutoff defaults to zero -- see the defaults below. The schedule +is configured through ``--attention-backend-config``, which overrides +individual keys of the defaults below:: + + --attention-backend subblock_sparse_attn \ + --attention-backend-config '{"sparsity": 0.85}' + +Requirements inherited from the kernel: compute capability 10.0 (B200 / GB200 +class -- it is built for ``sm_100a``, which does not forward-run on 10.3 or +12.x), bf16, head_dim 128. Inside the DiT, any call the kernel cannot serve -- +cross/refiner attention, short sequences, non-bf16 -- runs dense instead. On any +other GPU the resolver refuses the backend at startup rather than falling back. + +``--attention-backend`` reaches every component, and the text encoder admits +only fa / torch_sdpa / sage_attn_3, so pair it with +``--component-attention-backends text_encoder=fa``; see the README. +""" + +from __future__ import annotations + +import functools +import re +from dataclasses import dataclass +from typing import Any + +import msgspec +import torch + +from sglang.multimodal_gen.runtime.layers.attention.backends.attention_backend import ( + AttentionBackend, + AttentionImpl, + AttentionMetadata, + AttentionMetadataBuilder, +) +from sglang.multimodal_gen.runtime.layers.attention.backends.subblock_sparse import ( + SubBlockRouter, + load_bsa_attn_blk64_fwd, +) +from sglang.multimodal_gen.runtime.managers.forward_context import get_forward_context +from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum +from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger + +logger = init_logger(__name__) + +# The kernel is fixed at 64-token blocks and 128-wide heads. +SUBBLOCK_SPARSE_HEAD_DIM = 128 + +# Defaults for the schedule; override through --attention-backend-config. +# Sparsity is the speed lever, and it saturates: on MiniMax-H3 t2va at 37.7k +# tokens, 0.75 gives 1.14x, 0.80 gives 1.18x and 0.85 gives 1.21x -- cutting the +# block budget by 40% past 0.75 buys 6%, because attention is no longer the bulk +# of the step. 0.85 was the worst arm on cosine-vs-dense on both clips rendered +# across all three grades, and 0.80 costs 0.017 / 0.006 cos_c against 0.75 on +# those same two clips for 3.5% of the time, so the default takes the quality. +DEFAULT_SPARSITY = 0.75 +# The two cutoffs were swept independently on MiniMax-H3 t2va (1344x768, 5 s, +# 50 steps, n_k=4, sparsity 0.75) and behave nothing alike. Lowering the step +# cutoff from 10 to 5 halves cosine-vs-dense (0.558 -> 0.310 on two clips) and +# visibly re-frames the shot; going to 0 leaves the sample essentially +# uncorrelated with dense for 1.20x -> 1.30x. Lowering the layer cutoff from 2 +# to 0 costs 0.0013 of that cosine -- inside the 0.02 run-to-run noise floor -- +# and is worth ~1%, so the first DiT blocks get no special treatment. +DEFAULT_SKIP_FIRST_STEPS = 10 +DEFAULT_SKIP_FIRST_LAYERS = 0 +DEFAULT_N_K = 4 +# Query-side splitting. Splitting Q *alone* is worse than not splitting -- with +# one key vector to score against, the query detail averages out -- which is +# where the "n_q is worthless" reading came from. Splitting both sides together +# is a different estimator: the log-sum-exp then runs over query-key sub-block +# pairs. It is the only estimator change in this family that has reproduced end +# to end, and it costs 0.5% of the denoise time. Measured against n_q=1 on +# fifteen t2va prompts, every arm rendered in one session against that session's +# own dense render, as cosine of the decoded video: +# sparsity 0.90 +0.062 paired t = +2.6 better on 13/15 +# sparsity 0.75 +0.008 paired t = +2.1 better on 10/15 +# The margin shrinks as the budget loosens, which is the pattern every estimator +# comparison here has followed: at the shipped 148 of 590 blocks the rules mostly +# agree on what to keep. +DEFAULT_N_Q = 4 +# Below this many keys the router costs more than the blocks it saves, and the +# top-k budget collapses to a handful of blocks. +DEFAULT_MIN_SEQ_LEN = 4096 + +# ``blocks..attn`` is a DiT layer; ``token_refiner.blocks..attn`` and +# anything else is not and stays dense. +_DIT_LAYER_PREFIX = re.compile(r"^blocks\.(\d+)\.") + + +def _dit_layer_index(prefix: str) -> int | None: + match = _DIT_LAYER_PREFIX.match(prefix) + return int(match.group(1)) if match else None + + +@functools.lru_cache(maxsize=8) +def _cached_block_sizes(seq_len: int, device: torch.device) -> torch.Tensor: + """Per-block real token counts; identical for every layer and step. + + Rebuilding it per call costs an arange plus a clamp launch on the critical + path for a tensor that only depends on the sequence length. + """ + return SubBlockRouter.block_sizes(seq_len, device) + + +class SubBlockSparseAttentionBackend(AttentionBackend): + + @staticmethod + def get_supported_head_sizes() -> list[int]: + return [SUBBLOCK_SPARSE_HEAD_DIM] + + @staticmethod + def get_enum() -> AttentionBackendEnum: + return AttentionBackendEnum.SUBBLOCK_SPARSE_ATTN + + @staticmethod + def get_impl_cls() -> type[SubBlockSparseAttentionImpl]: + return SubBlockSparseAttentionImpl + + @staticmethod + def get_metadata_cls() -> type[SubBlockSparseAttentionMetadata]: + return SubBlockSparseAttentionMetadata + + @staticmethod + def get_builder_cls() -> type[SubBlockSparseAttentionMetadataBuilder]: + return SubBlockSparseAttentionMetadataBuilder + + +@dataclass +class SubBlockSparseAttentionMetadata(AttentionMetadata): + current_timestep: int + + +class SubBlockSparseAttentionMetadataBuilder(AttentionMetadataBuilder): + # The base class declares __init__ abstract, so a builder that does not + # override it cannot be instantiated at all. + def __init__(self) -> None: + pass + + def prepare(self) -> None: + pass + + def build( # type: ignore[override] + self, current_timestep: int, **kwargs: dict[str, Any] + ) -> SubBlockSparseAttentionMetadata: + return SubBlockSparseAttentionMetadata(current_timestep=current_timestep) + + +class SubBlockSparseSchedule(msgspec.Struct, frozen=True): + """When sparsity is allowed to apply, and how much of it.""" + + sparsity: float + skip_first_steps: int + skip_first_layers: int + n_k: int + n_q: int + min_seq_len: int + + @classmethod + def from_server_args(cls) -> SubBlockSparseSchedule: + from sglang.multimodal_gen.runtime.server_args import get_global_server_args + + config = get_global_server_args().attention_backend_config or {} + schedule = SubBlockSparseSchedule( + sparsity=float(config.get("sparsity", DEFAULT_SPARSITY)), + skip_first_steps=int( + config.get("skip_first_steps", DEFAULT_SKIP_FIRST_STEPS) + ), + skip_first_layers=int( + config.get("skip_first_layers", DEFAULT_SKIP_FIRST_LAYERS) + ), + n_k=int(config.get("n_k", DEFAULT_N_K)), + n_q=int(config.get("n_q", DEFAULT_N_Q)), + min_seq_len=int(config.get("min_seq_len", DEFAULT_MIN_SEQ_LEN)), + ) + if not 0.0 <= schedule.sparsity < 1.0: + raise ValueError( + f"subblock sparsity must be in [0, 1), got {schedule.sparsity}" + ) + for name, value in (("n_k", schedule.n_k), ("n_q", schedule.n_q)): + if value not in (1, 2, 4, 8): + raise ValueError(f"subblock {name} must be 1, 2, 4 or 8, got {value}") + if schedule.skip_first_steps < 0 or schedule.skip_first_layers < 0: + raise ValueError("subblock skip_first_* must be non-negative") + return schedule + + +class SubBlockSparseAttentionImpl(AttentionImpl): + """Block-sparse attention with a dense fallback for the excluded region. + + One impl instance is built per attention module, so ``prefix`` fixes the + layer for the lifetime of the object; only the denoise step varies per + call and it comes from the forward context. + """ + + def __init__( + self, + num_heads: int, + head_size: int, + causal: bool = False, + softmax_scale: float | None = None, + num_kv_heads: int | None = None, + prefix: str = "", + **extra_impl_args, + ) -> None: + self.prefix = prefix + self.num_heads = num_heads + self.head_size = head_size + self.causal = causal + self.softmax_scale = ( + softmax_scale if softmax_scale is not None else head_size**-0.5 + ) + self.num_kv_heads = num_kv_heads if num_kv_heads is not None else num_heads + + self.schedule = SubBlockSparseSchedule.from_server_args() + self.layer_idx = _dit_layer_index(prefix) + # A layer outside the DiT stack (token refiner, cross attention) never + # runs sparse: its sequences are short and its budget meaningless. + self.layer_enabled = ( + self.layer_idx is not None + and self.layer_idx >= self.schedule.skip_first_layers + and head_size == SUBBLOCK_SPARSE_HEAD_DIM + and self.schedule.sparsity > 0.0 + ) + self.router = ( + SubBlockRouter(n_k=self.schedule.n_k, n_q=self.schedule.n_q) + if self.layer_enabled + else None + ) + self.dense_impl = self._build_dense_impl(causal=causal) + if self.layer_enabled: + logger.info_once( + f"SubBlock sparse attention: sparsity={self.schedule.sparsity:.3f} " + f"n_k={self.schedule.n_k} n_q={self.schedule.n_q}, dense for the first " + f"{self.schedule.skip_first_steps} denoise steps and the first " + f"{self.schedule.skip_first_layers} DiT layers" + ) + + def _build_dense_impl(self, *, causal: bool) -> AttentionImpl: + """Flash attention, used wherever the schedule excludes sparsity.""" + from sglang.multimodal_gen.runtime.layers.attention.selector import ( + get_attn_backend, + ) + + backend = get_attn_backend( + self.head_size, + torch.bfloat16, + supported_attention_backends={ + AttentionBackendEnum.FA, + AttentionBackendEnum.TORCH_SDPA, + }, + selected_attention_backend=AttentionBackendEnum.FA, + ) + return backend.get_impl_cls()( + num_heads=self.num_heads, + head_size=self.head_size, + causal=causal, + softmax_scale=self.softmax_scale, + num_kv_heads=self.num_kv_heads, + prefix=f"{self.prefix}.dense", + ) + + def _step_enabled(self) -> bool: + return get_forward_context().current_timestep >= self.schedule.skip_first_steps + + def _sparse_ready(self, q: torch.Tensor, k: torch.Tensor) -> bool: + return ( + self.layer_enabled + and self._step_enabled() + and q.dtype == torch.bfloat16 + and k.shape[-3] >= self.schedule.min_seq_len + and not self.causal + ) + + def _sparse_attention( + self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor + ) -> torch.Tensor: + """q, k, v: ``[1, S, H, 128]`` bf16 -> same shape.""" + bsa_attn_blk64_fwd = load_bsa_attn_blk64_fwd() + plan = self.router.route( + q, k, sparsity=self.schedule.sparsity, softmax_scale=self.softmax_scale + ) + # Proof that the sparse path actually ran, with the shape it ran on -- + # the construction-time log above only says the layer was eligible. + logger.info_once( + f"SubBlock sparse attention active: S={k.shape[1]} heads={q.shape[2]} " + f"keeping {plan.topk}/{plan.num_blocks} key blocks per query block " + f"(sparsity {1 - plan.density:.4f})" + ) + out = bsa_attn_blk64_fwd( + q, + k, + v, + plan.index, + plan.topk, + block_sizes=_cached_block_sizes(k.shape[1], k.device), + q2k_block_nums=None, # the budget is uniform across rows + softmax_scale=self.softmax_scale, + ) + return out[0] if isinstance(out, tuple) else out + + def forward( + self, + query: torch.Tensor, + key: torch.Tensor, + value: torch.Tensor, + attn_metadata: SubBlockSparseAttentionMetadata | None = None, + ) -> torch.Tensor: + """query/key/value: ``[B, S, H, D]``.""" + if not self._sparse_ready(query, key): + return self.dense_impl.forward(query, key, value, attn_metadata) + return self._sparse_attention(query, key, value) + + def forward_varlen( + self, + query: torch.Tensor, + key: torch.Tensor, + value: torch.Tensor, + *, + cu_seqlens: torch.Tensor, + max_seqlen: int, + cu_seqlens_host: tuple[int, ...] | None = None, + ) -> torch.Tensor: + """Packed ``[T, H, D]`` rows split into documents by ``cu_seqlens``. + + The block-sparse kernel takes one contiguous sequence, so each packed + document is routed on its own. Documents shorter than ``min_seq_len`` + -- in MiniMax H3 the padding tail -- go through the dense kernel. + """ + + def all_dense() -> torch.Tensor: + return self.dense_impl.forward_varlen( + query, + key, + value, + cu_seqlens=cu_seqlens, + max_seqlen=max_seqlen, + cu_seqlens_host=cu_seqlens_host, + ) + + if cu_seqlens_host is None or not self._sparse_ready(query, key): + return all_dense() + + segments = [ + (start, stop) + for start, stop in zip(cu_seqlens_host[:-1], cu_seqlens_host[1:]) + if stop > start + ] + sparse_segments = { + (start, stop) + for start, stop in segments + if stop - start >= self.schedule.min_seq_len + } + if not sparse_segments: + return all_dense() + + out = torch.empty_like(query) + # cu_seqlens covers every packed row in practice; a caller that leaves + # a tail outside the last document would otherwise read uninitialized + # memory back out. + if segments[-1][1] < query.shape[0]: + out[segments[-1][1] :].zero_() + for start, stop in segments: + # Deliberately not `.contiguous()`. After the Ulysses all-to-all, + # q/k/v are last-dim slices of one packed buffer, so they are + # strided; both the block-sparse kernel and SDPA permute them + # anyway, and forcing contiguity here measured as a wasted + # full-tensor copy (0.46 ms per call at S=37.7k on B200). + q_seg = query[start:stop].unsqueeze(0) + k_seg = key[start:stop].unsqueeze(0) + v_seg = value[start:stop].unsqueeze(0) + if (start, stop) in sparse_segments: + seg_out = self._sparse_attention(q_seg, k_seg, v_seg) + else: + seg_out = self._dense_segment(q_seg, k_seg, v_seg) + out[start:stop] = seg_out[0] + return out + + def _dense_segment( + self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor + ) -> torch.Tensor: + """Dense attention over one packed document, ``[1, S, H, D]``.""" + return torch.nn.functional.scaled_dot_product_attention( + q.transpose(1, 2), + k.transpose(1, 2), + v.transpose(1, 2), + is_causal=self.causal, + scale=self.softmax_scale, + ).transpose(1, 2) + + +__all__ = [ + "SubBlockSparseAttentionBackend", + "SubBlockSparseAttentionImpl", + "SubBlockSparseAttentionMetadata", + "SubBlockSparseAttentionMetadataBuilder", + "SubBlockSparseSchedule", +] diff --git a/python/sglang/multimodal_gen/runtime/platforms/cuda.py b/python/sglang/multimodal_gen/runtime/platforms/cuda.py index e2b4518fa..5b89fb765 100644 --- a/python/sglang/multimodal_gen/runtime/platforms/cuda.py +++ b/python/sglang/multimodal_gen/runtime/platforms/cuda.py @@ -297,6 +297,44 @@ class _VMOBAAttentionBackendResolver(_CudaAttentionBackendResolver): raise ImportError("Video MoBA Attention backend is not installed. ") from e +class _SubBlockSparseAttentionBackendResolver(_CudaAttentionBackendResolver): + backend = AttentionBackendEnum.SUBBLOCK_SPARSE_ATTN + + # The blk64 kernel is built `-gencode=arch=compute_100a,code=sm_100a`, which + # is arch-specific: 10.3 (B300 / GB300) and 12.x have no cubin. Its own guard + # only compares the major version, so it would accept 10.3 and fail later. + required_capability = (10, 0) + + @classmethod + def resolve(cls, platform) -> str: + capability = platform.get_device_capability() + if capability is None or capability != cls.required_capability: + found = capability.as_version_str() if capability else "unknown" + raise ValueError( + "SubBlock sparse attention needs compute capability " + f"{'.'.join(map(str, cls.required_capability))} (B200 / GB200); " + f"this device reports {found}." + ) + try: + from sglang.multimodal_gen.runtime.layers.attention.backends.subblock_sparse import ( # noqa: F401 + load_bsa_attn_blk64_fwd, + ) + from sglang.multimodal_gen.runtime.layers.attention.backends.subblock_sparse_attn import ( # noqa: F401 + SubBlockSparseAttentionBackend, + ) + + # Importing the entry point catches a missing or broken FlashInfer; + # the CUDA extension itself is built lazily on the first call. + load_bsa_attn_blk64_fwd() + return "sglang.multimodal_gen.runtime.layers.attention.backends.subblock_sparse_attn.SubBlockSparseAttentionBackend" + except Exception as e: + logger.error("Failed to import SubBlock sparse attention: %s", str(e)) + raise ImportError( + "SubBlock sparse attention needs FlashInfer with the blk64 " + "block-sparse kernel (flashinfer.cute_dsl.sparse.bsa_attn_blk64_fwd)." + ) from e + + class _FlashAttention2BackendResolver(_CudaAttentionBackendResolver): backend = AttentionBackendEnum.FA2 @@ -338,6 +376,7 @@ _CUDA_ATTENTION_BACKEND_RESOLVERS = { _SparseVideoGen2AttentionBackendResolver, _SolAttnBackendResolver, _VMOBAAttentionBackendResolver, + _SubBlockSparseAttentionBackendResolver, _FlashAttention2BackendResolver, _FlashAttentionBackendResolver, ) diff --git a/python/sglang/multimodal_gen/runtime/platforms/interface.py b/python/sglang/multimodal_gen/runtime/platforms/interface.py index 2bd5d7a1f..a0be9394e 100644 --- a/python/sglang/multimodal_gen/runtime/platforms/interface.py +++ b/python/sglang/multimodal_gen/runtime/platforms/interface.py @@ -44,6 +44,7 @@ class AttentionBackendEnum(enum.Enum): BLOCK_SPARSE_ATTN = enum.auto() RAIN_FUSION_ATTN = enum.auto() SOL_ATTN = enum.auto() + SUBBLOCK_SPARSE_ATTN = enum.auto() NO_ATTENTION = enum.auto() def __str__(self): @@ -62,6 +63,7 @@ class AttentionBackendEnum(enum.Enum): AttentionBackendEnum.BLOCK_SPARSE_ATTN, AttentionBackendEnum.RAIN_FUSION_ATTN, AttentionBackendEnum.SOL_ATTN, + AttentionBackendEnum.SUBBLOCK_SPARSE_ATTN, } diff --git a/python/sglang/multimodal_gen/test/unit/test_subblock_sparse_attention.py b/python/sglang/multimodal_gen/test/unit/test_subblock_sparse_attention.py new file mode 100644 index 000000000..ffa1feb9b --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_subblock_sparse_attention.py @@ -0,0 +1,357 @@ +# SPDX-License-Identifier: Apache-2.0 +"""SubBlock block-sparse attention backend. + +The schedule tests are pure CPU. The numerical tests need an SM100 GPU with +FlashInfer's ``bsa_attn_blk64_fwd`` and are skipped otherwise. + +The trick that makes the sparse kernel checkable against dense attention: at +``sparsity`` just above 0 every block is inside the budget, so the block-sparse +result must reproduce dense attention up to bf16 rounding. That covers the +routing indices, the tail block sizes, and the softmax scale in one assertion, +none of which an accuracy-only comparison at real sparsity would pin down. +""" + +from __future__ import annotations + +import unittest +from unittest.mock import patch + +import torch + +from sglang.multimodal_gen.runtime.layers.attention.backends.subblock_sparse.router import ( + _snap_up_to_8, +) +from sglang.multimodal_gen.runtime.layers.attention.backends.subblock_sparse_attn import ( + SubBlockSparseAttentionBackend, + SubBlockSparseAttentionImpl, + SubBlockSparseSchedule, + _dit_layer_index, +) + +HEAD_DIM = 128 +NUM_HEADS = 4 + + +def _sm100_available() -> bool: + if not torch.cuda.is_available(): + return False + # Exactly 10.0: the kernel is built for sm_100a, and 10.3 has no cubin. + if torch.cuda.get_device_capability(0) != (10, 0): + return False + try: + from sglang.multimodal_gen.runtime.layers.attention.backends.subblock_sparse import ( + load_bsa_attn_blk64_fwd, + ) + + load_bsa_attn_blk64_fwd() + except Exception: + return False + return True + + +requires_sm100 = unittest.skipUnless( + _sm100_available(), "needs SM100 and FlashInfer bsa_attn_blk64_fwd" +) + + +class _FakeServerArgs: + def __init__(self, config): + self.attention_backend_config = config + + +def _patch_schedule(config): + return patch( + "sglang.multimodal_gen.runtime.server_args.get_global_server_args", + return_value=_FakeServerArgs(config), + ) + + +def _patch_step(step: int): + class _Ctx: + current_timestep = step + + return patch( + "sglang.multimodal_gen.runtime.layers.attention.backends.subblock_sparse_attn.get_forward_context", + return_value=_Ctx(), + ) + + +def _dense_reference( + q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, scale: float +) -> torch.Tensor: + """q, k, v: [1, S, H, D].""" + return torch.nn.functional.scaled_dot_product_attention( + q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2), scale=scale + ).transpose(1, 2) + + +def _structured_qkv( + seq_len: int, device: torch.device, n_topic: int = 64, seed: int = 0 +): + """Attention-like q/k/v: ``[1, S, H, 128]`` bf16. + + Random q/k produces a near-uniform attention map, and under a uniform map + no block-selection rule can work -- dropping blocks drops mass wherever you + cut. Real video attention concentrates on a limited region of keys, so each + query here prefers one narrow key topic (a couple of 64-token blocks), and + each topic carries its own value, so attending to the wrong blocks gives a + visibly wrong answer instead of the same blurred average. + """ + gen = torch.Generator(device=device).manual_seed(seed) + shape = (n_topic, NUM_HEADS, HEAD_DIM) + key_centers = torch.randn(shape, device=device, generator=gen) + value_centers = torch.randn(shape, device=device, generator=gen) + topic = torch.arange(seq_len, device=device) * n_topic // seq_len + + def _noise(scale: float) -> torch.Tensor: + return scale * torch.randn( + seq_len, NUM_HEADS, HEAD_DIM, device=device, generator=gen + ) + + q = key_centers[topic] + _noise(0.3) + # the topic a query wants sits half a sequence away, so a rule that simply + # keeps the diagonal cannot pass + k = key_centers[(topic + n_topic // 2) % n_topic] + _noise(0.3) + v = value_centers[topic] + _noise(0.1) + return ( + q[None].to(torch.bfloat16), + k[None].to(torch.bfloat16), + v[None].to(torch.bfloat16), + ) + + +def _cosine(a: torch.Tensor, b: torch.Tensor) -> float: + return float( + torch.nn.functional.cosine_similarity( + a.float().flatten(), b.float().flatten(), dim=0 + ) + ) + + +class TestSubBlockSparseSchedule(unittest.TestCase): + def test_dit_layer_index_only_matches_top_level_blocks(self): + self.assertEqual(_dit_layer_index("blocks.7.attn"), 7) + self.assertEqual(_dit_layer_index("blocks.0.attn"), 0) + self.assertIsNone(_dit_layer_index("token_refiner.blocks.1.attn")) + self.assertIsNone(_dit_layer_index("")) + self.assertIsNone(_dit_layer_index("blocks_extra.3.attn")) + + def test_defaults_when_config_is_empty(self): + with _patch_schedule({}): + schedule = SubBlockSparseSchedule.from_server_args() + self.assertEqual(schedule.sparsity, 0.75) + self.assertEqual(schedule.skip_first_steps, 10) + # Depth is not protected by default; the early steps are. See the + # sweep recorded next to the constants. + self.assertEqual(schedule.skip_first_layers, 0) + self.assertEqual(schedule.n_k, 4) + self.assertEqual(schedule.n_q, 4) + + def test_rejects_out_of_range_values(self): + for config in ({"sparsity": 1.0}, {"n_k": 3}, {"skip_first_steps": -1}): + with self.subTest(config=config), _patch_schedule(config): + with self.assertRaises(ValueError): + SubBlockSparseSchedule.from_server_args() + + +class TestBudgetGranularity(unittest.TestCase): + """The kernel bills in groups of 8 blocks; the budget should collect them.""" + + def test_snaps_up_to_the_billed_count(self): + for topk, expected in ((148, 152), (118, 120), (1, 8), (0, 8)): + with self.subTest(topk=topk): + self.assertEqual(_snap_up_to_8(topk, 590), expected) + + def test_never_exceeds_the_blocks_that_exist(self): + """The cap wins over the granularity: 590 blocks means at most 590.""" + self.assertEqual(_snap_up_to_8(586, 590), 590) + self.assertEqual(_snap_up_to_8(3, 5), 5) + + +class TestSubBlockSparseBackend(unittest.TestCase): + def test_the_advertised_builder_can_be_built(self): + """`AttentionMetadataBuilder.__init__` is abstract; a builder that does + not override it makes `get_builder_cls()()` a TypeError.""" + builder = SubBlockSparseAttentionBackend.get_builder_cls()() + builder.prepare() + metadata = builder.build(current_timestep=7) + self.assertIsInstance( + metadata, SubBlockSparseAttentionBackend.get_metadata_cls() + ) + self.assertEqual(metadata.current_timestep, 7) + + +class TestSubBlockGating(unittest.TestCase): + """The schedule must decide sparsity from the layer and the step alone.""" + + def _impl(self, prefix: str, **config) -> SubBlockSparseAttentionImpl: + with _patch_schedule(config), patch.object( + SubBlockSparseAttentionImpl, "_build_dense_impl", return_value=None + ): + return SubBlockSparseAttentionImpl( + num_heads=NUM_HEADS, + head_size=HEAD_DIM, + causal=False, + softmax_scale=HEAD_DIM**-0.5, + prefix=prefix, + ) + + def test_skip_first_layers_gates_the_bottom_of_the_stack(self): + for prefix, expected in ( + ("blocks.0.attn", False), + ("blocks.1.attn", False), + ("blocks.2.attn", True), + ): + with self.subTest(prefix=prefix): + impl = self._impl(prefix, skip_first_layers=2) + self.assertEqual(impl.layer_enabled, expected) + + def test_token_refiner_is_dense(self): + self.assertFalse(self._impl("token_refiner.blocks.0.attn").layer_enabled) + + def test_head_dim_other_than_128_is_dense(self): + with _patch_schedule({}), patch.object( + SubBlockSparseAttentionImpl, "_build_dense_impl", return_value=None + ): + impl = SubBlockSparseAttentionImpl( + num_heads=NUM_HEADS, + head_size=64, + causal=False, + softmax_scale=64**-0.5, + prefix="blocks.9.attn", + ) + self.assertFalse(impl.layer_enabled) + + def test_first_steps_are_dense(self): + impl = self._impl("blocks.9.attn") + for step, expected in ((0, False), (9, False), (10, True), (49, True)): + with self.subTest(step=step), _patch_step(step): + self.assertEqual(impl._step_enabled(), expected) + + def test_short_sequences_are_dense(self): + impl = self._impl("blocks.9.attn") + q = torch.empty(1, 1024, NUM_HEADS, HEAD_DIM, dtype=torch.bfloat16) + with _patch_step(20): + self.assertFalse(impl._sparse_ready(q, q)) + + def test_fp32_is_dense(self): + impl = self._impl("blocks.9.attn") + q = torch.empty(1, 8192, NUM_HEADS, HEAD_DIM, dtype=torch.float32) + with _patch_step(20): + self.assertFalse(impl._sparse_ready(q, q)) + + +@requires_sm100 +class TestSubBlockNumerics(unittest.TestCase): + seq_len = 8192 + + def _impl(self, **config) -> SubBlockSparseAttentionImpl: + with _patch_schedule(config): + return SubBlockSparseAttentionImpl( + num_heads=NUM_HEADS, + head_size=HEAD_DIM, + causal=False, + softmax_scale=HEAD_DIM**-0.5, + prefix="blocks.9.attn", + ) + + def test_full_budget_reproduces_dense(self): + device = torch.device("cuda") + q, k, v = _structured_qkv(self.seq_len, device) + impl = self._impl(sparsity=1e-6) + with _patch_step(20): + out = impl.forward(q, k, v, None) + ref = _dense_reference(q, k, v, HEAD_DIM**-0.5) + self.assertGreater(_cosine(out, ref), 0.999) + + def test_ragged_tail_reproduces_dense(self): + """A sequence that is not a multiple of the 64-token block.""" + device = torch.device("cuda") + seq_len = self.seq_len + 37 + q, k, v = _structured_qkv(seq_len, device) + impl = self._impl(sparsity=1e-6) + with _patch_step(20): + out = impl.forward(q, k, v, None) + ref = _dense_reference(q, k, v, HEAD_DIM**-0.5) + self.assertGreater(_cosine(out, ref), 0.999) + + def test_routing_finds_the_blocks_that_carry_the_mass(self): + """At 0.75 sparsity the router must keep the blocks that matter. + + The random-budget control is the point of this test: with the same + number of blocks but chosen at random the output collapses, so a high + cosine here measures the routing, not a forgiving fixture. + """ + from sglang.multimodal_gen.runtime.layers.attention.backends.subblock_sparse import ( + SubBlockRouter, + load_bsa_attn_blk64_fwd, + ) + + device = torch.device("cuda") + q, k, v = _structured_qkv(self.seq_len, device) + ref = _dense_reference(q, k, v, HEAD_DIM**-0.5) + + impl = self._impl(sparsity=0.75) + with _patch_step(20): + routed = impl.forward(q, k, v, None) + self.assertTrue(torch.isfinite(routed.float()).all()) + self.assertGreater(_cosine(routed, ref), 0.99) + + num_blocks = (self.seq_len + 63) // 64 + topk = impl.router.route(q, k, sparsity=0.75, softmax_scale=HEAD_DIM**-0.5).topk + # A random permutation per row, not `randint`: sampling with replacement + # would leave the control holding duplicate blocks, so it would attend + # fewer distinct blocks than the router at the same budget, and the + # repeats would distort the softmax mass on top of that. + random_index = ( + torch.rand(1, NUM_HEADS, num_blocks, num_blocks, device=device) + .argsort(dim=-1)[..., :topk] + .to(torch.int32) + ) + random_out = load_bsa_attn_blk64_fwd()( + q, + k, + v, + random_index, + topk, + block_sizes=SubBlockRouter.block_sizes(self.seq_len, device), + q2k_block_nums=None, + softmax_scale=HEAD_DIM**-0.5, + ) + random_out = random_out[0] if isinstance(random_out, tuple) else random_out + self.assertLess(_cosine(random_out, ref), 0.9) + + def test_skipped_step_is_bitwise_dense(self): + device = torch.device("cuda") + q, k, v = _structured_qkv(self.seq_len, device) + impl = self._impl(sparsity=0.75, skip_first_steps=10) + with _patch_step(3): + skipped = impl.forward(q, k, v, None) + dense = impl.dense_impl.forward(q, k, v, None) + torch.testing.assert_close(skipped, dense, rtol=0, atol=0) + + def test_varlen_routes_each_document(self): + """Packed [real | padding] layout, exactly MiniMax H3's cu_seqlens.""" + device = torch.device("cuda") + used, total = self.seq_len, self.seq_len + 512 + q, k, v = _structured_qkv(total, device) + q, k, v = q[0], k[0], v[0] + cu_host = (0, used, total) + cu = torch.tensor(cu_host, dtype=torch.int32, device=device) + impl = self._impl(sparsity=1e-6) + with _patch_step(20): + out = impl.forward_varlen( + q, k, v, cu_seqlens=cu, max_seqlen=used, cu_seqlens_host=cu_host + ) + for start, stop in ((0, used), (used, total)): + ref = _dense_reference( + q[start:stop][None], + k[start:stop][None], + v[start:stop][None], + HEAD_DIM**-0.5, + )[0] + self.assertGreater(_cosine(out[start:stop], ref), 0.999) + + +if __name__ == "__main__": + unittest.main()