[diffusion] feat: support quant-videogen prq kv-cache quantization (memory-saving) for causal-dit (#32581)

Co-authored-by: Claude Opus 4.8 <noreply@anthropic.com>
Co-authored-by: Mick <mickjagger19@icloud.com>
This commit is contained in:
Haocheng Xi
2026-08-08 12:56:54 +08:00
committed by GitHub
co-authored by Claude Opus 4.8 Mick
parent 24c84dfa68
commit f64328c7f6
16 changed files with 1174 additions and 54 deletions
@@ -32,6 +32,16 @@ sglang serve \
--enable-torch-compile false
```
### Optional KV-Cache Compression
Long-running sessions can enable lossy int4 PRQ compression for completed
causal KV-cache chunks by installing `quant-videogen` as described in the
quantization guide and adding `--kv-cache-quant int4` to the server command.
The current and most recent completed chunks remain in BF16. See
[Causal KV-Cache Quantization](/docs/sglang-diffusion/quantization#causal-kv-cache-quantization)
for the algorithm, tuning options, measured memory-latency tradeoff, and
support limits.
## 3. Realtime WebUI
The lightweight local WebUI is useful for validating latency, frame transport, and camera control behavior.
@@ -27,6 +27,16 @@ sglang serve \
--text-encoder-cpu-offload false
```
### Optional KV-Cache Compression
Long-running sessions can enable lossy int4 PRQ compression for completed
causal KV-cache chunks by installing `quant-videogen` as described in the
quantization guide and adding `--kv-cache-quant int4` to the server command.
The current and most recent completed chunks remain in BF16. See
[Causal KV-Cache Quantization](/docs/sglang-diffusion/quantization#causal-kv-cache-quantization)
for the algorithm, tuning options, measured memory-latency tradeoff, and
support limits.
## 3. Realtime WebUI
The lightweight local WebUI is useful for validating latency, frame transport, and camera control behavior.
+1
View File
@@ -1558,6 +1558,7 @@
"pages": [
"docs/sglang-diffusion/api/cli",
"docs/sglang-diffusion/api/openai_api",
"docs/sglang-diffusion/realtime_models",
"docs/sglang-diffusion/models_with_ar",
"docs/sglang-diffusion/models_with_pe",
"docs/sglang-diffusion/api/post_processing"
+8 -2
View File
@@ -114,7 +114,7 @@ Use `sglang generate --help` and `sglang serve --help` for the full argument lis
For frame interpolation and upscaling, see [Post-Processing](./post_processing).
### Quantized transformers
### Quantization
For quantized transformer checkpoints, prefer:
@@ -124,7 +124,13 @@ For quantized transformer checkpoints, prefer:
- `--quantization` for online quantization (apply quantization to unquantized models at load time, activations are quantized dynamically)
- `--quantization-ignored-layers` layer name patterns to keep unquantized (e.g. `attention.to_`)
See [Quantization](../quantization) for supported quantization families and examples.
For supported realtime causal video models, `--kv-cache-quant {off|int4|int2}`
compresses completed KV-cache chunks independently of transformer weight
quantization. It is lossy and disabled by default.
See [Realtime and Causal Video Models](../realtime_models) for the runtime and
model scope, and [Quantization](../quantization) for supported quantization
families and examples.
### Request logging
+2 -1
View File
@@ -35,7 +35,8 @@ sglang serve --model-path Qwen/Qwen-Image --port 30010
- [OpenAI-Compatible API](/docs/sglang-diffusion/api/openai_api): send image and video requests to the HTTP server
- [Performance Overview](/docs/sglang-diffusion/performance-optimization): choose speed, memory, parallelism, caching, and quality-tradeoff levers
- [Caching Acceleration](/docs/sglang-diffusion/caching-acceleration): use Cache-DiT, TeaCache, or Spectrum to reduce denoising cost
- [Quantization](/docs/sglang-diffusion/quantization): load quantized transformer checkpoints
- [Quantization](/docs/sglang-diffusion/quantization): configure transformer weight and causal KV-cache quantization
- [Realtime and Causal Video Models](/docs/sglang-diffusion/realtime_models): understand session state, causal caches, and realtime-only controls
- [Contributing](/docs/sglang-diffusion/contributing): contribution workflow, adding new models, and CI perf baselines
## Additional Documentation
+106 -1
View File
@@ -2,7 +2,7 @@
title: "Quantization"
tag: "approx"
metatags:
description: "SGLang-Diffusion supports quantized transformer checkpoints. In most cases, keep the base model and the quantized transformer override separate."
description: "Configure transformer weight quantization and Quant-VideoGen causal KV-cache quantization in SGLang-Diffusion."
---
SGLang-Diffusion supports quantized transformer checkpoints. In most cases, keep
@@ -17,6 +17,7 @@ Use these paths:
- `--transformer-weights-path`: quantized transformer weights provided as a single safetensors file, a sharded safetensors directory, a local path, or a Hugging Face repo ID
- `--quantization`: apply online quantization to unquantized models at load time (activations are quantized dynamically)
- `--quantization-ignored-layers` layer name patterns to keep unquantized (e.g. `attention.to_`)
- `--kv-cache-quant`: compress completed causal KV-cache chunks for supported realtime models
Recommended example for pre-quantized checkpoints:
@@ -99,6 +100,14 @@ backend.
<td>None</td>
<td>Mixed override repos keep the base model separate; full Qwen Image exports can be loaded directly as <code>--model-path</code>; raw exports such as <code>black-forest-labs/FLUX.2-dev-NVFP4</code> still use the weights-path flow</td>
</tr>
<tr>
<td><code>qvg-kv</code></td>
<td>Unquantized model with runtime causal KV-cache compression</td>
<td><code>--kv-cache-quant &#123;int4,int2&#125;</code></td>
<td>LingBot World realtime causal path</td>
<td><code>quant-videogen</code></td>
<td>CUDA only; compresses completed cache chunks rather than model weights; lossy and disabled by default</td>
</tr>
<tr>
<td><code>nunchaku-svdq</code></td>
<td>Pre-quantized Nunchaku transformer weights, usually named <code>svdq-&#123;int4\|fp4&#125;_r&#123;rank&#125;-...</code></td>
@@ -118,6 +127,102 @@ backend.
</tbody>
</table>
## Causal KV-Cache Quantization
Quant-VideoGen KV-cache quantization targets long-running autoregressive video
sessions, where the causal self-attention cache can become comparable to the
model weights. It does not change or quantize the checkpoint weights.
See [Realtime and Causal Video Models](./realtime_models) for the session
lifecycle, supported pipelines, and the distinction between realtime and
request-based causal generation.
Install the optional dependency without allowing its stale Torch requirement to
replace SGLang's pinned Torch version, then enable int4 compression when serving
a supported LingBot World realtime pipeline:
```bash
pip install "sglang[diffusion,diffusion-qvg]"
pip install --no-deps quant-videogen==0.1.0
sglang serve \
--model-path robbyant/lingbot-world-fast-diffusers \
--pipeline-class-name LingBotWorldCausalDMDPipeline \
--num-gpus 4 \
--ulysses-degree 4 \
--kv-cache-quant int4 \
--dit-cpu-offload false \
--text-encoder-cpu-offload false
```
### Storage Policy
The current chunk is rewritten at every denoising step, so it remains in BF16.
The newest `--kv-cache-quant-keep-recent` completed chunks also remain in BF16.
Older completed chunks are stable and are packed once with Progressive Residual
Quantization (PRQ); their dense BF16 tensors are then released.
When a transformer layer runs attention, its packed visible chunks are
dequantized and concatenated with the recent BF16 chunks. This creates one
layer's dense attention view at a time instead of keeping dense windows
resident for every transformer layer.
### How PRQ Works
For each K or V vector, PRQ uses k-means to select a centroid, then quantizes
the remaining error:
```text
x = centroid_1 + residual_1
residual_1 = centroid_2 + residual_2
...
x_hat = centroid_1 + centroid_2 + ... + dequantize(low_bit_residual)
```
Each additional stage applies another centroid lookup to the previous stage's
residual. SGLang's default uses one stage, 128 centroids, and an int4 or int2
block-quantized residual. More stages or centroids can reduce reconstruction
error but add codebook storage and packing work.
PRQ is the compression algorithm; selecting older completed chunks is the
runtime storage policy that makes it practical. Stable chunks are compressed
once, while mutable and recent chunks avoid repeated packing and retain higher
precision.
### Quality And Performance
<Warning>
KV-cache quantization is lossy. Disabling it uses the original dense BF16 cache
and is bit-exact with the unmodified path. Enabling int4 or int2 reconstructs an
approximation of K and V, so fixed-seed generated frames are not expected to be
pixel-identical to BF16.
</Warning>
In the initial LingBot measurements, int4 used about 47% of the dense resident
KV-cache memory for a 24-frame window and added about 18% per-chunk latency.
Int2 used about 37% of the dense resident KV-cache memory but introduces more
quantization error. These measurements are configuration-specific; benchmark
memory, latency, temporal consistency, identity stability, and motion quality
on the intended session length. Start with int4 unless capacity requires int2.
The current implementation is limited to the LingBot realtime
sliding-window-and-sink path, including Ulysses sequence sharding. It does not
support LongLive2 pinned sinks, global sinks, or dynamically growing caches.
### Tuning Options
| Option | Default | Effect |
| --- | ---: | --- |
| `--kv-cache-quant {off,int4,int2}` | `off` | Enables QVG KV-cache compression and selects residual precision. |
| `--kv-cache-quant-stages` | `1` | Number of progressive centroid-residual stages. |
| `--kv-cache-quant-centroids` | `128` | Number of k-means centroids per stage. |
| `--kv-cache-quant-block-size` | `64` | Block size used to quantize the final residual. |
| `--kv-cache-quant-iters` | `2` | K-means iterations used while packing a chunk. |
| `--kv-cache-quant-asymmetric` | disabled | Uses asymmetric residual quantization. |
| `--kv-cache-quant-keep-recent` | `1` | Number of newest completed chunks retained in BF16. |
| `--kv-cache-quant-sink {0,1}` | `1` | Whether to quantize completed sink chunks. |
| `--kv-cache-quant-sink-keep` | `0` | Number of leading sink chunks retained in BF16. |
## Online Quantization
Online quantization applies quantization to unquantized models at load time. This is useful for when pre-quantized checkpoints are not available.
@@ -0,0 +1,78 @@
---
title: "Realtime and Causal Video Models"
metatags:
description: "Deploy session-based realtime and request-based causal video models with SGLang Diffusion."
---
Realtime and causal video pipelines generate video incrementally and reuse state
across chunks. This differs from offline diffusion pipelines, which denoise one
bounded latent sequence and release all request state when generation finishes.
## Execution Modes
SGLang Diffusion exposes two related but distinct modes:
| Mode | Lifetime | Interface | Examples |
| --- | --- | --- | --- |
| Realtime session | State persists until the client disconnects or the session ends | `/v1/realtime_video/generate` WebSocket | LingBot World, SANA-WM realtime |
| Request-based causal generation | State is reused across chunks within one request, then released | Standard video generation API | LongLive 2.0, batch-streaming SANA-WM |
The realtime server retains model-specific state such as the causal self-attention
KV cache, cross-attention cache, decoder history, and pending control events.
State is isolated per session and is not reused by unrelated requests.
<Note>
A causal DiT is not automatically a realtime session model. The pipeline must
also register a realtime adapter and implement the WebSocket session lifecycle.
</Note>
## Supported Realtime Pipelines
| Model family | Pipeline | Live controls | QVG KV-cache quantization |
| --- | --- | --- | --- |
| LingBot World | `LingBotWorldCausalDMDPipeline` | Camera actions and prompt updates | Supported |
| SANA-WM | `SanaWMRealtimePipeline` | Camera actions | Not supported |
Use the model cookbooks for launch commands, request schemas, and control-token
details:
- [LingBot World](/cookbook/diffusion/LingBot-World/LingBot-World)
- [LingBot World 2.0](/cookbook/diffusion/LingBot-World/LingBot-World-2.0)
- [SANA-WM](/cookbook/diffusion/SANA-WM/SANA-WM)
For the complete model list, see
[Supported Models and Optimization Compatibility](./compatibility_matrix).
## Causal Cache Controls
Realtime requests can override two model defaults:
- `realtime_causal_sink_size`: amount of stable prefix history retained as an attention sink
- `realtime_causal_kv_cache_num_frames`: recent causal history retained in the rolling KV-cache window
Larger windows preserve more history but increase resident memory and attention
work. These fields are request/session controls; supported ranges and defaults
remain model-specific.
For supported LingBot World deployments, the server-level
`--kv-cache-quant {off,int4,int2}` option compresses completed cache chunks.
It is disabled by default and is lossy when enabled. Start with `int4`; use
`int2` only when the additional memory reduction is worth the larger quality
risk.
See [Causal KV-Cache Quantization](./quantization#causal-kv-cache-quantization)
for installation, storage policy, tuning options, memory/latency tradeoffs, and
current limitations.
<Warning>
QVG KV-cache quantization currently supports only the LingBot realtime
sliding-window-and-sink path. It does not apply to SANA-WM realtime, LongLive 2.0
pinned sinks, global sinks, or dynamically growing caches.
</Warning>
## Deployment Considerations
- Keep `--kv-cache-quant off` when bit-exact BF16 cache behavior is required.
- Benchmark a representative session length. Short clips may not exercise cold-cache packing and can hide both its memory benefit and packing overhead.
- Treat sequence parallelism as model-specific. Follow the model cookbook and the [Sequence Parallelism](./ring_sp_performance) guide instead of assuming one mesh is best for every realtime pipeline.
- Realtime WebSocket clients must send an initialization message before control events. The exact MessagePack schema and output encoding are documented in each model cookbook.
+8
View File
@@ -123,6 +123,14 @@ diffusion = [
"xatlas",
]
diffusion-qvg = [
# quant-videogen 0.1.0 pins an incompatible Torch version, so it is
# installed separately with --no-deps; this extra provides its missing
# import-time dependencies.
"loguru>=0.7",
"termcolor>=2.3",
]
ray = [
"ray[default]>=2.55.1",
]
@@ -0,0 +1,101 @@
# SPDX-License-Identifier: Apache-2.0
"""CLI-facing configuration for Quant-VideoGen PRQ KV-cache quantization.
Mirrors the SRT `--kv-cache-dtype` pattern: the on/off + tuning knobs live on a
typed config object carried by ServerArgs (see `kv_cache_quant_config`), instead
of a pile of raw environment variables.
Defaults match the tuned per-chunk setting
(kmeans 1 stage, 128 centroids, block 64, symmetric, 2 iters, recent 1,
per-chunk sink) — i.e. `--kv-cache-quant int4` alone reproduces it.
"""
from __future__ import annotations
from dataclasses import dataclass
_BITS = {"off": None, "none": None, "bf16": None, "int4": 4, "int2": 2}
def _parse_bits(val: str | None) -> int | None:
if val is None:
return None
key = str(val).strip().lower()
if key not in _BITS:
raise ValueError(f"kv-cache-quant must be one of {list(_BITS)}, got {val!r}")
return _BITS[key]
@dataclass
class QVGKVQuantArgs:
"""PRQ (multi-stage k-means) KV-cache quantization settings.
``bits is None`` means quantization is OFF (plain bf16 cache). Defaults
reproduce the tuned per-chunk config from offline sweeps.
"""
bits: int | None = None # None => off; 2 or 4 (master switch)
centroids: int = 128 # k-means centroids per stage
block_size: int = 64 # residual scale block size
stages: int = 1 # PRQ k-means stages
kmeans_iters: int = 2 # k-means iterations
asymmetric: bool = False # KIVI-style asymmetric residual quant
keep_recent_chunks: int = 1 # completed chunks kept bf16 (recency guard)
sink: bool = True # quantize the attention sink too
sink_keep_chunks: int = 0 # leading sink chunks kept bf16 forever
@property
def enabled(self) -> bool:
return self.bits is not None
def validate(self) -> QVGKVQuantArgs:
if self.bits not in (None, 2, 4):
raise ValueError(f"kv-cache-quant bits must be 2 or 4, got {self.bits}")
if self.centroids <= 0:
raise ValueError("kv-cache-quant-centroids must be > 0")
if self.block_size <= 0:
raise ValueError("kv-cache-quant-block-size must be > 0")
if self.stages <= 0:
raise ValueError("kv-cache-quant-stages must be > 0")
if self.kmeans_iters <= 0:
raise ValueError("kv-cache-quant-iters must be > 0")
if self.keep_recent_chunks < 0 or self.sink_keep_chunks < 0:
raise ValueError("keep_recent_chunks / sink_keep_chunks must be >= 0")
return self
def describe(self) -> str:
if not self.enabled:
return "off"
return (
f"int{self.bits} centroids={self.centroids} block={self.block_size} "
f"stages={self.stages} iters={self.kmeans_iters} "
f"asym={self.asymmetric} recent={self.keep_recent_chunks} "
f"sink={self.sink} sink_keep={self.sink_keep_chunks}"
)
@classmethod
def from_dict(cls, kwargs: dict) -> QVGKVQuantArgs:
"""Build from flat CLI kwargs (dest names ``kv_cache_quant*``)."""
master = kwargs.get("kv_cache_quant")
if master is None:
return cls()
inst = cls(bits=_parse_bits(master))
if kwargs.get("kv_cache_quant_centroids") is not None:
inst.centroids = kwargs["kv_cache_quant_centroids"]
if kwargs.get("kv_cache_quant_block_size") is not None:
inst.block_size = kwargs["kv_cache_quant_block_size"]
if kwargs.get("kv_cache_quant_stages") is not None:
inst.stages = kwargs["kv_cache_quant_stages"]
if kwargs.get("kv_cache_quant_iters") is not None:
inst.kmeans_iters = kwargs["kv_cache_quant_iters"]
if kwargs.get("kv_cache_quant_keep_recent") is not None:
inst.keep_recent_chunks = kwargs["kv_cache_quant_keep_recent"]
if kwargs.get("kv_cache_quant_sink_keep") is not None:
inst.sink_keep_chunks = kwargs["kv_cache_quant_sink_keep"]
if kwargs.get("kv_cache_quant_asymmetric") is not None:
inst.asymmetric = kwargs["kv_cache_quant_asymmetric"]
if kwargs.get("kv_cache_quant_sink") is not None:
inst.sink = kwargs["kv_cache_quant_sink"]
inst.sink = bool(inst.sink)
inst.asymmetric = bool(inst.asymmetric)
return inst.validate()
@@ -42,6 +42,10 @@ class CausalSelfAttentionKVCache:
if self.attention_window_size == 0:
self.attention_window_size = self.cache_size
@property
def num_cache_heads(self) -> int:
return self.k.shape[2]
def reset_indices(self) -> None:
self.global_end_index.zero_()
self.local_end_index.zero_()
@@ -0,0 +1,401 @@
# SPDX-License-Identifier: Apache-2.0
"""Quant-VideoGen KV cache with PRQ-packed storage for completed chunks.
Storage model (mirrors Quant-VideoGen's ChunkedKVCache, fitted to SGLang's
``update_and_get_attention_kv`` contract):
* The retained window is split into frame/chunk-aligned *segments* in global
token order. Each segment is either BF16 (`k`/`v` tensors resident) or
PRQ-packed (`packed_k`/`packed_v` dicts resident, BF16 freed).
* The current (still-denoising) chunk and the newest ``keep_recent_chunks``
completed chunks stay BF16 (rewritten each denoise step / attended cleanly).
* Older completed segments are PRQ-packed once and their BF16 freed -> the
resident footprint drops to ~(sink + recent) BF16 + packed tail.
* On read the visible window is reconstructed densely on the fly (dequantize
packed segments + cat BF16 ones) and returned to attention; that transient
dense tensor is freed after the layer's attention, so only ONE layer is
dense at a time vs. all ``num_layers`` resident dense windows before.
Scope: the LingBot realtime causal path only (sliding window + sink, chunk-
aligned writes, optional ulysses head-slice, ``recent_window_tokens`` None or a
non-negative int). Unsupported base-class features raise NotImplementedError
rather than silently diverge.
"""
from __future__ import annotations
from dataclasses import dataclass
from functools import cache
import torch
from sglang.multimodal_gen.configs.quantization.qvg_kv import QVGKVQuantArgs
from sglang.multimodal_gen.runtime.layers.kvcache.causal_attention_cache import (
CausalAttentionKVView,
)
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__)
@cache
def _qvg_functions():
try:
from quant_videogen.functions import (
triton_prq_dequantize_tensor,
triton_prq_quantize_tensor,
)
except ImportError as e:
raise ImportError(
"Quant-VideoGen KV-cache quantization requires its optional "
"runtime dependencies. Install them with: "
"pip install 'sglang[diffusion-qvg]' && "
"pip install --no-deps quant-videogen==0.1.0."
) from e
return triton_prq_quantize_tensor, triton_prq_dequantize_tensor
@dataclass
class _Segment:
g0: int # global start token (inclusive)
g1: int # global end token (exclusive)
is_sink: bool # sink segments are never evicted
k: torch.Tensor | None = None
v: torch.Tensor | None = None
packed_k: dict | None = None
packed_v: dict | None = None
@property
def packed(self) -> bool:
return self.packed_k is not None
def nbytes(self) -> int:
if self.packed:
return _packed_nbytes(self.packed_k) + _packed_nbytes(self.packed_v)
return (
self.k.numel() * self.k.element_size()
+ self.v.numel() * self.v.element_size()
)
def _packed_nbytes(packed: dict) -> int:
total = 0
for key in ("centroids_list", "cluster_ids_list"):
for t in packed.get(key) or []:
total += t.numel() * t.element_size()
for key in ("residual_quant", "scales", "zeros", "residual", "scale_factor"):
t = packed.get(key)
if isinstance(t, torch.Tensor):
total += t.numel() * t.element_size()
return total
class QVGPackedCausalKVCache:
"""Chunk-segmented causal KV cache with PRQ-packed cold segments."""
def __init__(
self,
*,
batch_size: int,
cache_size: int,
num_heads: int,
head_dim: int,
dtype: torch.dtype,
device: torch.device,
global_end_index: torch.Tensor,
local_end_index: torch.Tensor,
use_int_indices: bool = False,
sink_tokens: int = 0,
attention_window_size: int | None = None,
quant_args: QVGKVQuantArgs,
) -> None:
self.batch_size = batch_size
self.cache_size = cache_size
self.num_heads = num_heads
self.head_dim = head_dim
self.dtype = dtype
self.device = device
self.sink_tokens = sink_tokens
self.global_sink_tokens = 0
self.attention_window_size = attention_window_size or cache_size
self.q = quant_args
# kept for API compatibility with the dense cache (indices are not read
# by consumers, but reset/patterns touch them)
self.global_end_index = global_end_index
self.local_end_index = local_end_index
self.global_end_index_int = 0 if use_int_indices else None
self.local_end_index_int = 0 if use_int_indices else None
self._segments: list[_Segment] = [] # completed, global-ordered
self._cur: _Segment | None = None # current (mutable) chunk
self._global_end = 0
self._chunk_tokens = 0 # inferred from first advance
# ------------------------------------------------------------------ api
def reset_indices(self) -> None:
self._segments = []
self._cur = None
self._global_end = 0
if self.global_end_index_int is not None:
self.global_end_index_int = 0
self.local_end_index_int = 0
self.global_end_index.zero_()
self.local_end_index.zero_()
def can_direct_current_attention(self, num_new_tokens: int) -> bool:
return (
self.sink_tokens == 0
and self.cache_size == num_new_tokens
and self.attention_window_size == num_new_tokens
)
@property
def num_cache_heads(self) -> int:
return self.num_heads
def pin_current_chunk(self, current_num_tokens: int) -> None:
raise NotImplementedError(
"QVGPackedCausalKVCache does not support pinned-sink (longlive2); "
"packed KV quant is scoped to the LingBot realtime path."
)
def resident_nbytes(self) -> int:
total = sum(s.nbytes() for s in self._segments)
if self._cur is not None:
total += self._cur.nbytes()
return total
# -------------------------------------------------------------- helpers
def _new_bf16_segment(self, g0: int, g1: int, is_sink: bool) -> _Segment:
n = g1 - g0
k = torch.zeros(
self.batch_size,
n,
self.num_heads,
self.head_dim,
dtype=self.dtype,
device=self.device,
)
v = torch.zeros_like(k)
return _Segment(g0=g0, g1=g1, is_sink=is_sink, k=k, v=v)
def _write(self, seg: _Segment, key, value, head_slice) -> None:
if head_slice is None:
seg.k.copy_(key)
seg.v.copy_(value)
else:
seg.k[:, :, head_slice, :] = key
seg.v[:, :, head_slice, :] = value
def _pack(self, seg: _Segment) -> None:
if seg.packed or seg.k is None:
return
triton_prq_quantize_tensor, _ = _qvg_functions()
def q(x):
xb = x.permute(0, 2, 1, 3).contiguous() # [B,S,H,D]->[B,H,S,D]
devices = [xb.device] if xb.is_cuda else []
with torch.random.fork_rng(devices=devices):
torch.manual_seed(1234)
return triton_prq_quantize_tensor(
xb,
num_stages=self.q.stages,
num_clusters=self.q.centroids,
block_size=self.q.block_size,
max_iters=self.q.kmeans_iters,
quantize_fn=lambda _t: self.q.bits,
asymmetric=self.q.asymmetric,
)
seg.packed_k = q(seg.k)
seg.packed_v = q(seg.v)
seg.k = None
seg.v = None
logger.info_once(f"Using QVG packed KV cache: {self.q.describe()}")
def _dequant(self, packed: dict) -> torch.Tensor:
_, triton_prq_dequantize_tensor = _qvg_functions()
return triton_prq_dequantize_tensor(
packed, self.q.block_size, self.q.bits, output_dtype=self.dtype
) # [B,H,S,D]
def _all_segments(self) -> list[_Segment]:
segs = list(self._segments)
if self._cur is not None:
segs.append(self._cur)
return segs
def _sink_end(self) -> int:
return min(self.sink_tokens, self._global_end)
def _tail_start(self) -> int:
"""Global start of the rolling recent tail (sink occupies its own
budget at the window front, matching the dense cache's roll)."""
sink_end = self._sink_end()
recent_budget = max(0, self.attention_window_size - sink_end)
return max(sink_end, self._global_end - recent_budget)
def _pack_and_evict(self) -> None:
"""Pack completed segments older than the recency guard; drop segments
that have slid entirely out of the window (sink is never evicted)."""
tail_start = self._tail_start()
# eviction: drop non-sink segments fully left of the rolling tail
kept = []
for s in self._segments:
if not s.is_sink and s.g1 <= tail_start:
continue
kept.append(s)
self._segments = kept
if not self.q.enabled:
return
# recency guard: keep the newest `keep_recent_chunks` completed
# non-sink chunks in bf16; pack everything older.
recent = self.q.keep_recent_chunks
nonsink = [s for s in self._segments if not s.is_sink]
cutoff_idx = len(nonsink) - recent
for i, s in enumerate(nonsink):
if i < cutoff_idx:
self._pack(s)
# sink packing policy
if self.q.sink:
sink_keep_tokens = self.q.sink_keep_chunks * max(1, self._chunk_tokens)
for s in self._segments:
if s.is_sink and s.g0 >= sink_keep_tokens:
# only pack sink chunks past the protected prefix, and only
# once they are no longer the current recency-recent region
self._pack(s)
# -------------------------------------------------------------- contract
def update_and_get_attention_kv(
self,
*,
key: torch.Tensor,
value: torch.Tensor,
current_chunk_start: int,
cache_head_start: int | None = None,
recent_window_tokens: int | None = None,
debug_name: str = "QVG packed KV cache",
) -> CausalAttentionKVView:
num_new = key.shape[1]
num_input_heads = key.shape[2]
head_slice = None
if num_input_heads != self.num_heads:
if cache_head_start is None:
raise ValueError(
f"{debug_name}: cache_head_start required for head slice"
)
head_slice = slice(cache_head_start, cache_head_start + num_input_heads)
cend = current_chunk_start + num_new
if self._cur is not None and current_chunk_start == self._cur.g0:
# rewrite current chunk in place (denoise step)
if cend != self._cur.g1:
raise NotImplementedError(
f"{debug_name}: current-chunk rewrite size changed"
)
self._write(self._cur, key, value, head_slice)
elif current_chunk_start == self._global_end:
# advance: finalize current chunk, start a new one
if self._cur is not None:
self._segments.append(self._cur)
is_sink = current_chunk_start < self.sink_tokens
if self._chunk_tokens == 0:
self._chunk_tokens = num_new
self._cur = self._new_bf16_segment(current_chunk_start, cend, is_sink)
self._write(self._cur, key, value, head_slice)
self._global_end = cend
self._pack_and_evict()
else:
raise NotImplementedError(
f"{debug_name}: non-sequential write current_start="
f"{current_chunk_start} global_end={self._global_end} "
f"cur={None if self._cur is None else self._cur.g0}"
)
local_end = min(self._global_end, self.cache_size)
if self.global_end_index_int is not None:
self.global_end_index_int = self._global_end
self.local_end_index_int = local_end
else:
self.global_end_index.fill_(self._global_end)
self.local_end_index.fill_(local_end)
vk, vv = self._reconstruct(current_chunk_start, recent_window_tokens)
return CausalAttentionKVView(
k=vk,
v=vv,
local_start_index=0,
local_end_index=num_new,
visible_local_end=min(self._global_end, self.cache_size),
visible_global_end=self._global_end,
)
def _reconstruct(
self,
current_chunk_start: int,
recent_window_tokens: int | None,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Dense visible window = sink prefix ++ rolling recent tail, matching
the dense cache's [sink | rolled-recent] buffer content."""
sink_end = self._sink_end()
if recent_window_tokens is None:
tail_start = self._tail_start()
else:
if recent_window_tokens < 0:
raise ValueError("recent_window_tokens must be >= 0 or None")
tail_start = max(sink_end, current_chunk_start - recent_window_tokens)
if tail_start <= sink_end:
ranges = [(0, self._global_end)]
else:
ranges = [(0, sink_end), (tail_start, self._global_end)]
visible_segments: list[tuple[_Segment, int, int]] = []
visible_tokens = 0
for g_lo, g_hi in ranges:
for seg in self._all_segments():
a = max(g_lo, seg.g0)
b = min(g_hi, seg.g1)
if b <= a:
continue
visible_segments.append((seg, a - seg.g0, b - seg.g0))
visible_tokens += b - a
if len(visible_segments) == 1 and not visible_segments[0][0].packed:
seg, i0, i1 = visible_segments[0]
return seg.k[:, i0:i1], seg.v[:, i0:i1]
output_shape = (
self.batch_size,
visible_tokens,
self.num_heads,
self.head_dim,
)
vk = torch.empty(output_shape, dtype=self.dtype, device=self.device)
vv = torch.empty_like(vk)
output_start = 0
# Dequantize one tensor at a time so reconstruction needs only the
# final dense view plus one segment-sized temporary.
for seg, i0, i1 in visible_segments:
output_end = output_start + i1 - i0
if seg.packed:
dequantized = self._dequant(seg.packed_k)
vk[:, output_start:output_end].copy_(
dequantized[:, :, i0:i1].permute(0, 2, 1, 3)
)
del dequantized
dequantized = self._dequant(seg.packed_v)
vv[:, output_start:output_end].copy_(
dequantized[:, :, i0:i1].permute(0, 2, 1, 3)
)
del dequantized
else:
vk[:, output_start:output_end].copy_(seg.k[:, i0:i1])
vv[:, output_start:output_end].copy_(seg.v[:, i0:i1])
output_start = output_end
return vk, vv
@@ -12,6 +12,9 @@ from sglang.multimodal_gen.runtime.layers.kvcache.causal_attention_cache import
CausalSelfAttentionKVCache,
CrossAttentionKVCache,
)
from sglang.multimodal_gen.runtime.layers.kvcache.qvg_packed_cache import (
QVGPackedCausalKVCache,
)
from sglang.multimodal_gen.runtime.managers.forward_context import set_forward_context
from sglang.multimodal_gen.runtime.pipelines_core.diffusion_scheduler_utils import (
get_or_create_request_scheduler,
@@ -46,6 +49,7 @@ logger = init_logger(__name__)
CAUSAL_BLOCK_PROMPTS_KEY = "causal_block_prompts"
CAUSAL_SCENE_CUT_MASK_KEY = "causal_scene_cut_mask"
CAUSAL_SHOT_INDICES_KEY = "causal_shot_indices"
CausalKVCache = CausalSelfAttentionKVCache | QVGPackedCausalKVCache
def expand_causal_block_prompts(
@@ -134,7 +138,7 @@ class CausalDMDCachePolicy:
class CausalDMDRealtimeCacheContext:
cache_state: RealtimeCausalDiTState
persist_state: bool
kv_cache: list[CausalSelfAttentionKVCache]
kv_cache: list[CausalKVCache]
crossattn_cache: list[CrossAttentionKVCache]
current_start_frame: int
chunk_idx: int
@@ -415,6 +419,17 @@ class CausalDMDDenoisingStage(DenoisingStage):
raise ValueError("realtime_causal_kv_cache_num_frames must be positive")
self.sliding_window_num_frames = int(kv_cache_num_frames)
if (
server_args.kv_cache_quant_config.enabled
and not self._supports_qvg_kv_cache_quantization()
):
raise ValueError(
f"{type(self).__name__} does not support QVG KV-cache quantization"
)
def _supports_qvg_kv_cache_quantization(self) -> bool:
return False
def _causal_sequence_shard_enabled(self, batch: Req) -> bool:
return False
@@ -487,8 +502,8 @@ class CausalDMDDenoisingStage(DenoisingStage):
or crossattn_cache is None
or len(causal_kv_cache) != self.num_transformer_blocks
or len(crossattn_cache) != self.num_transformer_blocks
or causal_kv_cache[0].k.shape[1] != policy.expected_cache_tokens
or causal_kv_cache[0].k.shape[2] != policy.num_attention_heads
or causal_kv_cache[0].cache_size != policy.expected_cache_tokens
or causal_kv_cache[0].num_cache_heads != policy.num_attention_heads
or causal_kv_cache[0].sink_tokens != policy.expected_sink_tokens
)
@@ -1111,46 +1126,27 @@ class CausalDMDDenoisingStage(DenoisingStage):
global_sink_tokens: int = 0,
attention_window_size: int | None = None,
allow_growth: bool = False,
) -> list[CausalSelfAttentionKVCache]:
causal_kv_cache = []
int_index = 0 if use_int_indices else None
) -> list[CausalKVCache]:
if attention_window_size is None:
attention_window_size = kv_cache_size
for _ in range(self.num_transformer_blocks):
causal_kv_cache.append(
CausalSelfAttentionKVCache(
k=torch.zeros(
[
batch_size,
kv_cache_size,
num_attention_heads,
attention_head_dim,
],
dtype=dtype,
device=device,
),
v=torch.zeros(
[
batch_size,
kv_cache_size,
num_attention_heads,
attention_head_dim,
],
dtype=dtype,
device=device,
),
global_end_index=torch.zeros(1, dtype=torch.long, device=device),
local_end_index=torch.zeros(1, dtype=torch.long, device=device),
global_end_index_int=int_index,
local_end_index_int=int_index,
cache_size=kv_cache_size,
sink_tokens=sink_tokens,
global_sink_tokens=global_sink_tokens,
attention_window_size=attention_window_size,
allow_growth=allow_growth,
)
int_index = 0 if use_int_indices else None
shape = [batch_size, kv_cache_size, num_attention_heads, attention_head_dim]
return [
CausalSelfAttentionKVCache(
k=torch.zeros(shape, dtype=dtype, device=device),
v=torch.zeros(shape, dtype=dtype, device=device),
global_end_index=torch.zeros(1, dtype=torch.long, device=device),
local_end_index=torch.zeros(1, dtype=torch.long, device=device),
global_end_index_int=int_index,
local_end_index_int=int_index,
cache_size=kv_cache_size,
sink_tokens=sink_tokens,
global_sink_tokens=global_sink_tokens,
attention_window_size=attention_window_size,
allow_growth=allow_growth,
)
return causal_kv_cache
for _ in range(self.num_transformer_blocks)
]
@torch.no_grad()
def forward(
@@ -12,6 +12,9 @@ from sglang.multimodal_gen.runtime.distributed.parallel_state import (
get_ring_parallel_world_size,
get_ulysses_parallel_world_size,
)
from sglang.multimodal_gen.runtime.layers.kvcache.qvg_packed_cache import (
QVGPackedCausalKVCache,
)
from sglang.multimodal_gen.runtime.managers.forward_context import set_forward_context
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.pipelines_core.stages.causal_denoising import (
@@ -19,6 +22,7 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.causal_denoising import
CausalDMDDenoisingStage,
CausalDMDForwardContext,
CausalDMDRealtimeCacheContext,
CausalKVCache,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.lingbot_world.constants import (
LINGBOT_C2WS_PLUCKER_EMB_CACHE,
@@ -48,6 +52,63 @@ class LingBotWorldCausalDMDDenoisingStage(CausalDMDDenoisingStage):
Each call processes one chunk (num_frames_per_block frames).
"""
def _supports_qvg_kv_cache_quantization(self) -> bool:
return True
def _allocate_causal_kv_cache(
self,
*,
batch_size: int,
kv_cache_size: int,
num_attention_heads: int,
attention_head_dim: int,
dtype: torch.dtype,
device,
use_int_indices: bool = False,
sink_tokens: int = 0,
global_sink_tokens: int = 0,
attention_window_size: int | None = None,
allow_growth: bool = False,
) -> list[CausalKVCache]:
if not self._kv_quant_args.enabled:
return super()._allocate_causal_kv_cache(
batch_size=batch_size,
kv_cache_size=kv_cache_size,
num_attention_heads=num_attention_heads,
attention_head_dim=attention_head_dim,
dtype=dtype,
device=device,
use_int_indices=use_int_indices,
sink_tokens=sink_tokens,
global_sink_tokens=global_sink_tokens,
attention_window_size=attention_window_size,
allow_growth=allow_growth,
)
if global_sink_tokens or allow_growth:
raise NotImplementedError(
"QVG packed KV cache supports only the LingBot realtime "
"sliding-window and sink path"
)
if attention_window_size is None:
attention_window_size = kv_cache_size
return [
QVGPackedCausalKVCache(
batch_size=batch_size,
cache_size=kv_cache_size,
num_heads=num_attention_heads,
head_dim=attention_head_dim,
dtype=dtype,
device=device,
use_int_indices=use_int_indices,
global_end_index=torch.zeros(1, dtype=torch.long, device=device),
local_end_index=torch.zeros(1, dtype=torch.long, device=device),
sink_tokens=sink_tokens,
attention_window_size=attention_window_size,
quant_args=self._kv_quant_args,
)
for _ in range(self.num_transformer_blocks)
]
def _get_causal_kv_cache_size(
self,
*,
@@ -141,6 +202,7 @@ class LingBotWorldCausalDMDDenoisingStage(CausalDMDDenoisingStage):
) -> None:
self._reset_causal_cache_config_defaults()
super()._apply_causal_cache_overrides(batch, server_args)
self._kv_quant_args = server_args.kv_cache_quant_config
self._sync_interactive_kv_cache_window(server_args)
def _reset_causal_cache_config_defaults(self) -> None:
@@ -26,6 +26,7 @@ from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import (
is_ltx23_native_variant,
)
from sglang.multimodal_gen.configs.quantization.nunchaku import NunchakuSVDQuantArgs
from sglang.multimodal_gen.configs.quantization.qvg_kv import QVGKVQuantArgs
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
from sglang.multimodal_gen.runtime.layers.quantization.configs.nunchaku_config import (
NunchakuConfig,
@@ -358,6 +359,12 @@ class ServerArgs(DisaggServerArgsMixin):
default_factory=NunchakuSVDQuantArgs, repr=False
)
# KV-cache quantization (Quant-VideoGen PRQ). Off by default; mirrors the
# SRT --kv-cache-dtype pattern (typed config, not a pile of env vars).
kv_cache_quant_config: QVGKVQuantArgs = field(
default_factory=QVGKVQuantArgs, repr=False
)
# Master port for distributed inference
master_port: int = 30005
@@ -1814,6 +1821,67 @@ class ServerArgs(DisaggServerArgsMixin):
help="Disable autocast for denoising loop and vae decoding in pipeline sampling",
)
# KV-cache quantization (Quant-VideoGen PRQ)
parser.add_argument(
"--kv-cache-quant",
type=str,
default=None,
choices=["off", "int4", "int2"],
help="Enable Quant-VideoGen PRQ KV-cache quantization (off|int4|int2). "
"Defaults reproduce the tuned per-chunk config (stages=1, "
"centroids=128, block=64, symmetric, iters=2, recent=1, "
"per-chunk sink).",
)
parser.add_argument(
"--kv-cache-quant-centroids",
type=int,
default=None,
help="PRQ k-means centroids per stage (default 128).",
)
parser.add_argument(
"--kv-cache-quant-block-size",
type=int,
default=None,
help="PRQ residual scale block size (default 64).",
)
parser.add_argument(
"--kv-cache-quant-stages",
type=int,
default=None,
help="PRQ k-means stages (default 1).",
)
parser.add_argument(
"--kv-cache-quant-iters",
type=int,
default=None,
help="PRQ k-means iterations (default 2).",
)
parser.add_argument(
"--kv-cache-quant-asymmetric",
action="store_true",
default=None,
help="Use KIVI-style asymmetric residual quantization.",
)
parser.add_argument(
"--kv-cache-quant-keep-recent",
type=int,
default=None,
help="Completed chunks kept bf16 before quantizing (default 1).",
)
parser.add_argument(
"--kv-cache-quant-sink",
type=int,
default=None,
choices=[0, 1],
help="Quantize the attention sink too (1, default) " "or keep it bf16 (0).",
)
parser.add_argument(
"--kv-cache-quant-sink-keep",
type=int,
default=None,
help="Leading sink chunks kept bf16 forever (default 0).",
)
# quantization
parser.add_argument(
"--quantization",
@@ -2311,6 +2379,19 @@ class ServerArgs(DisaggServerArgsMixin):
elif attr == "nunchaku_config":
nunchaku_config = NunchakuSVDQuantArgs.from_dict(kwargs)
server_args_kwargs["nunchaku_config"] = nunchaku_config
elif attr == "kv_cache_quant_config":
kv_quant_config = kwargs.get("kv_cache_quant_config")
if kv_quant_config is None:
kv_quant_config = QVGKVQuantArgs.from_dict(kwargs)
elif isinstance(kv_quant_config, dict):
kv_quant_config = QVGKVQuantArgs(**kv_quant_config).validate()
elif isinstance(kv_quant_config, QVGKVQuantArgs):
kv_quant_config.validate()
else:
raise TypeError(
"kv_cache_quant_config must be QVGKVQuantArgs or a dict"
)
server_args_kwargs["kv_cache_quant_config"] = kv_quant_config
elif attr in kwargs:
server_args_kwargs[attr] = kwargs[attr]
@@ -7,6 +7,7 @@ import torch
from sglang.multimodal_gen.configs.pipeline_configs.lingbot_world import (
LingBotWorldCausalDMDConfig,
)
from sglang.multimodal_gen.configs.quantization.qvg_kv import QVGKVQuantArgs
from sglang.multimodal_gen.runtime.layers.kvcache.causal_attention_cache import (
CausalSelfAttentionKVCache,
CrossAttentionKVCache,
@@ -72,12 +73,16 @@ def test_lingbot_realtime_cache_config_overrides_checkpoint_defaults():
stage.sink_size = 9
stage.sliding_window_num_frames = 18
stage.num_token_per_frame = 10
stage.transformer = SimpleNamespace(
config=SimpleNamespace(arch_config=SimpleNamespace())
)
server_args = SimpleNamespace(
kv_cache_quant_config=QVGKVQuantArgs(),
pipeline_config=SimpleNamespace(
realtime_causal_sink_size=3,
realtime_causal_kv_cache_num_frames=45,
)
),
)
stage._apply_causal_cache_overrides(SimpleNamespace(), server_args)
@@ -93,16 +98,20 @@ def test_lingbot_realtime_cache_config_uses_request_overrides():
stage.sink_size = 9
stage.sliding_window_num_frames = 18
stage.num_token_per_frame = 10
stage.transformer = SimpleNamespace(
config=SimpleNamespace(arch_config=SimpleNamespace())
)
batch = SimpleNamespace(
realtime_causal_sink_size=4,
realtime_causal_kv_cache_num_frames=12,
)
server_args = SimpleNamespace(
kv_cache_quant_config=QVGKVQuantArgs(),
pipeline_config=SimpleNamespace(
realtime_causal_sink_size=3,
realtime_causal_kv_cache_num_frames=45,
)
),
)
stage._apply_causal_cache_overrides(batch, server_args)
@@ -120,6 +129,7 @@ def test_lingbot_realtime_attention_cache_rolls_with_sink_window():
stage.num_token_per_frame = 1
stage.num_frames_per_block = 3
stage.sliding_window_num_frames = 6
stage._kv_quant_args = QVGKVQuantArgs()
stage.transformer = SimpleNamespace(
num_attention_heads=1,
attention_head_dim=1,
@@ -279,12 +289,13 @@ def test_lingbot_interactive_kv_window_samples_base_moving_and_still(monkeypatch
stage.sliding_window_num_frames = 18
stage.transformer = SimpleNamespace(num_attention_heads=1)
server_args = SimpleNamespace(
kv_cache_quant_config=QVGKVQuantArgs(),
pipeline_config=SimpleNamespace(
interactive_kv_window_enable=True,
interactive_kv_moving_window=12,
interactive_kv_still_window=3,
interactive_kv_still_chunks=2,
)
),
)
cache_state = RealtimeCausalDiTState()
@@ -323,6 +334,7 @@ def test_lingbot_interactive_kv_window_none_disables_moving_window(monkeypatch):
stage.sliding_window_num_frames = 18
stage.transformer = SimpleNamespace(num_attention_heads=1)
server_args = SimpleNamespace(
kv_cache_quant_config=QVGKVQuantArgs(),
pipeline_config=SimpleNamespace(
realtime_causal_sink_size=9,
realtime_causal_kv_cache_num_frames=18,
@@ -330,7 +342,7 @@ def test_lingbot_interactive_kv_window_none_disables_moving_window(monkeypatch):
interactive_kv_moving_window=None,
interactive_kv_still_window=3,
interactive_kv_still_chunks=2,
)
),
)
cache_state = RealtimeCausalDiTState()
batch = SimpleNamespace(condition_inputs={"camera_actions": [["w"], [], []]})
@@ -353,6 +365,7 @@ def test_lingbot_interactive_kv_window_zero_is_valid_moving_window(monkeypatch):
stage.sliding_window_num_frames = 18
stage.transformer = SimpleNamespace(num_attention_heads=1)
server_args = SimpleNamespace(
kv_cache_quant_config=QVGKVQuantArgs(),
pipeline_config=SimpleNamespace(
realtime_causal_sink_size=9,
realtime_causal_kv_cache_num_frames=18,
@@ -360,7 +373,7 @@ def test_lingbot_interactive_kv_window_zero_is_valid_moving_window(monkeypatch):
interactive_kv_moving_window=0,
interactive_kv_still_window=None,
interactive_kv_still_chunks=2,
)
),
)
cache_state = RealtimeCausalDiTState()
batch = SimpleNamespace(condition_inputs={"camera_actions": [["w"], [], []]})
@@ -385,6 +398,7 @@ def test_lingbot_interactive_kv_window_updates_total_window_for_moving_default(
stage.sliding_window_num_frames = 18
stage.transformer = SimpleNamespace(num_attention_heads=1)
server_args = SimpleNamespace(
kv_cache_quant_config=QVGKVQuantArgs(),
pipeline_config=SimpleNamespace(
realtime_causal_sink_size=9,
realtime_causal_kv_cache_num_frames=18,
@@ -392,7 +406,7 @@ def test_lingbot_interactive_kv_window_updates_total_window_for_moving_default(
interactive_kv_moving_window=12,
interactive_kv_still_window=3,
interactive_kv_still_chunks=2,
)
),
)
batch = SimpleNamespace(condition_inputs={"camera_actions": [["w"], [], []]})
@@ -424,6 +438,7 @@ def test_lingbot_interactive_kv_window_resets_stage_window_between_requests(
),
)
dynamic_args = SimpleNamespace(
kv_cache_quant_config=QVGKVQuantArgs(),
pipeline_config=SimpleNamespace(
realtime_causal_sink_size=9,
realtime_causal_kv_cache_num_frames=18,
@@ -431,15 +446,16 @@ def test_lingbot_interactive_kv_window_resets_stage_window_between_requests(
interactive_kv_moving_window=12,
interactive_kv_still_window=3,
interactive_kv_still_chunks=2,
)
),
)
disabled_args = SimpleNamespace(
kv_cache_quant_config=QVGKVQuantArgs(),
pipeline_config=SimpleNamespace(
interactive_kv_window_enable=False,
interactive_kv_moving_window=12,
interactive_kv_still_window=3,
interactive_kv_still_chunks=2,
)
),
)
dynamic_batch = SimpleNamespace(condition_inputs={"camera_actions": [["w"]]})
@@ -469,6 +485,7 @@ def test_lingbot_interactive_kv_window_default_disabled(monkeypatch):
stage.sliding_window_num_frames = 18
stage.transformer = SimpleNamespace(num_attention_heads=1)
server_args = SimpleNamespace(
kv_cache_quant_config=QVGKVQuantArgs(),
pipeline_config=SimpleNamespace(
realtime_causal_sink_size=9,
realtime_causal_kv_cache_num_frames=18,
@@ -476,7 +493,7 @@ def test_lingbot_interactive_kv_window_default_disabled(monkeypatch):
interactive_kv_moving_window=12,
interactive_kv_still_window=3,
interactive_kv_still_chunks=2,
)
),
)
cache_state = RealtimeCausalDiTState()
batch = SimpleNamespace(condition_inputs={"camera_actions": [["w"], [], []]})
@@ -501,6 +518,7 @@ def test_lingbot_interactive_kv_window_env_can_enable_default(monkeypatch):
stage.sliding_window_num_frames = 18
stage.transformer = SimpleNamespace(num_attention_heads=1)
server_args = SimpleNamespace(
kv_cache_quant_config=QVGKVQuantArgs(),
pipeline_config=SimpleNamespace(
realtime_causal_sink_size=9,
realtime_causal_kv_cache_num_frames=18,
@@ -508,7 +526,7 @@ def test_lingbot_interactive_kv_window_env_can_enable_default(monkeypatch):
interactive_kv_moving_window=12,
interactive_kv_still_window=3,
interactive_kv_still_chunks=2,
)
),
)
cache_state = RealtimeCausalDiTState()
batch = SimpleNamespace(condition_inputs={"camera_actions": [["w"], [], []]})
@@ -531,6 +549,7 @@ def test_lingbot_interactive_kv_window_allocates_expected_cache_size():
stage.num_token_per_frame = 10
stage.num_frames_per_block = 3
stage.sliding_window_num_frames = 18
stage._kv_quant_args = QVGKVQuantArgs()
stage.transformer = SimpleNamespace(num_attention_heads=1, attention_head_dim=1)
policy = CausalDMDCachePolicy(
sequence_shard_enabled=False,
@@ -0,0 +1,237 @@
# SPDX-License-Identifier: Apache-2.0
"""Unit tests for Quant-VideoGen packed KV-cache (QVGPackedCausalKVCache).
- config parsing + defaults: CPU, no deps.
- packed-storage equivalence vs the dense cache (quant disabled -> bf16
segments): CPU-only, exercises the segment / sliding-window / eviction /
reconstruction logic bit-exactly without needing quant-videogen or a GPU.
- quant memory saving + reconstruction: requires CUDA + quant-videogen.
"""
import importlib.util
import unittest
from types import SimpleNamespace
import torch
from sglang.multimodal_gen.configs.quantization.qvg_kv import QVGKVQuantArgs
from sglang.multimodal_gen.runtime.layers.kvcache.causal_attention_cache import (
CausalSelfAttentionKVCache,
)
from sglang.multimodal_gen.runtime.layers.kvcache.qvg_packed_cache import (
QVGPackedCausalKVCache,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.causal_denoising import (
CausalDMDDenoisingStage,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.lingbot_world.lingbot_world_causal_denoising import (
LingBotWorldCausalDMDDenoisingStage,
)
_HAS_QVG = importlib.util.find_spec("quant_videogen") is not None
_HAS_CUDA = torch.cuda.is_available()
def _base(B, W, H, D, sink, dev):
return CausalSelfAttentionKVCache(
k=torch.zeros(B, W, H, D, device=dev),
v=torch.zeros(B, W, H, D, device=dev),
global_end_index=torch.zeros(1, dtype=torch.long, device=dev),
local_end_index=torch.zeros(1, dtype=torch.long, device=dev),
cache_size=W,
sink_tokens=sink,
attention_window_size=W,
)
def _packed(B, W, H, D, sink, dev, quant, use_int_indices=False):
return QVGPackedCausalKVCache(
batch_size=B,
cache_size=W,
num_heads=H,
head_dim=D,
dtype=torch.float32,
device=torch.device(dev),
use_int_indices=use_int_indices,
global_end_index=torch.zeros(1, dtype=torch.long, device=dev),
local_end_index=torch.zeros(1, dtype=torch.long, device=dev),
sink_tokens=sink,
attention_window_size=W,
quant_args=quant,
)
def _replay(cache, chunks, C, H, D, B, nsteps, rwt, data, dev):
views = []
for ci in range(chunks):
g = ci * C
for step in range(nsteps):
k, v = data[(ci, step)]
view = cache.update_and_get_attention_kv(
key=k, value=v, current_chunk_start=g, recent_window_tokens=rwt
)
views.append((view.k.clone(), view.v.clone()))
return views
class TestQVGKVQuantArgs(unittest.TestCase):
def test_off_by_default(self):
self.assertFalse(QVGKVQuantArgs().enabled)
self.assertFalse(QVGKVQuantArgs.from_dict({}).enabled)
def test_master_flag_defaults(self):
a = QVGKVQuantArgs.from_dict({"kv_cache_quant": "int4"})
self.assertTrue(a.enabled)
self.assertEqual(
(
a.bits,
a.stages,
a.centroids,
a.block_size,
a.kmeans_iters,
a.asymmetric,
a.keep_recent_chunks,
a.sink,
a.sink_keep_chunks,
),
(4, 1, 128, 64, 2, False, 1, True, 0),
)
def test_overrides_and_validation(self):
a = QVGKVQuantArgs.from_dict(
{
"kv_cache_quant": "int2",
"kv_cache_quant_stages": 3,
"kv_cache_quant_sink": 0,
}
)
self.assertEqual((a.bits, a.stages, a.sink), (2, 3, False))
with self.assertRaises(ValueError):
QVGKVQuantArgs.from_dict({"kv_cache_quant": "int8"})
class TestPackedStorageEquivalence(unittest.TestCase):
"""quant disabled -> packed cache stores bf16 segments; must match the
dense cache bit-for-bit across the sliding-window + eviction lifecycle."""
def _run(self, rwt):
B, H, D, C, sink, chunks, nsteps = 1, 2, 4, 8, 8, 8, 3
W = C * 5
dev = "cpu"
torch.manual_seed(0)
data = {
(ci, s): (torch.randn(B, C, H, D), torch.randn(B, C, H, D))
for ci in range(chunks)
for s in range(nsteps)
}
base = _base(B, W, H, D, sink, dev)
packed = _packed(B, W, H, D, sink, dev, QVGKVQuantArgs()) # disabled
ob = _replay(base, chunks, C, H, D, B, nsteps, rwt, data, dev)
op = _replay(packed, chunks, C, H, D, B, nsteps, rwt, data, dev)
for i, ((bk, bv), (pk, pv)) in enumerate(zip(ob, op)):
self.assertEqual(bk.shape, pk.shape, f"chunk {i} shape")
self.assertTrue(torch.equal(bk, pk), f"chunk {i} K mismatch")
self.assertTrue(torch.equal(bv, pv), f"chunk {i} V mismatch")
def test_full_window(self):
self._run(rwt=None)
def test_recent_window_selection(self):
self._run(rwt=16)
def test_host_index_cursors_reset_without_device_updates(self):
cache = _packed(
B=1,
W=16,
H=2,
D=4,
sink=0,
dev="cpu",
quant=QVGKVQuantArgs(),
use_int_indices=True,
)
cache.update_and_get_attention_kv(
key=torch.ones(1, 4, 2, 4),
value=torch.ones(1, 4, 2, 4),
current_chunk_start=0,
)
self.assertEqual(cache.global_end_index_int, 4)
self.assertEqual(cache.local_end_index_int, 4)
self.assertEqual(int(cache.global_end_index.item()), 0)
self.assertEqual(int(cache.local_end_index.item()), 0)
cache.reset_indices()
self.assertEqual(cache.global_end_index_int, 0)
self.assertEqual(cache.local_end_index_int, 0)
def test_quantization_is_rejected_by_unsupported_causal_stages(self):
stage = CausalDMDDenoisingStage.__new__(CausalDMDDenoisingStage)
stage.sink_size = 0
stage.sliding_window_num_frames = 1
with self.assertRaisesRegex(ValueError, "does not support QVG"):
stage._apply_causal_cache_overrides(
SimpleNamespace(),
SimpleNamespace(
pipeline_config=SimpleNamespace(),
kv_cache_quant_config=QVGKVQuantArgs(bits=4),
),
)
def test_quantization_is_available_for_lingbot(self):
stage = LingBotWorldCausalDMDDenoisingStage.__new__(
LingBotWorldCausalDMDDenoisingStage
)
stage._kv_quant_args = QVGKVQuantArgs(bits=4)
stage.num_transformer_blocks = 1
caches = stage._allocate_causal_kv_cache(
batch_size=1,
kv_cache_size=8,
num_attention_heads=1,
attention_head_dim=4,
dtype=torch.float32,
device=torch.device("cpu"),
)
self.assertIsInstance(caches[0], QVGPackedCausalKVCache)
@unittest.skipUnless(
_HAS_CUDA and _HAS_QVG,
"needs CUDA + quant-videogen",
)
class TestPackedQuantMemory(unittest.TestCase):
def test_saves_memory_and_reconstructs(self):
B, H, D, C, sink, chunks, nsteps = 1, 8, 128, 512, 512, 12, 2
W = C * 6
dev = "cuda"
q = QVGKVQuantArgs.from_dict({"kv_cache_quant": "int4"})
torch.manual_seed(1)
data = {
(ci, s): (
torch.randn(B, C, H, D, device=dev),
torch.randn(B, C, H, D, device=dev),
)
for ci in range(chunks)
for s in range(nsteps)
}
base = _base(B, W, H, D, sink, dev)
packed = _packed(B, W, H, D, sink, dev, q)
ob = _replay(base, chunks, C, H, D, B, nsteps, None, data, dev)
op = _replay(packed, chunks, C, H, D, B, nsteps, None, data, dev)
dense_bytes = base.k.numel() * base.k.element_size() * 2
self.assertLess(
packed.resident_nbytes(),
dense_bytes * 0.7,
"packed cache should save >30% resident memory",
)
bk, _ = ob[-1]
pk, _ = op[-1]
self.assertEqual(bk.shape, pk.shape)
mse = ((bk - pk) ** 2).mean().item()
self.assertLess(mse, 1.0, "int4 reconstruction error too large")
if __name__ == "__main__":
unittest.main()