From 3340f4e3dae66807abbbe3acaa3ebc1174a1eb7c Mon Sep 17 00:00:00 2001 From: Yuan Luo Date: Thu, 18 Jun 2026 11:41:46 +0800 Subject: [PATCH] [GDN][KDA][mem_cache] int8 checkpoint pool for the linear-attn prefix cache (#28185) Co-authored-by: luoyuan.luo --- .../bench_int8_checkpoint_reuse.py | 165 ++++++++ .../attention/hybrid_linear_attn_backend.py | 17 +- .../scheduler_components/invariant_checker.py | 34 ++ .../pool_stats_observer.py | 13 +- .../srt/mem_cache/mamba_checkpoint_pool.py | 373 ++++++++++++++++++ .../sglang/srt/mem_cache/mamba_radix_cache.py | 93 ++++- python/sglang/srt/mem_cache/memory_pool.py | 20 + python/sglang/srt/server_args.py | 42 ++ .../test_int8_mamba_checkpoint_e2e.py | 93 +++++ .../mem_cache/test_int8_checkpoint_store.py | 178 +++++++++ 10 files changed, 1010 insertions(+), 18 deletions(-) create mode 100644 benchmark/bench_linear_attention/bench_int8_checkpoint_reuse.py create mode 100644 python/sglang/srt/mem_cache/mamba_checkpoint_pool.py create mode 100644 test/registered/radix_cache/test_int8_mamba_checkpoint_e2e.py create mode 100644 test/srt/mem_cache/test_int8_checkpoint_store.py diff --git a/benchmark/bench_linear_attention/bench_int8_checkpoint_reuse.py b/benchmark/bench_linear_attention/bench_int8_checkpoint_reuse.py new file mode 100644 index 000000000..12d5f26e6 --- /dev/null +++ b/benchmark/bench_linear_attention/bench_int8_checkpoint_reuse.py @@ -0,0 +1,165 @@ +"""Benchmark: int8 linear-attention checkpoint pool — prefix-reuse capacity & latency. + +Drives a *running* SGLang server that serves a linear-attention (KDA / GDN) hybrid +model, and measures how prefix reuse — and the probe-phase prefill latency that +depends on it — holds up as the number of DISTINCT cached prefixes grows. + +The active bf16 mamba state pool caches one state per distinct prefix and is sized +to the running set; once the number of distinct cached prefixes exceeds it, reuse +collapses (states get evicted and recomputed). With ``--enable-int8-mamba-checkpoint`` +the radix-cached states live in a separate int8 pool holding ~2x more slots at ~the +same memory, so the collapse knee moves out ~2x and the probe-phase prefill stays a +cheap cache hit well past the bf16 pool size. + +Method, per K in ``--num-prefixes``: + flush -> WARM: send K distinct ~P-token prefixes once (populate the cache) + -> PROBE: re-send each prefix with a short different suffix; read + meta_info.cached_tokens (= reused prefix length) and time each request. + reuse_frac = sum(cached) / sum(prefix_tokens); probe throughput = K / wall_time. + +Run the server twice and compare (same flags, toggle int8): + + python -m sglang.launch_server --model-path --tp 4 \ + --trust-remote-code --mamba-scheduler-strategy extra_buffer \ + --max-mamba-cache-size 256 [--enable-int8-mamba-checkpoint] --port 30000 + + python benchmark/bench_linear_attention/bench_int8_checkpoint_reuse.py \ + --port 30000 --prefix-tokens 1000 --num-prefixes 128 384 640 --parallel 8 + +NOTE: prefix-tokens must cross the mamba cache chunk granularity (typically ~512), +otherwise nothing is cacheable and reuse is 0 by construction (not a regression). +Use ``--mamba-scheduler-strategy extra_buffer`` on the server: ``no_buffer`` only +snapshots state at the full-sequence leaf, so a divergent-suffix probe never reuses. +""" + +import argparse +import random +import statistics +import time +from concurrent.futures import ThreadPoolExecutor + +import requests + +# A small word pool so each prefix is distinct but realistic English text. +_VOCAB = ( + "time year people way day man thing woman life child world school state family " + "student group country problem hand part place case week company system program " + "question work government number night point home water room mother area money " + "story fact month lot right study book eye job word business issue side kind head " + "house service friend father power hour game line end member law car city community " + "name president team minute idea body information back parent face level office door " + "health person art war history party result change morning reason research girl guy " + "moment air teacher force education foot boy age policy process music market sense " + "nation plan college interest death course someone experience behavior career goal" +).split() + + +def make_prefix(i: int, n_words: int) -> str: + rng = random.Random(1000 + i) + head = f"Document {i} unique tag {i * 7919 % 100000}. " + return head + " ".join(rng.choice(_VOCAB) for _ in range(n_words)) + + +def make_suffix(i: int, salt: int, n_words: int) -> str: + rng = random.Random(salt * i + salt) + return " ".join(rng.choice(_VOCAB) for _ in range(n_words)) + + +def main(): + ap = argparse.ArgumentParser() + ap.add_argument("--host", default="127.0.0.1") + ap.add_argument("--port", type=int, default=30000) + ap.add_argument( + "--num-prefixes", + type=int, + nargs="+", + default=[128, 384, 640], + help="distinct-prefix counts (K) to sweep", + ) + ap.add_argument( + "--prefix-tokens", + type=int, + default=1000, + help="approx prefix length in words (must cross the ~512 chunk granularity)", + ) + ap.add_argument("--suffix-tokens", type=int, default=8) + ap.add_argument("--parallel", type=int, default=8) + ap.add_argument("--timeout", type=int, default=600) + args = ap.parse_args() + + base = f"http://{args.host}:{args.port}" + gen_url = base + "/generate" + + def send(prompt): + t0 = time.perf_counter() + try: + r = requests.post( + gen_url, + headers={"Content-Type": "application/json"}, + json={ + "text": prompt, + "sampling_params": {"max_new_tokens": 1, "temperature": 0.0}, + }, + timeout=args.timeout, + ) + dt = time.perf_counter() - t0 + mi = r.json().get("meta_info", {}) + return mi.get("prompt_tokens"), mi.get("cached_tokens"), dt, None + except Exception as e: # noqa: BLE001 + return None, None, time.perf_counter() - t0, str(e)[:80] + + def flush(): + try: + requests.post(base + "/flush_cache", timeout=60) + time.sleep(1.5) + except Exception: + pass + + print( + f"server={base} prefix_tokens~{args.prefix_tokens} suffix_tokens={args.suffix_tokens} " + f"parallel={args.parallel} K_sweep={args.num_prefixes}" + ) + header = ( + f"{'K':>6} {'reuse_frac':>11} {'probe_p50_ms':>13} {'probe_p90_ms':>13} " + f"{'probe_thru_rps':>15} {'errors':>7}" + ) + print(header) + print("-" * len(header)) + + for K in args.num_prefixes: + flush() + prefixes = [make_prefix(i, args.prefix_tokens) for i in range(K)] + warm = [ + p + " " + make_suffix(i, 7, args.suffix_tokens) + for i, p in enumerate(prefixes) + ] + probe = [ + p + " " + make_suffix(i, 13, args.suffix_tokens) + for i, p in enumerate(prefixes) + ] + + with ThreadPoolExecutor(args.parallel) as ex: + list(ex.map(send, warm)) # WARM: populate the cache + t0 = time.perf_counter() + results = list(ex.map(send, probe)) # PROBE: measured + wall = time.perf_counter() - t0 + + ok = [(pt, ct, dt) for (pt, ct, dt, err) in results if err is None and pt] + errs = [r for r in results if r[3] is not None] + if not ok: + print(f"{K:>6} ALL-ERR e.g. {errs[:1]}") + continue + sum_prompt = sum(pt for pt, _, _ in ok) + sum_cached = sum((ct or 0) for _, ct, _ in ok) + lat_ms = sorted(dt * 1000.0 for _, _, dt in ok) + p50 = statistics.median(lat_ms) + p90 = lat_ms[min(len(lat_ms) - 1, int(0.9 * len(lat_ms)))] + thru = len(ok) / wall if wall > 0 else 0.0 + print( + f"{K:>6} {sum_cached / max(1, sum_prompt):>11.3f} {p50:>13.1f} " + f"{p90:>13.1f} {thru:>15.1f} {len(errs):>7}" + ) + + +if __name__ == "__main__": + main() diff --git a/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py b/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py index 75f4c890a..19ecb95cc 100644 --- a/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py +++ b/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py @@ -64,9 +64,20 @@ class MambaAttnBackendBase(AttentionBackend): forward_batch.mamba_cow_src_indices is not None and len(forward_batch.mamba_cow_src_indices) > 0 ): - self.req_to_token_pool.mamba_pool.copy_from( - forward_batch.mamba_cow_src_indices, forward_batch.mamba_cow_dst_indices - ) + ckpt_pool = getattr(self.req_to_token_pool, "mamba_ckpt_pool", None) + if ckpt_pool is not None: + # int8 checkpoints: dequantize the cached state (src = int8 ckpt slot) + # into the request's active bf16 slot (dst). + ckpt_pool.load_to_active( + self.req_to_token_pool.mamba_pool, + forward_batch.mamba_cow_src_indices, + forward_batch.mamba_cow_dst_indices, + ) + else: + self.req_to_token_pool.mamba_pool.copy_from( + forward_batch.mamba_cow_src_indices, + forward_batch.mamba_cow_dst_indices, + ) forward_batch.mamba_clear_indices = None forward_batch.mamba_cow_src_indices = None forward_batch.mamba_cow_dst_indices = None diff --git a/python/sglang/srt/managers/scheduler_components/invariant_checker.py b/python/sglang/srt/managers/scheduler_components/invariant_checker.py index 33013a424..4238ae281 100644 --- a/python/sglang/srt/managers/scheduler_components/invariant_checker.py +++ b/python/sglang/srt/managers/scheduler_components/invariant_checker.py @@ -114,6 +114,9 @@ class SchedulerInvariantChecker: ) def _check_mamba_pool(self, ps: PoolStats) -> Tuple[bool, str]: + ckpt_pool = getattr(self.req_to_token_pool, "mamba_ckpt_pool", None) + if ckpt_pool is not None: + return self._check_mamba_pool_with_int8(ps, ckpt_pool) leak, msg = self._check_pool_invariant( "mamba", ps.mamba_available_size, @@ -150,6 +153,37 @@ class SchedulerInvariantChecker: ) return leak, msg + def _check_mamba_pool_with_int8(self, ps: PoolStats, ckpt_pool) -> Tuple[bool, str]: + """Two-pool invariant for int8 mamba checkpoints. + + The radix-cached states live in the int8 checkpoint pool, NOT the active + bf16 pool. So the single-pool equation (active.available + radix_cached == + active.size) is wrong -- it double-counts the radix states against a pool + that does not hold them. Instead check the two pools independently: + + * active bf16 pool: backs running requests only; the radix owns ZERO + active slots. Checked at idle (in-flight == 0) -> available == total. + * int8 checkpoint pool: backs the radix-cached states; its occupancy is + exactly the radix evictable + protected counts. + """ + active_leak, active_msg = self._check_pool_invariant( + "mamba-active", + ps.mamba_available_size, + ps.mamba_evictable_size, # 0 in int8 mode (radix owns no active slots) + 0, + self.pool_stats_observer.session_held_mamba_slots(), + self.req_to_token_pool.mamba_pool.size, + ) + int8_leak, int8_msg = self._check_pool_invariant( + "mamba-int8", + ckpt_pool.available_size(), + self.tree_cache.mamba_evictable_size(), + self.tree_cache.mamba_protected_size(), + 0, + ckpt_pool.num_slots, + ) + return active_leak or int8_leak, active_msg + "\n" + int8_msg + def _get_total_uncached_sizes( self, ) -> Tuple[int, int]: diff --git a/python/sglang/srt/managers/scheduler_components/pool_stats_observer.py b/python/sglang/srt/managers/scheduler_components/pool_stats_observer.py index 5e811ad5a..7fcbda113 100644 --- a/python/sglang/srt/managers/scheduler_components/pool_stats_observer.py +++ b/python/sglang/srt/managers/scheduler_components/pool_stats_observer.py @@ -247,8 +247,19 @@ class SchedulerPoolStatsObserver: self.tree_cache.full_evictable_size() if is_mamba_radix_cache else 0 ) mamba_available_size = self.req_to_token_pool.mamba_allocator.available_size() + # `mamba_usage`/`mamba_num_used` track the ACTIVE bf16 pool occupancy (running + # requests) -- this feeds throttle decisions (get_max_pool_usage) which asserts + # usage >= 0. With int8 checkpoints the radix-cached states live in a SEPARATE + # int8 pool, so they own ZERO active slots: report evictable=0 against the active + # pool (otherwise active.size - (available + radix_cached) goes negative). The + # int8 cache pool's own occupancy is validated separately in the invariant check. + has_int8_ckpt = ( + getattr(self.req_to_token_pool, "mamba_ckpt_pool", None) is not None + ) mamba_evictable_size = ( - self.tree_cache.mamba_evictable_size() if is_mamba_radix_cache else 0 + self.tree_cache.mamba_evictable_size() + if (is_mamba_radix_cache and not has_int8_ckpt) + else 0 ) full_num_used = self.token_to_kv_pool_allocator.size - ( full_available_size + full_evictable_size diff --git a/python/sglang/srt/mem_cache/mamba_checkpoint_pool.py b/python/sglang/srt/mem_cache/mamba_checkpoint_pool.py new file mode 100644 index 000000000..13c30bdc9 --- /dev/null +++ b/python/sglang/srt/mem_cache/mamba_checkpoint_pool.py @@ -0,0 +1,373 @@ +""" +Copyright 2023-2026 SGLang Team +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. + +MambaCheckpointPool — the radix prefix cache's int8-compressed store for cached +linear-attention (KDA / GDN / Mamba2 gated-delta-rule) recurrent states. + +It decouples the *cached* states (radix-owned, idle, compressed) from the *active* +``MambaPool`` (running requests, full precision, kernel-facing). The radix stores +one cached state per node HERE; on a prefix-cache hit it is dequantized back into +a fresh active slot (copy-on-write). + +Per cached slot it holds: + * the SSM temporal state in **int8** (per-(head,k-channel) symmetric), via the + embedded ``Int8CheckpointStore`` — ~2x more cached states than bf16, + quality-safe (quantized once on store, dequantized once on a hit; never + re-enters the recurrence as a quant->dequant loop). + * the conv1d window state at its native dtype (tiny, W-1 tokens; not worth + quantizing). + +Why int8 (not fp8): a cached checkpoint is loaded ONCE on a cache hit, then +decoding continues at full precision, so the only error is a single rounding of +S. The temporal state is roughly uniformly distributed, so int8-per-(head, +k-channel) beats fp8-e4m3 at the same 1 byte (fp8 wastes bits on the exponent). +The scale axis (reduces over d_v) matches the per-k-channel decay diag(alpha), so +the large state entries keep ~bf16 precision and the error concentrates on small +entries that barely affect the readout. Storing cached states int8 gives ~2x the +cached-prefix capacity at fixed memory, and composes with host-offload +(HiMambaRadixCache) which it also halves. + +This is strategy-agnostic: whether the active slot to be cached was produced by +the ``no_buffer`` donate (copy_from) or the ``extra_buffer`` ping-pong track +buffer (spec path), both converge on "an active slot becomes the cached +``mamba_value``" — which is exactly the (store_from_active) hook here. Slot +lifecycle is owned by the caller via the embedded ``MambaSlotAllocator``. +""" + +from __future__ import annotations + +import logging +from typing import List, Optional + +import torch + +from sglang.srt.mem_cache.allocator.mamba import MambaSlotAllocator + +logger = logging.getLogger(__name__) + + +class Int8CheckpointStore: + """int8 store for cached multi-layer linear-attn states. + + Tensors (slot index handed out by the caller's allocator): + qdata : [L, num_slots, H, d_v, d_k] int8 (the quantized state) + scale : [L, num_slots, H, 1, d_k] scale_dtype (per layer,slot,head,k-chan) + + A "state" spans all L mamba layers for one cached point (matching how the + radix caches one full state per node). The reduction axis for the scale is + d_v (dim=-2), so each (head, k-channel) gets its own scale — aligned with the + per-k-channel decay diag(alpha). + + ``scale_dtype`` should match the source state's dtype (bf16 / fp16 / fp32) so + that quantize and dequantize use the identical scale — it is NOT required to + be bf16. + """ + + QMAX = 127 + + def __init__( + self, + *, + num_layers: int, + num_slots: int, + num_heads: int, + head_v_dim: int, + head_k_dim: int, + device: str, + scale_dtype: torch.dtype = torch.bfloat16, + ): + self.num_layers = num_layers + self.num_slots = num_slots + self.H = num_heads + self.d_v = head_v_dim + self.d_k = head_k_dim + self.device = device + self.qdata = torch.empty( + num_layers, + num_slots, + num_heads, + head_v_dim, + head_k_dim, + dtype=torch.int8, + device=device, + ) + self.scale = torch.empty( + num_layers, + num_slots, + num_heads, + 1, + head_k_dim, + dtype=scale_dtype, + device=device, + ) + + # ---- (de)quant math (also usable standalone for probes/tests) ---- + + @classmethod + def quantize(cls, state: torch.Tensor): + """state [..., H, d_v, d_k] -> (qint8, scale[..., H, 1, d_k]). + + amax / scale / round are computed in float32 so quantizing a low-precision + state doesn't lose precision in the intermediate (symmetric with + ``dequantize``, which is already float32). The scale is rounded to the + state dtype (its storage precision) BEFORE the division, so quantize and + dequantize use the identical scale.""" + state_fp32 = state.to(torch.float32) + amax = state_fp32.abs().amax(dim=-2, keepdim=True).clamp(min=1e-8) + scale = (amax / cls.QMAX).to(state.dtype) + q = ( + torch.round(state_fp32 / scale.to(torch.float32)) + .clamp(-cls.QMAX, cls.QMAX) + .to(torch.int8) + ) + return q, scale + + @staticmethod + def dequantize(q: torch.Tensor, scale: torch.Tensor, out_dtype: torch.dtype): + return (q.to(torch.float32) * scale.to(torch.float32)).to(out_dtype) + + # ---- store / load (caller supplies slot indices) ---- + + def store(self, slots: torch.Tensor, state: torch.Tensor) -> None: + """Quantize and write states. state: [L, N, H, d_v, d_k] for the N slots + (or [L, H, d_v, d_k] when slots is a scalar/len-1).""" + if state.dim() == 4: + state = state.unsqueeze(1) + q, scale = self.quantize(state) + self.qdata[:, slots] = q + self.scale[:, slots] = scale.to(self.scale.dtype) + + def load(self, slots: torch.Tensor, out_dtype: torch.dtype) -> torch.Tensor: + """Dequantize states at slots -> [L, N, H, d_v, d_k] in out_dtype.""" + return self.dequantize(self.qdata[:, slots], self.scale[:, slots], out_dtype) + + def copy_to_pool( + self, + dst_temporal: torch.Tensor, + src_slots: torch.Tensor, + dst_slots: torch.Tensor, + ) -> None: + """Dequantize checkpoints at ``src_slots`` directly into the active pool + tensor ``dst_temporal`` [L, pool_slots, H, d_v, d_k] at ``dst_slots`` (the + copy-on-write on a cache hit). Output dtype follows ``dst_temporal``.""" + dst_temporal[:, dst_slots] = self.load(src_slots, dst_temporal.dtype) + + def store_from_pool( + self, + src_temporal: torch.Tensor, + src_slots: torch.Tensor, + dst_slots: torch.Tensor, + ) -> None: + """Quantize states from an active pool tensor into checkpoint slots (cache + store / donate).""" + self.store(dst_slots, src_temporal[:, src_slots]) + + def mem_usage_bytes(self) -> int: + return ( + self.qdata.numel() * self.qdata.element_size() + + self.scale.numel() * self.scale.element_size() + ) + + def bytes_per_slot(self) -> int: + return self.mem_usage_bytes() // max(1, self.num_slots) + + +class MambaCheckpointPool: + def __init__( + self, + *, + num_layers: int, + num_slots: int, + num_heads: int, + head_v_dim: int, + head_k_dim: int, + conv_shapes: List[tuple], + conv_dtype: torch.dtype, + device: str, + temporal_dtype: Optional[torch.dtype] = None, + ): + self.num_slots = num_slots + self.device = device + self.temporal = Int8CheckpointStore( + num_layers=num_layers, + num_slots=num_slots + 1, # slot 0 reserved (matches MambaSlotAllocator) + num_heads=num_heads, + head_v_dim=head_v_dim, + head_k_dim=head_k_dim, + device=device, + # store the scale in the temporal state's own dtype so quantize and + # dequantize use the identical scale (not hard-coded to bf16) + scale_dtype=( + temporal_dtype if temporal_dtype is not None else torch.bfloat16 + ), + ) + # conv windows stay at their native dtype (small); one buffer per conv + # tensor in the State + self.conv = [ + torch.empty( + (num_layers, num_slots + 1) + tuple(shape), + dtype=conv_dtype, + device=device, + ) + for shape in conv_shapes + ] + self.allocator = MambaSlotAllocator(size=num_slots, device=device) + + # ---- lifecycle (delegates to the embedded allocator) ---- + + def alloc(self, n: int = 1): + return self.allocator.alloc(n) + + def free(self, slots: torch.Tensor): + self.allocator.free(slots) + + def available_size(self) -> int: + return self.allocator.available_size() + + def clear(self) -> None: + """Release every checkpoint slot (radix flush/reset). The int8 qdata is + left as-is; slots are reused/overwritten on the next store.""" + self.allocator.clear() + + # ---- state transfer between the active MambaPool and this store ---- + + def store_from_active(self, active_mamba_pool, active_slots, ckpt_slots) -> None: + """Quantize temporal + copy conv from the active pool into checkpoint slots + (the radix donate / cache-store).""" + cache = active_mamba_pool.mamba_cache + self.temporal.store_from_pool(cache.temporal, active_slots, ckpt_slots) + for i, c in enumerate(self.conv): + c[:, ckpt_slots] = cache.conv[i][:, active_slots] + + def load_to_active(self, active_mamba_pool, ckpt_slots, active_slots) -> None: + """Dequantize temporal + copy conv from checkpoint slots into the active pool + (the cache-hit copy-on-write).""" + cache = active_mamba_pool.mamba_cache + self.temporal.copy_to_pool(cache.temporal, ckpt_slots, active_slots) + for i, c in enumerate(self.conv): + cache.conv[i][:, active_slots] = c[:, ckpt_slots].to(cache.conv[i].dtype) + + @staticmethod + def estimate_mem_usage_bytes( + *, + num_layers: int, + num_slots: int, + num_heads: int, + head_v_dim: int, + head_k_dim: int, + conv_shapes: List[tuple], + conv_dtype: torch.dtype, + temporal_dtype: torch.dtype, + ) -> dict: + """Estimate the pool's HBM footprint (bytes) WITHOUT allocating, so a + caller can check it against free memory before construction. Mirrors the + real layout: int8 qdata + per-(head,k) scale + bf16 conv windows, including + the reserved slot 0.""" + slots = num_slots + 1 # slot 0 reserved (matches MambaSlotAllocator) + scale_isz = torch.empty((), dtype=temporal_dtype).element_size() + conv_isz = torch.empty((), dtype=conv_dtype).element_size() + qdata = num_layers * slots * num_heads * head_v_dim * head_k_dim # int8 = 1B + scale = num_layers * slots * num_heads * head_k_dim * scale_isz + conv = 0 + for shape in conv_shapes: + n = 1 + for s in shape: + n *= int(s) + conv += num_layers * slots * n * conv_isz + return { + "qdata": qdata, + "scale": scale, + "conv": conv, + "total": qdata + scale + conv, + } + + def mem_usage_bytes(self) -> int: + conv_bytes = sum(c.numel() * c.element_size() for c in self.conv) + return self.temporal.mem_usage_bytes() + conv_bytes + + +def maybe_init_int8_mamba_checkpoint_pool( + *, + mamba_size: int, + cache_params, + mamba_layer_ids: List[int], + device: str, +) -> Optional[MambaCheckpointPool]: + """Build the optional int8 ``MambaCheckpointPool`` when + ``--enable-int8-mamba-checkpoint`` is set (and a global server-args context + exists), else return ``None``. The radix caches states here (int8) instead of + in the active bf16 pool -> ~2x cached-prefix capacity at fixed memory. + + Estimates the pool's HBM footprint and checks it against free memory BEFORE + allocating, so an oversized ``--int8-mamba-ckpt-size`` fails with an actionable + message instead of a cryptic mid-allocation CUDA OOM. + """ + from sglang.srt.server_args import get_global_server_args + + try: + _sa = get_global_server_args() + except ValueError: + # Some unit-test / mock runners construct HybridReqToTokenPool directly + # without a global server-args context. The int8 checkpoint pool is opt-in + # via a CLI flag, so an unset context unambiguously means it is off. + _sa = None + if not getattr(_sa, "enable_int8_mamba_checkpoint", False): + return None + + GB = 1 << 30 + H, d_v, d_k = cache_params.shape.temporal + ckpt_size = _sa.int8_mamba_ckpt_size or (2 * mamba_size) + kwargs = dict( + num_layers=len(mamba_layer_ids), + num_slots=ckpt_size, + num_heads=H, + head_v_dim=d_v, + head_k_dim=d_k, + conv_shapes=list(cache_params.shape.conv), + conv_dtype=cache_params.dtype.conv, + temporal_dtype=cache_params.dtype.temporal, + ) + + est = MambaCheckpointPool.estimate_mem_usage_bytes(**kwargs) + free_bytes = None + if isinstance(device, str) and device.startswith("cuda"): + try: + free_bytes, _ = torch.cuda.mem_get_info(device) + except Exception: + free_bytes = None + logger.info( + f"int8 mamba checkpoint pool: {ckpt_size} slots, " + f"{est['total'] / GB:.2f}GB (qdata {est['qdata'] / GB:.2f} + scale " + f"{est['scale'] / GB:.2f} + conv {est['conv'] / GB:.2f}); active mamba " + f"pool {mamba_size} slots" + + (f"; free HBM {free_bytes / GB:.2f}GB" if free_bytes is not None else "") + ) + if free_bytes is not None and est["total"] >= free_bytes: + raise RuntimeError( + f"int8 mamba checkpoint pool needs ~{est['total'] / GB:.2f}GB but only " + f"{free_bytes / GB:.2f}GB HBM is free. Lower --int8-mamba-ckpt-size " + f"(currently {ckpt_size}) or --mem-fraction-static." + ) + + pool = MambaCheckpointPool(device=device, **kwargs) + # NOTE: this pool's HBM is NOT subtracted from the KV-cache budget + # (max_total_num_tokens); it is allocated from --mem-fraction-static headroom. + # The estimate check above guards against an oversized pool; accounting it in + # the KV budget is a follow-up. + logger.warning( + f"int8 mamba checkpoint pool ({est['total'] / GB:.2f}GB) is allocated from " + f"--mem-fraction-static headroom and is not reflected in " + f"max_total_num_tokens; ensure headroom covers it." + ) + return pool diff --git a/python/sglang/srt/mem_cache/mamba_radix_cache.py b/python/sglang/srt/mem_cache/mamba_radix_cache.py index 4a8e53834..e01c41446 100644 --- a/python/sglang/srt/mem_cache/mamba_radix_cache.py +++ b/python/sglang/srt/mem_cache/mamba_radix_cache.py @@ -562,22 +562,29 @@ class MambaRadixCache(KVCacheEventMixin, BasePrefixCache): mamba_ping_pong_track_buffer_to_keep = ( self.req_to_token_pool.get_mamba_ping_pong_keep_idx(req) ) - mamba_value = ( - req.mamba_ping_pong_track_buffer[ - mamba_ping_pong_track_buffer_to_keep - ] - .unsqueeze(-1) - .clone() - ) - assert mamba_value.item() != -1, ( + src_active = req.mamba_ping_pong_track_buffer[ + mamba_ping_pong_track_buffer_to_keep + ].unsqueeze(-1) + assert src_active.item() != -1, ( f"Cached mamba slot is -1: keep_idx={mamba_ping_pong_track_buffer_to_keep}, " f"buf={req.mamba_ping_pong_track_buffer.tolist()}, " f"next_track_idx={req.mamba_next_track_idx}, " f"last_track_seqlen={req.mamba_last_track_seqlen}, " f"rid={req.rid}" ) + if self.int8_ckpt_pool is not None: + mamba_value = self._commit_int8_checkpoint(src_active) + # quantized -> no ping-pong slot needs keeping + mamba_ping_pong_track_buffer_to_keep = None + else: + mamba_value = src_active.clone() else: - mamba_value = req.mamba_pool_idx.unsqueeze(-1).clone() + if self.int8_ckpt_pool is not None: + mamba_value = self._commit_int8_checkpoint( + req.mamba_pool_idx.unsqueeze(-1) + ) + else: + mamba_value = req.mamba_pool_idx.unsqueeze(-1).clone() mamba_ping_pong_track_buffer_to_keep = None result = self.insert( @@ -589,6 +596,9 @@ class MambaRadixCache(KVCacheEventMixin, BasePrefixCache): ) ) mamba_exist = result.mamba_exist + if mamba_exist and self.int8_ckpt_pool is not None: + # state already cached -> the int8 slot we just allocated is a duplicate + self.int8_ckpt_pool.free(mamba_value) else: self.token_to_kv_pool_allocator.free(kv_indices[req.cache_protected_len :]) mamba_exist = True @@ -596,7 +606,13 @@ class MambaRadixCache(KVCacheEventMixin, BasePrefixCache): if mamba_exist: mamba_ping_pong_track_buffer_to_keep = None - free_mamba_cache = True if self.enable_mamba_extra_buffer else mamba_exist + # With int8 checkpoints the radix owns an int8 slot (not the request's active + # slot), so the active mamba slot must always be returned to the active pool. + free_mamba_cache = ( + True + if (self.enable_mamba_extra_buffer or self.int8_ckpt_pool is not None) + else mamba_exist + ) if free_mamba_cache: self.req_to_token_pool.free_mamba_cache( @@ -649,7 +665,21 @@ class MambaRadixCache(KVCacheEventMixin, BasePrefixCache): # Donate the mamba index to the radix cache instead of copying. # This avoids a data copy that would race with the forward stream. - if self.enable_mamba_extra_buffer: + if self.int8_ckpt_pool is not None: + # int8 path: quantize the to-be-cached active state into an int8 slot + # (strategy-agnostic donate hook). + if self.enable_mamba_extra_buffer: + new_slot = self._alloc_mamba_slot() + src_active = self.req_to_token_pool.donate_mamba_ping_pong_slot( + req, new_slot + ) + mamba_value_donated = self._commit_int8_checkpoint(src_active) + self.req_to_token_pool.mamba_allocator.free(src_active) + else: + mamba_value_donated = self._commit_int8_checkpoint( + req.mamba_pool_idx.view(-1) + ) + elif self.enable_mamba_extra_buffer: new_slot = self._alloc_mamba_slot() mamba_value_donated = self.req_to_token_pool.donate_mamba_ping_pong_slot( req, new_slot @@ -671,7 +701,7 @@ class MambaRadixCache(KVCacheEventMixin, BasePrefixCache): ) new_prefix_len, mamba_exist = result.prefix_len, result.mamba_exist if mamba_exist: - self.req_to_token_pool.mamba_allocator.free(mamba_value_donated) + self._free_mamba_value(mamba_value_donated) # The prefix indices could be updated, reuse it match_result = self.match_prefix( @@ -729,7 +759,7 @@ class MambaRadixCache(KVCacheEventMixin, BasePrefixCache): self._record_remove_event(x) self.token_to_kv_pool_allocator.free(x.value) full_num_evicted = len(x.value) - self.req_to_token_pool.mamba_allocator.free(x.mamba_value) + self._free_mamba_value(x.mamba_value) mamba_num_evicted = len(x.mamba_value) # 2. get the next node, update the lru lists @@ -782,7 +812,7 @@ class MambaRadixCache(KVCacheEventMixin, BasePrefixCache): if len(x.children) > 0: # 1. an internal node, free mamba tokens. - self.req_to_token_pool.mamba_allocator.free(x.mamba_value) + self._free_mamba_value(x.mamba_value) mamba_num_evicted += len(x.mamba_value) # 2. get the next node, update the lru lists @@ -954,6 +984,41 @@ class MambaRadixCache(KVCacheEventMixin, BasePrefixCache): assert slot is not None, "Can not alloc mamba cache" return slot + @property + def int8_ckpt_pool(self): + """The int8 checkpoint pool, or None when --enable-int8-mamba-checkpoint is off. + When enabled, radix-cached mamba states live HERE (int8), not in the active + bf16 pool -> ~2x cached-prefix capacity at fixed memory.""" + return getattr(self.req_to_token_pool, "mamba_ckpt_pool", None) + + def _alloc_int8_ckpt_slot(self) -> torch.Tensor: + """Allocate one int8 checkpoint slot, evicting cached states if the pool is full.""" + slot = self.int8_ckpt_pool.alloc(1) + if slot is None: + self.evict(EvictParams(num_tokens=0, mamba_num=1)) + slot = self.int8_ckpt_pool.alloc(1) + assert slot is not None, "Can not alloc int8 mamba checkpoint slot" + return slot + + def _commit_int8_checkpoint(self, active_slots: torch.Tensor) -> torch.Tensor: + """Quantize the active-pool state at ``active_slots`` into a fresh int8 + checkpoint slot and return that slot. Strategy-agnostic donate hook: both + no_buffer (copy_from) and extra_buffer (ping-pong) converge here. The caller + frees ``active_slots`` separately.""" + ckpt_slot = self._alloc_int8_ckpt_slot() + self.int8_ckpt_pool.store_from_active( + self.req_to_token_pool.mamba_pool, active_slots, ckpt_slot + ) + return ckpt_slot + + def _free_mamba_value(self, mamba_value: torch.Tensor) -> None: + """Free a node's mamba_value to the right allocator (int8 ckpt pool or the + active mamba allocator).""" + if self.int8_ckpt_pool is not None: + self.int8_ckpt_pool.free(mamba_value) + else: + self.req_to_token_pool.mamba_allocator.free(mamba_value) + def _match_prefix_helper( self, key: RadixKey ) -> Tuple[List[torch.Tensor], TreeNode, int]: diff --git a/python/sglang/srt/mem_cache/memory_pool.py b/python/sglang/srt/mem_cache/memory_pool.py index 31db39e25..ec8479e1f 100644 --- a/python/sglang/srt/mem_cache/memory_pool.py +++ b/python/sglang/srt/mem_cache/memory_pool.py @@ -612,6 +612,20 @@ class HybridReqToTokenPool(ReqToTokenPool): ) self.mamba_map = {layer_id: i for i, layer_id in enumerate(mamba_layer_ids)} + # Optional int8 checkpoint pool: the radix caches states here (int8) instead + # of holding them in the active bf16 pool -> ~2x cached-prefix capacity at + # fixed memory. Strategy-agnostic (no_buffer / extra_buffer / spec). + from sglang.srt.mem_cache.mamba_checkpoint_pool import ( + maybe_init_int8_mamba_checkpoint_pool, + ) + + self.mamba_ckpt_pool = maybe_init_int8_mamba_checkpoint_pool( + mamba_size=mamba_size, + cache_params=cache_params, + mamba_layer_ids=mamba_layer_ids, + device=device, + ) + self.device = device req_pool_size = self.req_to_token.shape[0] self.req_index_to_mamba_index_mapping: torch.Tensor = torch.zeros( @@ -821,6 +835,12 @@ class HybridReqToTokenPool(ReqToTokenPool): logger.info("Reset HybridReqToTokenPool") super().clear() self.mamba_allocator.clear() + # The int8 checkpoint pool holds radix-cached states in its own slots; a + # flush/reset drops the radix tree, so its slots must be released too, + # otherwise the (now unreferenced) slots leak and break the int8-pool + # invariant (int8_available + radix_cached != int8_total). + if self.mamba_ckpt_pool is not None: + self.mamba_ckpt_pool.clear() self.req_index_to_mamba_index_mapping.zero_() if self.enable_mamba_extra_buffer: self.req_index_to_mamba_ping_pong_track_buffer_mapping.zero_() diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 4a0eb4421..752b09c51 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -694,6 +694,12 @@ class ServerArgs: mamba_full_memory_ratio: float = 0.9 mamba_scheduler_strategy: str = "auto" mamba_track_interval: int = 256 + # int8-compress radix-cached linear-attn (mamba) checkpoints -> ~2x cached + # prefixes at fixed memory (quality-safe; see mem_cache/mamba_checkpoint_pool.py). + enable_int8_mamba_checkpoint: bool = False + int8_mamba_ckpt_size: Optional[int] = ( + None # #int8 checkpoint slots; default 2x the active pool + ) linear_attn_backend: str = "triton" linear_attn_decode_backend: Optional[str] = None linear_attn_prefill_backend: Optional[str] = None @@ -1009,6 +1015,7 @@ class ServerArgs: self._handle_deterministic_inference() self._handle_attention_backend_compatibility() self._handle_mamba_backend() + self._handle_int8_mamba_checkpoint() self._handle_linear_attn_backend() self._handle_kv4_compatibility() self._handle_page_size() @@ -3376,6 +3383,28 @@ class ServerArgs: "FlashInfer mamba module not available, please check flashinfer installation." ) + def _handle_int8_mamba_checkpoint(self): + # The int8 mamba checkpoint pool is only wired into the built-in + # MambaRadixCache. The host-offload variant (HiMambaRadixCache, enabled by + # --enable-hierarchical-cache) and custom radix-cache backends are NOT + # int8-aware: they would read int8 checkpoint slots as bf16 active slots + # (wrong pool / out-of-range). Reject the combination up front rather than + # silently corrupting state. + if not self.enable_int8_mamba_checkpoint: + return + if self.enable_hierarchical_cache: + raise ValueError( + "--enable-int8-mamba-checkpoint is not supported together with " + "--enable-hierarchical-cache: the host-offload path " + "(HiMambaRadixCache) is not int8-aware. Disable one of them." + ) + if self.radix_cache_backend is not None: + raise ValueError( + "--enable-int8-mamba-checkpoint only supports the built-in mamba " + f"radix cache; --radix-cache-backend={self.radix_cache_backend!r} " + "is not int8-aware. Omit --radix-cache-backend." + ) + def _handle_linear_attn_backend(self): import torch @@ -6444,6 +6473,19 @@ class ServerArgs: default=ServerArgs.max_mamba_cache_size, help="The maximum size of the mamba cache.", ) + parser.add_argument( + "--enable-int8-mamba-checkpoint", + action="store_true", + help="Store radix-cached linear-attn (mamba) states in int8 (separate " + "checkpoint pool) for ~2x cached-prefix capacity at fixed memory.", + ) + parser.add_argument( + "--int8-mamba-ckpt-size", + type=int, + default=ServerArgs.int8_mamba_ckpt_size, + help="Number of int8 mamba checkpoint slots (default: 2x the active " + "mamba pool size).", + ) parser.add_argument( "--mamba-ssm-dtype", type=str, diff --git a/test/registered/radix_cache/test_int8_mamba_checkpoint_e2e.py b/test/registered/radix_cache/test_int8_mamba_checkpoint_e2e.py new file mode 100644 index 000000000..02db313c6 --- /dev/null +++ b/test/registered/radix_cache/test_int8_mamba_checkpoint_e2e.py @@ -0,0 +1,93 @@ +""" +End-to-end test for the int8 mamba checkpoint pool on a real GDN-hybrid model. + +Launches Qwen3-Next-80B-A3B (a gated-delta-net / linear-attention hybrid) with +``--enable-int8-mamba-checkpoint`` and checks two things while the int8 dequant +path is exercised: + + * KLDivergenceMixin — on a prefix/decode cache HIT the generated logprobs are + compared (KL) against a full recompute. This is the *sensitive* precision + guard: it directly bounds how far the int8-reused state moves the output + distribution from the exact-recompute distribution. + * test_gsm8k — end-to-end task accuracy holds. + +NOTE: the int8 checkpoint is only engaged when a cached prefix is reused FROM the +int8 pool, which requires ``--mamba-scheduler-strategy extra_buffer`` — the default +``no_buffer`` only snapshots the recurrent state at the full-sequence leaf, so a +fixed-prefix / divergent-question workload reuses ~0 mamba state and the int8 path +would never fire. + +Usage: + python3 -m unittest test_int8_mamba_checkpoint_e2e +""" + +import unittest +from types import SimpleNamespace +from urllib.parse import urlparse + +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.kits.kl_divergence_kit import KLDivergenceMixin +from sglang.test.server_fixtures.default_fixture import DefaultServerBase +from sglang.test.test_utils import DEFAULT_HYBRID_MAMBA_MODEL_NAME_FOR_TEST + +register_cuda_ci(est_time=400, stage="base-c", runner_config="4-gpu-h100") + + +class TestInt8MambaCheckpointE2E(KLDivergenceMixin, DefaultServerBase): + """int8 mamba checkpoint pool on Qwen3-Next-80B-A3B (GDN-hybrid).""" + + model = DEFAULT_HYBRID_MAMBA_MODEL_NAME_FOR_TEST + + # Cache-hit KL: int8 is a lossy codec, so its cache-hit divergence is + # inherently larger than the bf16/fp8 reuse the other KL tests bound (~0.005), + # and it grows with context length (a longer prefix = a fuller state = larger + # absolute rounding error in the logits). Measured on a Qwen3.5-35B stand-in + # over LongBench-V2 prompts: prefill ~0.044, decode ~0.024. Thresholds are set + # to ~2x that, to cover model differences (80B) and the reuse path's + # run-to-run noise while still catching a real int8 regression. + kl_div_thres = 0.06 + kl_div_thres_prefill = 0.10 + kl_div_thres_decode = 0.06 + kl_div_max_samples = 16 + kl_div_prefill_max_new_tokens = 512 + kl_div_decode_max_new_tokens = 512 + + gsm8k_threshold = 0.90 + num_gsm8k_questions = 100 + num_shots = 8 + parallel = 8 + + other_args = [ + "--trust-remote-code", + "--tp-size", + "4", + "--mem-fraction-static", + "0.7", + "--enable-int8-mamba-checkpoint", + "--mamba-scheduler-strategy", + "extra_buffer", + ] + + def test_gsm8k(self): + from sglang.test.few_shot_gsm8k import run_eval as run_few_shot_gsm8k + + url = urlparse(self.base_url) + args = SimpleNamespace( + num_shots=self.num_shots, + data_path=None, + num_questions=self.num_gsm8k_questions, + max_new_tokens=512, + parallel=self.parallel, + host=f"http://{url.hostname}", + port=int(url.port), + ) + metrics = run_few_shot_gsm8k(args) + print( + f"[{self.__class__.__name__}] GSM8K accuracy: {metrics['accuracy']:.3f} " + f"(threshold: {self.gsm8k_threshold})" + ) + self.assertGreaterEqual(metrics["accuracy"], self.gsm8k_threshold) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/srt/mem_cache/test_int8_checkpoint_store.py b/test/srt/mem_cache/test_int8_checkpoint_store.py new file mode 100644 index 000000000..11d38fb3a --- /dev/null +++ b/test/srt/mem_cache/test_int8_checkpoint_store.py @@ -0,0 +1,178 @@ +"""Tests for Int8CheckpointStore (int8-compressed cached linear-attn states). + +CPU tests cover the codec error bound, store/load round-trip, and the active-pool +copy-on-write helpers. The CUDA test reproduces the validated decode-output error +(int8 checkpoint loaded once then decoded continues bf16) ~ 0.5%, far below the +bf16-baseline-relative threshold that GSM8K showed is quality-safe. + + python -m pytest test/srt/mem_cache/test_int8_checkpoint_store.py -v +""" + +import unittest + +import torch + +from sglang.srt.mem_cache.mamba_checkpoint_pool import ( + Int8CheckpointStore, + MambaCheckpointPool, +) + +H, V, K = 32, 128, 128 +L = 4 + + +def _rand_state(n, device="cpu"): + # KDA-like state magnitudes (see fp8_checkpoint_probe: |S| mean ~6e-2) + return torch.randn(L, n, H, V, K, device=device) * 6e-2 + + +class TestInt8CheckpointCodec(unittest.TestCase): + def test_quantize_dequantize_error_bound(self): + s = _rand_state(8) + q, scale = Int8CheckpointStore.quantize(s) + self.assertEqual(q.dtype, torch.int8) + self.assertEqual(scale.shape, (L, 8, H, 1, K)) # per (layer,slot,head,k-chan) + deq = Int8CheckpointStore.dequantize(q, scale, torch.float32) + rel = (deq - s).norm() / s.norm() + # uniform int8 per-channel on a ~uniform state: well under 1% + self.assertLess(rel.item(), 1e-2, f"int8 codec rel err too high: {rel}") + + def test_symmetric_and_zero(self): + s = torch.zeros(L, 1, H, V, K) + q, scale = Int8CheckpointStore.quantize(s) + self.assertTrue(torch.equal(q, torch.zeros_like(q))) + deq = Int8CheckpointStore.dequantize(q, scale, torch.float32) + self.assertTrue(torch.equal(deq, s)) + + def test_store_load_roundtrip(self): + store = Int8CheckpointStore( + num_layers=L, + num_slots=16, + num_heads=H, + head_v_dim=V, + head_k_dim=K, + device="cpu", + ) + s = _rand_state(4) + slots = torch.tensor([1, 3, 5, 7]) + store.store(slots, s) + out = store.load(slots, torch.float32) + # load == dequant of stored + q, scale = Int8CheckpointStore.quantize(s) + ref = Int8CheckpointStore.dequantize( + q, scale.to(store.scale.dtype), torch.float32 + ) + self.assertLess((out - ref).abs().max().item(), 1e-3) + + def test_cow_helpers(self): + store = Int8CheckpointStore( + num_layers=L, + num_slots=16, + num_heads=H, + head_v_dim=V, + head_k_dim=K, + device="cpu", + ) + active = torch.zeros(L, 10, H, V, K) # bf16/fp32 active pool + active[:, 2] = _rand_state(1).squeeze(1) + # store active slot 2 -> ckpt slot 4 + store.store_from_pool(active, torch.tensor([2]), torch.tensor([4])) + # load ckpt slot 4 -> active slot 6 (cache-hit COW) + store.copy_to_pool(active, torch.tensor([4]), torch.tensor([6])) + rel = (active[:, 6] - active[:, 2]).norm() / active[:, 2].norm() + self.assertLess(rel.item(), 1e-2) + + def test_memory_is_half_of_bf16(self): + store = Int8CheckpointStore( + num_layers=L, + num_slots=100, + num_heads=H, + head_v_dim=V, + head_k_dim=K, + device="cpu", + ) + bf16_per_slot = L * H * V * K * 2 + # int8 data (1B) + small per-(head,k) bf16 scale -> well under bf16; ~2x slots + self.assertLess(store.bytes_per_slot(), bf16_per_slot * 0.6) + + def test_estimate_matches_actual_mem(self): + # the pre-allocation estimate (used to fit-check HBM before building the + # pool) must equal the real allocated footprint, for any temporal dtype + for tdt in (torch.bfloat16, torch.float32): + kw = dict( + num_layers=L, + num_slots=64, + num_heads=H, + head_v_dim=V, + head_k_dim=K, + conv_shapes=[(4, K)], + conv_dtype=torch.bfloat16, + temporal_dtype=tdt, + ) + est = MambaCheckpointPool.estimate_mem_usage_bytes(**kw) + pool = MambaCheckpointPool(**kw, device="cpu") + self.assertEqual(est["qdata"] + est["scale"] + est["conv"], est["total"]) + self.assertEqual(est["total"], pool.mem_usage_bytes()) + + +@unittest.skipUnless(torch.cuda.is_available(), "needs CUDA + fla kernels") +class TestInt8CheckpointDecodeError(unittest.TestCase): + def test_decode_error_within_bound(self): + try: + from sglang.srt.layers.attention.fla.kda import fused_recurrent_kda + except (ImportError, ModuleNotFoundError) as e: + self.skipTest(f"fla kernels unavailable: {e}") + + dev = "cuda" + torch.manual_seed(0) + + def synth(T, s): + torch.manual_seed(s) + q = torch.randn(1, T, H, K, device=dev, dtype=torch.bfloat16) * 0.5 + k = torch.randn(1, T, H, K, device=dev, dtype=torch.bfloat16) * 0.5 + v = (torch.randn(1, T, H, V, device=dev) * 0.5).bfloat16() + beta = torch.rand(1, T, H, device=dev, dtype=torch.bfloat16) + g = -torch.rand(1, T, H, K, device=dev, dtype=torch.float32) * 0.1 - 0.005 + return q, k, v, g, beta + + def decode(state, inp): + st = state.clone() + o, _ = fused_recurrent_kda( + q=inp[0], + k=inp[1], + v=inp[2], + g=inp[3], + beta=inp[4], + scale=K**-0.5, + initial_state=st, + inplace_final_state=True, + use_qk_l2norm_in_kernel=True, + cu_seqlens=None, + ) + return o.float() + + S = torch.zeros(1, H, V, K, device=dev, dtype=torch.float32) + pre = synth(512, 0) + fused_recurrent_kda( + q=pre[0], + k=pre[1], + v=pre[2], + g=pre[3], + beta=pre[4], + scale=K**-0.5, + initial_state=S, + inplace_final_state=True, + use_qk_l2norm_in_kernel=True, + cu_seqlens=None, + ) + dec = synth(128, 1) + o_ref = decode(S, dec) + q, scale = Int8CheckpointStore.quantize(S) # [1,H,V,K] + S_int8 = Int8CheckpointStore.dequantize(q, scale, torch.float32) + o_int8 = decode(S_int8, dec) + rel = (o_int8 - o_ref).norm() / o_ref.norm() + self.assertLess(rel.item(), 1.5e-2, f"int8 decode err {rel} too high") + + +if __name__ == "__main__": + unittest.main()