[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:
co-authored by
Claude Opus 4.8
Mick
parent
24c84dfa68
commit
f64328c7f6
@@ -32,6 +32,16 @@ sglang serve \
|
|||||||
--enable-torch-compile false
|
--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
|
## 3. Realtime WebUI
|
||||||
|
|
||||||
The lightweight local WebUI is useful for validating latency, frame transport, and camera control behavior.
|
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
|
--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
|
## 3. Realtime WebUI
|
||||||
|
|
||||||
The lightweight local WebUI is useful for validating latency, frame transport, and camera control behavior.
|
The lightweight local WebUI is useful for validating latency, frame transport, and camera control behavior.
|
||||||
|
|||||||
@@ -1558,6 +1558,7 @@
|
|||||||
"pages": [
|
"pages": [
|
||||||
"docs/sglang-diffusion/api/cli",
|
"docs/sglang-diffusion/api/cli",
|
||||||
"docs/sglang-diffusion/api/openai_api",
|
"docs/sglang-diffusion/api/openai_api",
|
||||||
|
"docs/sglang-diffusion/realtime_models",
|
||||||
"docs/sglang-diffusion/models_with_ar",
|
"docs/sglang-diffusion/models_with_ar",
|
||||||
"docs/sglang-diffusion/models_with_pe",
|
"docs/sglang-diffusion/models_with_pe",
|
||||||
"docs/sglang-diffusion/api/post_processing"
|
"docs/sglang-diffusion/api/post_processing"
|
||||||
|
|||||||
@@ -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).
|
For frame interpolation and upscaling, see [Post-Processing](./post_processing).
|
||||||
|
|
||||||
### Quantized transformers
|
### Quantization
|
||||||
|
|
||||||
For quantized transformer checkpoints, prefer:
|
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` 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_`)
|
- `--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
|
### Request logging
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
- [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
|
- [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
|
- [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
|
- [Contributing](/docs/sglang-diffusion/contributing): contribution workflow, adding new models, and CI perf baselines
|
||||||
|
|
||||||
## Additional Documentation
|
## Additional Documentation
|
||||||
|
|||||||
@@ -2,7 +2,7 @@
|
|||||||
title: "Quantization"
|
title: "Quantization"
|
||||||
tag: "approx"
|
tag: "approx"
|
||||||
metatags:
|
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
|
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
|
- `--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`: 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_`)
|
- `--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:
|
Recommended example for pre-quantized checkpoints:
|
||||||
|
|
||||||
@@ -99,6 +100,14 @@ backend.
|
|||||||
<td>None</td>
|
<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>
|
<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>
|
||||||
|
<tr>
|
||||||
|
<td><code>qvg-kv</code></td>
|
||||||
|
<td>Unquantized model with runtime causal KV-cache compression</td>
|
||||||
|
<td><code>--kv-cache-quant {int4,int2}</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>
|
<tr>
|
||||||
<td><code>nunchaku-svdq</code></td>
|
<td><code>nunchaku-svdq</code></td>
|
||||||
<td>Pre-quantized Nunchaku transformer weights, usually named <code>svdq-{int4\|fp4}_r{rank}-...</code></td>
|
<td>Pre-quantized Nunchaku transformer weights, usually named <code>svdq-{int4\|fp4}_r{rank}-...</code></td>
|
||||||
@@ -118,6 +127,102 @@ backend.
|
|||||||
</tbody>
|
</tbody>
|
||||||
</table>
|
</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
|
||||||
|
|
||||||
Online quantization applies quantization to unquantized models at load time. This is useful for when pre-quantized checkpoints are not available.
|
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.
|
||||||
@@ -123,6 +123,14 @@ diffusion = [
|
|||||||
"xatlas",
|
"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 = [
|
||||||
"ray[default]>=2.55.1",
|
"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:
|
if self.attention_window_size == 0:
|
||||||
self.attention_window_size = self.cache_size
|
self.attention_window_size = self.cache_size
|
||||||
|
|
||||||
|
@property
|
||||||
|
def num_cache_heads(self) -> int:
|
||||||
|
return self.k.shape[2]
|
||||||
|
|
||||||
def reset_indices(self) -> None:
|
def reset_indices(self) -> None:
|
||||||
self.global_end_index.zero_()
|
self.global_end_index.zero_()
|
||||||
self.local_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,
|
CausalSelfAttentionKVCache,
|
||||||
CrossAttentionKVCache,
|
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.managers.forward_context import set_forward_context
|
||||||
from sglang.multimodal_gen.runtime.pipelines_core.diffusion_scheduler_utils import (
|
from sglang.multimodal_gen.runtime.pipelines_core.diffusion_scheduler_utils import (
|
||||||
get_or_create_request_scheduler,
|
get_or_create_request_scheduler,
|
||||||
@@ -46,6 +49,7 @@ logger = init_logger(__name__)
|
|||||||
CAUSAL_BLOCK_PROMPTS_KEY = "causal_block_prompts"
|
CAUSAL_BLOCK_PROMPTS_KEY = "causal_block_prompts"
|
||||||
CAUSAL_SCENE_CUT_MASK_KEY = "causal_scene_cut_mask"
|
CAUSAL_SCENE_CUT_MASK_KEY = "causal_scene_cut_mask"
|
||||||
CAUSAL_SHOT_INDICES_KEY = "causal_shot_indices"
|
CAUSAL_SHOT_INDICES_KEY = "causal_shot_indices"
|
||||||
|
CausalKVCache = CausalSelfAttentionKVCache | QVGPackedCausalKVCache
|
||||||
|
|
||||||
|
|
||||||
def expand_causal_block_prompts(
|
def expand_causal_block_prompts(
|
||||||
@@ -134,7 +138,7 @@ class CausalDMDCachePolicy:
|
|||||||
class CausalDMDRealtimeCacheContext:
|
class CausalDMDRealtimeCacheContext:
|
||||||
cache_state: RealtimeCausalDiTState
|
cache_state: RealtimeCausalDiTState
|
||||||
persist_state: bool
|
persist_state: bool
|
||||||
kv_cache: list[CausalSelfAttentionKVCache]
|
kv_cache: list[CausalKVCache]
|
||||||
crossattn_cache: list[CrossAttentionKVCache]
|
crossattn_cache: list[CrossAttentionKVCache]
|
||||||
current_start_frame: int
|
current_start_frame: int
|
||||||
chunk_idx: int
|
chunk_idx: int
|
||||||
@@ -415,6 +419,17 @@ class CausalDMDDenoisingStage(DenoisingStage):
|
|||||||
raise ValueError("realtime_causal_kv_cache_num_frames must be positive")
|
raise ValueError("realtime_causal_kv_cache_num_frames must be positive")
|
||||||
self.sliding_window_num_frames = int(kv_cache_num_frames)
|
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:
|
def _causal_sequence_shard_enabled(self, batch: Req) -> bool:
|
||||||
return False
|
return False
|
||||||
|
|
||||||
@@ -487,8 +502,8 @@ class CausalDMDDenoisingStage(DenoisingStage):
|
|||||||
or crossattn_cache is None
|
or crossattn_cache is None
|
||||||
or len(causal_kv_cache) != self.num_transformer_blocks
|
or len(causal_kv_cache) != self.num_transformer_blocks
|
||||||
or len(crossattn_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].cache_size != policy.expected_cache_tokens
|
||||||
or causal_kv_cache[0].k.shape[2] != policy.num_attention_heads
|
or causal_kv_cache[0].num_cache_heads != policy.num_attention_heads
|
||||||
or causal_kv_cache[0].sink_tokens != policy.expected_sink_tokens
|
or causal_kv_cache[0].sink_tokens != policy.expected_sink_tokens
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -1111,34 +1126,15 @@ class CausalDMDDenoisingStage(DenoisingStage):
|
|||||||
global_sink_tokens: int = 0,
|
global_sink_tokens: int = 0,
|
||||||
attention_window_size: int | None = None,
|
attention_window_size: int | None = None,
|
||||||
allow_growth: bool = False,
|
allow_growth: bool = False,
|
||||||
) -> list[CausalSelfAttentionKVCache]:
|
) -> list[CausalKVCache]:
|
||||||
causal_kv_cache = []
|
|
||||||
int_index = 0 if use_int_indices else None
|
|
||||||
if attention_window_size is None:
|
if attention_window_size is None:
|
||||||
attention_window_size = kv_cache_size
|
attention_window_size = kv_cache_size
|
||||||
for _ in range(self.num_transformer_blocks):
|
int_index = 0 if use_int_indices else None
|
||||||
causal_kv_cache.append(
|
shape = [batch_size, kv_cache_size, num_attention_heads, attention_head_dim]
|
||||||
|
return [
|
||||||
CausalSelfAttentionKVCache(
|
CausalSelfAttentionKVCache(
|
||||||
k=torch.zeros(
|
k=torch.zeros(shape, dtype=dtype, device=device),
|
||||||
[
|
v=torch.zeros(shape, dtype=dtype, device=device),
|
||||||
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),
|
global_end_index=torch.zeros(1, dtype=torch.long, device=device),
|
||||||
local_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,
|
global_end_index_int=int_index,
|
||||||
@@ -1149,8 +1145,8 @@ class CausalDMDDenoisingStage(DenoisingStage):
|
|||||||
attention_window_size=attention_window_size,
|
attention_window_size=attention_window_size,
|
||||||
allow_growth=allow_growth,
|
allow_growth=allow_growth,
|
||||||
)
|
)
|
||||||
)
|
for _ in range(self.num_transformer_blocks)
|
||||||
return causal_kv_cache
|
]
|
||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def forward(
|
def forward(
|
||||||
|
|||||||
+62
@@ -12,6 +12,9 @@ from sglang.multimodal_gen.runtime.distributed.parallel_state import (
|
|||||||
get_ring_parallel_world_size,
|
get_ring_parallel_world_size,
|
||||||
get_ulysses_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.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.schedule_batch import Req
|
||||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.causal_denoising import (
|
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,
|
CausalDMDDenoisingStage,
|
||||||
CausalDMDForwardContext,
|
CausalDMDForwardContext,
|
||||||
CausalDMDRealtimeCacheContext,
|
CausalDMDRealtimeCacheContext,
|
||||||
|
CausalKVCache,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.lingbot_world.constants import (
|
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.lingbot_world.constants import (
|
||||||
LINGBOT_C2WS_PLUCKER_EMB_CACHE,
|
LINGBOT_C2WS_PLUCKER_EMB_CACHE,
|
||||||
@@ -48,6 +52,63 @@ class LingBotWorldCausalDMDDenoisingStage(CausalDMDDenoisingStage):
|
|||||||
Each call processes one chunk (num_frames_per_block frames).
|
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(
|
def _get_causal_kv_cache_size(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
@@ -141,6 +202,7 @@ class LingBotWorldCausalDMDDenoisingStage(CausalDMDDenoisingStage):
|
|||||||
) -> None:
|
) -> None:
|
||||||
self._reset_causal_cache_config_defaults()
|
self._reset_causal_cache_config_defaults()
|
||||||
super()._apply_causal_cache_overrides(batch, server_args)
|
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)
|
self._sync_interactive_kv_cache_window(server_args)
|
||||||
|
|
||||||
def _reset_causal_cache_config_defaults(self) -> None:
|
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,
|
is_ltx23_native_variant,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.configs.quantization.nunchaku import NunchakuSVDQuantArgs
|
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.disaggregation.roles import RoleType
|
||||||
from sglang.multimodal_gen.runtime.layers.quantization.configs.nunchaku_config import (
|
from sglang.multimodal_gen.runtime.layers.quantization.configs.nunchaku_config import (
|
||||||
NunchakuConfig,
|
NunchakuConfig,
|
||||||
@@ -358,6 +359,12 @@ class ServerArgs(DisaggServerArgsMixin):
|
|||||||
default_factory=NunchakuSVDQuantArgs, repr=False
|
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 for distributed inference
|
||||||
master_port: int = 30005
|
master_port: int = 30005
|
||||||
|
|
||||||
@@ -1814,6 +1821,67 @@ class ServerArgs(DisaggServerArgsMixin):
|
|||||||
help="Disable autocast for denoising loop and vae decoding in pipeline sampling",
|
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
|
# quantization
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--quantization",
|
"--quantization",
|
||||||
@@ -2311,6 +2379,19 @@ class ServerArgs(DisaggServerArgsMixin):
|
|||||||
elif attr == "nunchaku_config":
|
elif attr == "nunchaku_config":
|
||||||
nunchaku_config = NunchakuSVDQuantArgs.from_dict(kwargs)
|
nunchaku_config = NunchakuSVDQuantArgs.from_dict(kwargs)
|
||||||
server_args_kwargs["nunchaku_config"] = nunchaku_config
|
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:
|
elif attr in kwargs:
|
||||||
server_args_kwargs[attr] = kwargs[attr]
|
server_args_kwargs[attr] = kwargs[attr]
|
||||||
|
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ import torch
|
|||||||
from sglang.multimodal_gen.configs.pipeline_configs.lingbot_world import (
|
from sglang.multimodal_gen.configs.pipeline_configs.lingbot_world import (
|
||||||
LingBotWorldCausalDMDConfig,
|
LingBotWorldCausalDMDConfig,
|
||||||
)
|
)
|
||||||
|
from sglang.multimodal_gen.configs.quantization.qvg_kv import QVGKVQuantArgs
|
||||||
from sglang.multimodal_gen.runtime.layers.kvcache.causal_attention_cache import (
|
from sglang.multimodal_gen.runtime.layers.kvcache.causal_attention_cache import (
|
||||||
CausalSelfAttentionKVCache,
|
CausalSelfAttentionKVCache,
|
||||||
CrossAttentionKVCache,
|
CrossAttentionKVCache,
|
||||||
@@ -72,12 +73,16 @@ def test_lingbot_realtime_cache_config_overrides_checkpoint_defaults():
|
|||||||
stage.sink_size = 9
|
stage.sink_size = 9
|
||||||
stage.sliding_window_num_frames = 18
|
stage.sliding_window_num_frames = 18
|
||||||
stage.num_token_per_frame = 10
|
stage.num_token_per_frame = 10
|
||||||
|
stage.transformer = SimpleNamespace(
|
||||||
|
config=SimpleNamespace(arch_config=SimpleNamespace())
|
||||||
|
)
|
||||||
|
|
||||||
server_args = SimpleNamespace(
|
server_args = SimpleNamespace(
|
||||||
|
kv_cache_quant_config=QVGKVQuantArgs(),
|
||||||
pipeline_config=SimpleNamespace(
|
pipeline_config=SimpleNamespace(
|
||||||
realtime_causal_sink_size=3,
|
realtime_causal_sink_size=3,
|
||||||
realtime_causal_kv_cache_num_frames=45,
|
realtime_causal_kv_cache_num_frames=45,
|
||||||
)
|
),
|
||||||
)
|
)
|
||||||
stage._apply_causal_cache_overrides(SimpleNamespace(), server_args)
|
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.sink_size = 9
|
||||||
stage.sliding_window_num_frames = 18
|
stage.sliding_window_num_frames = 18
|
||||||
stage.num_token_per_frame = 10
|
stage.num_token_per_frame = 10
|
||||||
|
stage.transformer = SimpleNamespace(
|
||||||
|
config=SimpleNamespace(arch_config=SimpleNamespace())
|
||||||
|
)
|
||||||
|
|
||||||
batch = SimpleNamespace(
|
batch = SimpleNamespace(
|
||||||
realtime_causal_sink_size=4,
|
realtime_causal_sink_size=4,
|
||||||
realtime_causal_kv_cache_num_frames=12,
|
realtime_causal_kv_cache_num_frames=12,
|
||||||
)
|
)
|
||||||
server_args = SimpleNamespace(
|
server_args = SimpleNamespace(
|
||||||
|
kv_cache_quant_config=QVGKVQuantArgs(),
|
||||||
pipeline_config=SimpleNamespace(
|
pipeline_config=SimpleNamespace(
|
||||||
realtime_causal_sink_size=3,
|
realtime_causal_sink_size=3,
|
||||||
realtime_causal_kv_cache_num_frames=45,
|
realtime_causal_kv_cache_num_frames=45,
|
||||||
)
|
),
|
||||||
)
|
)
|
||||||
stage._apply_causal_cache_overrides(batch, server_args)
|
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_token_per_frame = 1
|
||||||
stage.num_frames_per_block = 3
|
stage.num_frames_per_block = 3
|
||||||
stage.sliding_window_num_frames = 6
|
stage.sliding_window_num_frames = 6
|
||||||
|
stage._kv_quant_args = QVGKVQuantArgs()
|
||||||
stage.transformer = SimpleNamespace(
|
stage.transformer = SimpleNamespace(
|
||||||
num_attention_heads=1,
|
num_attention_heads=1,
|
||||||
attention_head_dim=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.sliding_window_num_frames = 18
|
||||||
stage.transformer = SimpleNamespace(num_attention_heads=1)
|
stage.transformer = SimpleNamespace(num_attention_heads=1)
|
||||||
server_args = SimpleNamespace(
|
server_args = SimpleNamespace(
|
||||||
|
kv_cache_quant_config=QVGKVQuantArgs(),
|
||||||
pipeline_config=SimpleNamespace(
|
pipeline_config=SimpleNamespace(
|
||||||
interactive_kv_window_enable=True,
|
interactive_kv_window_enable=True,
|
||||||
interactive_kv_moving_window=12,
|
interactive_kv_moving_window=12,
|
||||||
interactive_kv_still_window=3,
|
interactive_kv_still_window=3,
|
||||||
interactive_kv_still_chunks=2,
|
interactive_kv_still_chunks=2,
|
||||||
)
|
),
|
||||||
)
|
)
|
||||||
cache_state = RealtimeCausalDiTState()
|
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.sliding_window_num_frames = 18
|
||||||
stage.transformer = SimpleNamespace(num_attention_heads=1)
|
stage.transformer = SimpleNamespace(num_attention_heads=1)
|
||||||
server_args = SimpleNamespace(
|
server_args = SimpleNamespace(
|
||||||
|
kv_cache_quant_config=QVGKVQuantArgs(),
|
||||||
pipeline_config=SimpleNamespace(
|
pipeline_config=SimpleNamespace(
|
||||||
realtime_causal_sink_size=9,
|
realtime_causal_sink_size=9,
|
||||||
realtime_causal_kv_cache_num_frames=18,
|
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_moving_window=None,
|
||||||
interactive_kv_still_window=3,
|
interactive_kv_still_window=3,
|
||||||
interactive_kv_still_chunks=2,
|
interactive_kv_still_chunks=2,
|
||||||
)
|
),
|
||||||
)
|
)
|
||||||
cache_state = RealtimeCausalDiTState()
|
cache_state = RealtimeCausalDiTState()
|
||||||
batch = SimpleNamespace(condition_inputs={"camera_actions": [["w"], [], []]})
|
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.sliding_window_num_frames = 18
|
||||||
stage.transformer = SimpleNamespace(num_attention_heads=1)
|
stage.transformer = SimpleNamespace(num_attention_heads=1)
|
||||||
server_args = SimpleNamespace(
|
server_args = SimpleNamespace(
|
||||||
|
kv_cache_quant_config=QVGKVQuantArgs(),
|
||||||
pipeline_config=SimpleNamespace(
|
pipeline_config=SimpleNamespace(
|
||||||
realtime_causal_sink_size=9,
|
realtime_causal_sink_size=9,
|
||||||
realtime_causal_kv_cache_num_frames=18,
|
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_moving_window=0,
|
||||||
interactive_kv_still_window=None,
|
interactive_kv_still_window=None,
|
||||||
interactive_kv_still_chunks=2,
|
interactive_kv_still_chunks=2,
|
||||||
)
|
),
|
||||||
)
|
)
|
||||||
cache_state = RealtimeCausalDiTState()
|
cache_state = RealtimeCausalDiTState()
|
||||||
batch = SimpleNamespace(condition_inputs={"camera_actions": [["w"], [], []]})
|
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.sliding_window_num_frames = 18
|
||||||
stage.transformer = SimpleNamespace(num_attention_heads=1)
|
stage.transformer = SimpleNamespace(num_attention_heads=1)
|
||||||
server_args = SimpleNamespace(
|
server_args = SimpleNamespace(
|
||||||
|
kv_cache_quant_config=QVGKVQuantArgs(),
|
||||||
pipeline_config=SimpleNamespace(
|
pipeline_config=SimpleNamespace(
|
||||||
realtime_causal_sink_size=9,
|
realtime_causal_sink_size=9,
|
||||||
realtime_causal_kv_cache_num_frames=18,
|
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_moving_window=12,
|
||||||
interactive_kv_still_window=3,
|
interactive_kv_still_window=3,
|
||||||
interactive_kv_still_chunks=2,
|
interactive_kv_still_chunks=2,
|
||||||
)
|
),
|
||||||
)
|
)
|
||||||
batch = SimpleNamespace(condition_inputs={"camera_actions": [["w"], [], []]})
|
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(
|
dynamic_args = SimpleNamespace(
|
||||||
|
kv_cache_quant_config=QVGKVQuantArgs(),
|
||||||
pipeline_config=SimpleNamespace(
|
pipeline_config=SimpleNamespace(
|
||||||
realtime_causal_sink_size=9,
|
realtime_causal_sink_size=9,
|
||||||
realtime_causal_kv_cache_num_frames=18,
|
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_moving_window=12,
|
||||||
interactive_kv_still_window=3,
|
interactive_kv_still_window=3,
|
||||||
interactive_kv_still_chunks=2,
|
interactive_kv_still_chunks=2,
|
||||||
)
|
),
|
||||||
)
|
)
|
||||||
disabled_args = SimpleNamespace(
|
disabled_args = SimpleNamespace(
|
||||||
|
kv_cache_quant_config=QVGKVQuantArgs(),
|
||||||
pipeline_config=SimpleNamespace(
|
pipeline_config=SimpleNamespace(
|
||||||
interactive_kv_window_enable=False,
|
interactive_kv_window_enable=False,
|
||||||
interactive_kv_moving_window=12,
|
interactive_kv_moving_window=12,
|
||||||
interactive_kv_still_window=3,
|
interactive_kv_still_window=3,
|
||||||
interactive_kv_still_chunks=2,
|
interactive_kv_still_chunks=2,
|
||||||
)
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
dynamic_batch = SimpleNamespace(condition_inputs={"camera_actions": [["w"]]})
|
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.sliding_window_num_frames = 18
|
||||||
stage.transformer = SimpleNamespace(num_attention_heads=1)
|
stage.transformer = SimpleNamespace(num_attention_heads=1)
|
||||||
server_args = SimpleNamespace(
|
server_args = SimpleNamespace(
|
||||||
|
kv_cache_quant_config=QVGKVQuantArgs(),
|
||||||
pipeline_config=SimpleNamespace(
|
pipeline_config=SimpleNamespace(
|
||||||
realtime_causal_sink_size=9,
|
realtime_causal_sink_size=9,
|
||||||
realtime_causal_kv_cache_num_frames=18,
|
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_moving_window=12,
|
||||||
interactive_kv_still_window=3,
|
interactive_kv_still_window=3,
|
||||||
interactive_kv_still_chunks=2,
|
interactive_kv_still_chunks=2,
|
||||||
)
|
),
|
||||||
)
|
)
|
||||||
cache_state = RealtimeCausalDiTState()
|
cache_state = RealtimeCausalDiTState()
|
||||||
batch = SimpleNamespace(condition_inputs={"camera_actions": [["w"], [], []]})
|
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.sliding_window_num_frames = 18
|
||||||
stage.transformer = SimpleNamespace(num_attention_heads=1)
|
stage.transformer = SimpleNamespace(num_attention_heads=1)
|
||||||
server_args = SimpleNamespace(
|
server_args = SimpleNamespace(
|
||||||
|
kv_cache_quant_config=QVGKVQuantArgs(),
|
||||||
pipeline_config=SimpleNamespace(
|
pipeline_config=SimpleNamespace(
|
||||||
realtime_causal_sink_size=9,
|
realtime_causal_sink_size=9,
|
||||||
realtime_causal_kv_cache_num_frames=18,
|
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_moving_window=12,
|
||||||
interactive_kv_still_window=3,
|
interactive_kv_still_window=3,
|
||||||
interactive_kv_still_chunks=2,
|
interactive_kv_still_chunks=2,
|
||||||
)
|
),
|
||||||
)
|
)
|
||||||
cache_state = RealtimeCausalDiTState()
|
cache_state = RealtimeCausalDiTState()
|
||||||
batch = SimpleNamespace(condition_inputs={"camera_actions": [["w"], [], []]})
|
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_token_per_frame = 10
|
||||||
stage.num_frames_per_block = 3
|
stage.num_frames_per_block = 3
|
||||||
stage.sliding_window_num_frames = 18
|
stage.sliding_window_num_frames = 18
|
||||||
|
stage._kv_quant_args = QVGKVQuantArgs()
|
||||||
stage.transformer = SimpleNamespace(num_attention_heads=1, attention_head_dim=1)
|
stage.transformer = SimpleNamespace(num_attention_heads=1, attention_head_dim=1)
|
||||||
policy = CausalDMDCachePolicy(
|
policy = CausalDMDCachePolicy(
|
||||||
sequence_shard_enabled=False,
|
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()
|
||||||
Reference in New Issue
Block a user