[diffusion] chore: update Cache-DiT to 1.5.1 for DMD Calibrator, SVDQuant DQ, etc (#40104)

Co-authored-by: copilot-swe-agent[bot] <198982749+Copilot@users.noreply.github.com>
This commit is contained in:
DefTruth
2026-09-19 11:14:33 +08:00
committed by GitHub
co-authored by copilot-swe-agent[bot]
parent 090263eff6
commit f1fbbd17bb
11 changed files with 1485 additions and 27 deletions
+1
View File
@@ -174,6 +174,7 @@ benchmark/llava_bench/mme_pack
!tools/sglang-simulator/examples/replay/trace.jsonl
tmp*.txt
/tmp/
.tmp/
# Torch Compile logs
tl_out/
+219 -5
View File
@@ -11,8 +11,12 @@ SGLang integrates [Cache-DiT](https://github.com/vipshop/cache-dit), a caching a
- **DBCache (Dual Block Cache)**: Dynamically decides when to cache transformer blocks based on residual differences
- **TaylorSeer**: Uses Taylor expansion for calibration to optimize caching decisions
- **DMD Calibrator**: An **exponential-basis** forecasting calibrator (Dynamic Mode Decomposition, not Distribution Matching Distillation) that serves as a drop-in alternative to TaylorSeer's polynomial basis; strongest on flow-matching models
- **SCM (Step Computation Masking)**: Step-level caching control for additional speedup
Cache-DiT also ships **SVDQuant** W4A4 (int4 / NVFP4) dynamic quantization, which can be combined
with DBCache caching (see [Quantization](#quantization)).
## Basic Usage
Cache-DiT is a **per-request** switch: each request decides whether to run
@@ -47,8 +51,11 @@ client.images.generate(
`SGLANG_CACHE_DIT_ENABLED` server default). `cache_dit_params` accepts the
DBCache knobs (`Fn_compute_blocks`, `Bn_compute_blocks`, `max_warmup_steps`,
`residual_diff_threshold`, `max_continuous_cached_steps`, `enable_taylorseer`,
`taylorseer_order`), the SCM knobs (`scm_preset`, `scm_compute_bins`,
`scm_cache_bins`, `scm_policy`), and a nested `secondary` dict with the DBCache
`taylorseer_order`), the DMD knobs (`enable_dmd`, `dmd_history`, `dmd_rank`,
`dmd_ridge`, `dmd_svd_precision`; DMD and TaylorSeer are mutually exclusive
calibrators and cannot be enabled together), the SCM knobs (`scm_preset`,
`scm_compute_bins`, `scm_cache_bins`, `scm_policy`), and a nested `secondary`
dict with the DBCache
knobs for the second transformer of dual-DiT models (unset secondary keys
inherit the request's primary values, then the
`SGLANG_CACHE_DIT_SECONDARY_*` defaults).
@@ -134,6 +141,49 @@ cache_config:
enable_sperate_cfg: true # e.g, Qwen-Image, Wan, Chroma, Ovis-Image, etc.
```
- DBCache + DMD Calibrator
Instead of TaylorSeer, you can use the DMD calibrator: an **exponential-basis** forecasting
calibrator that serves as a drop-in alternative to TaylorSeer's polynomial basis. DMD models
the cached feature stream as a linear dynamical system (`Y_{t+1} ~= A @ Y_t`), forecasts
cached features from the fitted eigen-modes, and stays accurate over longer cache skips where
polynomial extrapolation diverges. DMD here refers to Dynamic Mode Decomposition (Schmid
2010), **not** Distribution Matching Distillation. DMD works best on flow-matching models
(e.g., FLUX), while TaylorSeer is often better on DDPM-style models — try both. DMD and
TaylorSeer are mutually exclusive — enable only one of `enable_dmd` / `enable_taylorseer`:
```yaml Config
cache_config:
max_warmup_steps: 8
warmup_interval: 2
max_cached_steps: -1
max_continuous_cached_steps: 2
Fn_compute_blocks: 1
Bn_compute_blocks: 0 # Bn=0 since the DMD calibrator replaces the Bn calibrator
residual_diff_threshold: 0.12
enable_dmd: true
dmd_history: 6 # snapshot window length, 5-6 typical
dmd_svd_precision: "medium" # "low", "medium" or "high"
```
A `dmd_history` window of 56 snapshots is typically the sweet spot — longer histories do not
always help, because the feature dynamics drift across timesteps. With fewer than 4 uniformly
spaced snapshots available, DMD transparently falls back to the Taylor expansion it maintains
internally. See the
[Cache-DiT DMD documentation](https://cache-dit.readthedocs.io/en/latest/user_guide/CACHE_API/#dmd-calibrator-dynamic-mode-decomposition)
for the mathematical principle and quantitative comparisons. A ready-made config is available
at
[examples/configs/cache_dmd.yaml](https://github.com/vipshop/cache-dit/blob/main/examples/configs/cache_dmd.yaml)
in the Cache-DiT repository. Apply it with the same `--cache-dit-config` flag:
```bash
sglang generate \
--backend diffusers \
--model-path Qwen/Qwen-Image \
--cache-dit-config cache_dmd.yaml \
--prompt "A beautiful sunset over the mountains"
```
### Distributed inference
- 1D Parallelism
@@ -300,6 +350,81 @@ sglang generate \
--prompt "A beautiful sunset over the mountains"
```
#### SVDQuant (W4A4 int4 / NVFP4)
SVDQuant is Cache-DiT's built-in W4A4 PTQ quantization (weights and activations in int4 or
NVFP4, with smoothed low-rank branches). It can be freely combined with DBCache caching and
the DMD calibrator for the largest speedups.
::::note
SVDQuant requires a cache-dit build **with CUDA extension support** — a plain
`pip install cache-dit` does NOT include it. Install one of:
```bash Command
# Option 1: prebuilt CUDA 13 wheel
pip install cache-dit-cu13==<version> --no-deps
# Option 2: build from source with SVDQuant enabled
git clone https://github.com/vipshop/cache-dit
cd cache-dit
export CUDA_HOME=/usr/local/cuda
CACHE_DIT_BUILD_SVDQUANT=1 pip install ".[quantization]" --no-build-isolation
```
::::
Valid `quant_type` values are `svdq_int4_r{32,64,128,256}_dq` (int4 W4A4) and
`svdq_nvfp4_r{32,64,128,256}_dq` (NVFP4 W4A4; requires a Blackwell GPU). Example config
combining SVDQuant NVFP4 with DBCache + DMD (see
[examples/configs/blackwell/cache_dmd_svdq.yaml](https://github.com/vipshop/cache-dit/blob/main/examples/configs/blackwell/cache_dmd_svdq.yaml)):
```yaml Config
cache_config:
max_warmup_steps: 8
warmup_interval: 2
max_cached_steps: -1
max_continuous_cached_steps: 2
Fn_compute_blocks: 1
Bn_compute_blocks: 0
residual_diff_threshold: 0.12
enable_dmd: true
dmd_history: 6
dmd_svd_precision: "medium"
quantize_config:
quant_type: "svdq_nvfp4_r128_dq" # nvfp4 for Blackwell; use svdq_int4_r128_dq for int4
svdq_kwargs:
quantize_device: "cuda"
fused_mlp: true
exclude_layers:
- "embedder"
- "embed"
verbose: false
```
For int4 W4A4 (pre-Blackwell GPUs), the same config with
`quant_type: "svdq_int4_r128_dq"` is available at
[examples/configs/cache_dmd_svdq.yaml](https://github.com/vipshop/cache-dit/blob/main/examples/configs/cache_dmd_svdq.yaml)
(add `runtime_kernel: "v2"` to `svdq_kwargs`).
Enable `torch.compile` for the best SVDQuant performance, and make sure `--warmup-steps`
covers the compile warmup (use the same value as `--num-inference-steps`):
```bash Command
sglang generate \
--backend diffusers \
--model-path black-forest-labs/FLUX.1-dev \
--num-inference-steps=28 \
--warmup-mode request \
--warmup-steps 28 \
--cache-dit-config cache_dmd_svdq.yaml \
--enable-torch-compile \
--dit-cpu-offload false \
--text-encoder-cpu-offload false \
--prompt "A beautiful sunset over the mountains"
```
You can verify from the log that the quantization is active:
`[Cache-DiT] SVDQuant Type: svdq_nvfp4_r128_dq, Rank: 128`.
### Combined Configs: Cache + Parallelism + Quantization
You can also combine all the above configs together in a single yaml file `combined.yaml` that contains:
@@ -418,12 +543,89 @@ TaylorSeer improves caching accuracy using Taylor expansion:
</tbody>
</table>
### DMD Calibrator Configuration
DMD (Dynamic Mode Decomposition, Schmid 2010 — **not** Distribution Matching Distillation) is
an **exponential-basis** forecasting calibrator and a drop-in alternative to TaylorSeer's
polynomial basis. At each full-compute step it records a snapshot of the computed features; at
a cached step it identifies a linear propagator from the recent snapshot window (one economy
SVD with rank truncation, then eigendecomposition) and forecasts the current features via
eigenvalue powers — cheap to advance, and stable over longer cache skips where polynomial
extrapolation diverges. It typically improves both speed and quality over pure DBCache.
**DMD and TaylorSeer are mutually exclusive** (enabling both raises a `ValueError`); DMD is
best for flow-matching models, TaylorSeer for DDPM-style ones. See the
[Cache-DiT DMD documentation](https://cache-dit.readthedocs.io/en/latest/user_guide/CACHE_API/#dmd-calibrator-dynamic-mode-decomposition)
for details:
<table style={{width: "100%", borderCollapse: "collapse", tableLayout: "fixed"}}>
<colgroup>
<col style={{width: "14%"}} />
<col style={{width: "38%"}} />
<col style={{width: "14%"}} />
<col style={{width: "34%"}} />
</colgroup>
<thead>
<tr style={{borderBottom: "2px solid #d55816"}}>
<th style={{textAlign: "left", padding: "10px 12px", fontWeight: 700, whiteSpace: "nowrap", backgroundColor: "rgba(255,255,255,0.02)"}}>Parameter</th>
<th style={{textAlign: "left", padding: "10px 12px", fontWeight: 700, whiteSpace: "nowrap", backgroundColor: "rgba(255,255,255,0.05)"}}>Env Variable</th>
<th style={{textAlign: "left", padding: "10px 12px", fontWeight: 700, whiteSpace: "nowrap", backgroundColor: "rgba(255,255,255,0.02)"}}>Default</th>
<th style={{textAlign: "left", padding: "10px 12px", fontWeight: 700, whiteSpace: "nowrap", backgroundColor: "rgba(255,255,255,0.05)"}}>Description</th>
</tr>
</thead>
<tbody>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>Enable</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>`SGLANG_CACHE_DIT_DMD`</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>false</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Enable the DMD calibrator</td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>History</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>`SGLANG_CACHE_DIT_DMD_HISTORY`</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>6</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Snapshot window length; 5-6 typical. Needs >= 4 uniformly spaced snapshots, otherwise DMD falls back to TaylorSeer</td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>Rank</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>`SGLANG_CACHE_DIT_DMD_RANK`</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>0</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>SVD truncation rank; 0 = automatic (drop modes below 1e-4 of the leading singular value)</td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>Ridge</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>`SGLANG_CACHE_DIT_DMD_RIDGE`</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>1e-8</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Tikhonov regularization added to the inverted singular values</td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>SVD Precision</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>`SGLANG_CACHE_DIT_DMD_SVD_PRECISION`</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>medium</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>SVD precision: "low", "medium" or "high"</td>
</tr>
</tbody>
</table>
Usage (SGLD backend, env-driven):
```bash Command
SGLANG_CACHE_DIT_ENABLED=true \
SGLANG_CACHE_DIT_DMD=true \
sglang generate --model-path black-forest-labs/FLUX.1-dev \
--prompt "A curious raccoon in a forest"
```
On the diffusers backend, enable DMD from the yaml config instead
(`enable_dmd: true` in `cache_config`, see
[Diffusers Backend](#diffusers-backend)); DMD can also be set per request via
`cache_dit_params: {"enable_dmd": true}`.
### Combined Configuration Example
DBCache and TaylorSeer are complementary strategies that work together, you can configure both sets of parameters
simultaneously:
```bash
```bash Command
SGLANG_CACHE_DIT_ENABLED=true \
SGLANG_CACHE_DIT_FN=2 \
SGLANG_CACHE_DIT_BN=1 \
@@ -496,7 +698,7 @@ SCM is configured with presets:
**Usage**
```bash
```bash Command
SGLANG_CACHE_DIT_ENABLED=true \
SGLANG_CACHE_DIT_SCM_PRESET=medium \
sglang generate --model-path Qwen/Qwen-Image \
@@ -507,7 +709,7 @@ sglang generate --model-path Qwen/Qwen-Image \
For fine-grained control over which steps to compute vs cache:
```bash
```bash Command
SGLANG_CACHE_DIT_ENABLED=true \
SGLANG_CACHE_DIT_SCM_COMPUTE_BINS="8,3,3,2,2" \
SGLANG_CACHE_DIT_SCM_CACHE_BINS="1,2,2,2,3" \
@@ -617,6 +819,18 @@ SGLang Diffusion x Cache-DiT supports almost all models originally supported in
For models with < 8 inference steps (e.g., DMD distilled models), SCM will be automatically disabled. DBCache
acceleration still works.
### SVDQuant unavailable or load failure
SVDQuant cases raise `svdq_is_available() = False` or
`undefined symbol: ... materialize_cow_storage ...` when the installed cache-dit has no CUDA
extension, or the prebuilt wheel was compiled against an incompatible torch. Fix: reinstall
from the `cache-dit-cu13` wheel matching your torch version, or build cache-dit from source
with `CACHE_DIT_BUILD_SVDQUANT=1` (see [Quantization](#quantization)). Quick self-check:
```bash Command
python -c "from cache_dit.quantization.svdquant import svdq_is_available, svdq_get_load_error as e; print(svdq_is_available(), e())"
```
## References
- [Cache-DiT](https://github.com/vipshop/cache-dit)
@@ -268,6 +268,31 @@ See [cache-dit documentation](./cache_dit) for detailed configuration.
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>1</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>TaylorSeer order (1 or 2)</td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>`SGLANG_CACHE_DIT_DMD`</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>false</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>Enable the DMD (Dynamic Mode Decomposition) calibrator (mutually exclusive with TaylorSeer)</td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>`SGLANG_CACHE_DIT_DMD_HISTORY`</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>6</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>DMD snapshot window length (5-6 typical; needs >= 4 uniformly spaced snapshots, otherwise DMD falls back to TaylorSeer)</td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>`SGLANG_CACHE_DIT_DMD_RANK`</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>0</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>DMD SVD truncation rank (0 = automatic)</td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>`SGLANG_CACHE_DIT_DMD_RIDGE`</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>1e-8</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>DMD Tikhonov regularization added to the inverted singular values</td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>`SGLANG_CACHE_DIT_DMD_SVD_PRECISION`</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>medium</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>DMD SVD precision (low/medium/high)</td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>`SGLANG_CACHE_DIT_SCM_PRESET`</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>none</td>
@@ -344,6 +369,31 @@ For dual-transformer models (e.g., Wan2.2 with high/low-noise experts), these va
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>(from primary)</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>TaylorSeer order (1 or 2)</td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_CACHE_DIT_SECONDARY_DMD</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>(from primary)</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>Enable the DMD calibrator (mutually exclusive with TaylorSeer)</td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_CACHE_DIT_SECONDARY_DMD_HISTORY</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>(from primary)</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>DMD snapshot window length</td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_CACHE_DIT_SECONDARY_DMD_RANK</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>(from primary)</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>DMD SVD truncation rank (0 = automatic)</td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_CACHE_DIT_SECONDARY_DMD_RIDGE</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>(from primary)</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>DMD Tikhonov regularization term</td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_CACHE_DIT_SECONDARY_DMD_SVD_PRECISION</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>(from primary)</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>DMD SVD precision (low/medium/high)</td>
</tr>
</tbody>
</table>
+1 -1
View File
@@ -112,7 +112,7 @@ runai = ["runai-model-streamer[s3,gcs,azure]>=0.15.7"]
diffusion = [
"addict==2.4.0",
"av==16.1.0",
"cache-dit==1.3.0",
"cache-dit==1.5.1",
"cloudpickle==3.1.2",
"diffusers==0.37.0",
"imageio==2.36.0",
@@ -0,0 +1,405 @@
---
name: sglang-diffusion-cache-dit
description: "Workflow for upgrading/integrating cache-dit in SGLang diffusion (multimodal_gen): DBCache, DMD calibrator, TaylorSeer, SVDQuant DQ; porting upstream PRs and resolving conflicts against the per-request knob system; adding new cache knobs; building the sglang generate CLI test matrix; precision validation (PSNR / log evidence); troubleshooting environment issues (wheel ABI, svdq extension, flashinfer conflicts). Use when upgrading or integrating cache-dit in sglang diffusion, porting cache-dit PRs with conflicts, adding cache knobs, running the sglang generate CLI test matrix, or validating precision (PSNR) for DBCache/DMD/SVDQuant(DQ) paths."
user-invocable: true
---
# SGLang Diffusion × Cache-DiT Integration
Path placeholders used throughout: `<sglang_dir>` = sglang repo root, `<cache_dit_dir>` = local cache-dit repo root (if not present locally, clone it: `git clone https://github.com/vipshop/cache-dit`), `<flux_model_dir>` = the DiT model weights actually under test (**if the user has not specified the model, ask for the model name and checkpoint path first — the workflow is not FLUX-specific; FLUX.1-dev is only the reference run**), `<cuda_home>` = CUDA toolkit root (typically `/usr/local/cuda`), `<gpu_id>` = a free GPU index.
## GATE CHECK (confirm before starting)
```
STOP — are all of the following confirmed?
1. All work happens in a dedicated env for sglang diffusion testing (e.g. `conda activate sgl`),
fully isolated from cache-dit/ffpa dev envs: never touch other envs, and never let sglang
dependencies leak into them. If the env does not exist, create a new dedicated one first
(conda/venv both fine, see §3) — never reuse an existing env just to save effort. Install
missing dependencies directly into the dedicated env.
2. GPU: run sglang jobs on the GPUs allocated for them (e.g. CUDA_VISIBLE_DEVICES=<gpu_id>);
other GPUs may be busy with other jobs.
3. The test artifact directory <sglang_dir>/.tmp/{task}/ exists (.tmp/ is gitignored; never
write test outputs into the repo root).
4. CLI args have been verified once via `sglang generate --help | grep -- <every arg you use>`
(the old --warmup from earlier docs/PRs is gone; it is now
--warmup-mode {off,request,server} + --warmup-steps N).
5. Plan-time alignment — align ALL of the following with the user while drafting the plan
(before any run), not afterwards:
- Model under test: model name + weights path. Do NOT default to FLUX.1-dev; adapt the
case-name prefix and the PSNR baseline table to the actual model.
- Generation settings: resolution, step count, prompt/seed (reference run used
1024x1024 / 28 steps — follow the user's actual setup instead).
- Local cache-dit checkout: does `<cache_dit_dir>` exist? The yaml configs in §4 come from
`<cache_dit_dir>/examples/configs/`; if absent, agree with the user whether to
`git clone https://github.com/vipshop/cache-dit` or obtain the configs another way.
Anything unknown → ASK the user first.
NO → fix these before touching anything.
```
**Hard rules**
- Conflict resolution: **keep the target branch's refactored structure** (the knob system in §2). If the upstream PR's direct-write style targets code that has since been refactored, re-inject it following the new pattern; never revert the target branch's refactor.
- After touching cache-related modules, run **all** `test/unit/test_cache_dit*.py` (not just one file — a skipped stub test once left a 9/9 ImportError that only surfaced at CLI stage).
## 1. Integration map (4 files on the sglang side)
| File | Responsibility |
|------|----------------|
| `python/pyproject.toml` | `cache-dit==x.y.z` version pin in the diffusion extra |
| `python/sglang/multimodal_gen/envs.py` | env vars: annotation section + lazy getters + `_CACHE_DIT_SECONDARY_CONFIGS` + special bool getters |
| `python/sglang/multimodal_gen/runtime/cache/cache_dit_integration.py` | `CacheDitConfig`, `enable_cache_on_transformer` (single/dual transformer), custom BlockAdapter, per-request knob validation set, calibrator construction |
| `python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py` | `_build_cache_dit_config()` + `_cache_dit_knob()`**the single injection point for new cache parameters** |
Data flow: `env (with secondary fallback) → knob (request > env) → CacheDitConfig fields → mutual-exclusion guard → cache-dit Config (DBCacheConfig / DMDCalibratorConfig / TaylorSeerCalibratorConfig)`
Key facts:
- Knob priority: request override (`sampling_params.cache_dit_params`) > env; a secondary knob first inherits the request-level primary value, then falls back to `SGLANG_CACHE_DIT_SECONDARY_*`, then to the primary env default.
- A single change in `_build_cache_dit_config` automatically covers both the primary/secondary call sites **and** the minimax_h3 subclass (its override calls `super()`).
- Dual-transformer models (wan2.2 etc.) can reuse the primary config wholesale for the secondary transformer (`_cache_dit_secondary_uses_primary_config()`); new fields are inherited for free.
- The calibrator is a single slot: DMD × TaylorSeer are mutually exclusive; `_assert_calibrator_exclusive(config, label)` guards both enable entry points (enabling both raises ValueError).
## 2. Standard six steps to add a cache knob (DMD as the example)
1. **envs.py annotation section**: add the primary + secondary variables (e.g. `SGLANG_CACHE_DIT_DMD: bool = False`).
2. **envs.py getter**: non-bool → add a `(SUFFIX, type, default)` tuple to `_CACHE_DIT_SECONDARY_CONFIGS` (auto-generates the secondary fallback); bool → write a dedicated `_secondary_xxx_getter` (`get_bool_env_var(SECONDARY_X, default=os.getenv(PRIMARY_X, "false"))`) and register it.
3. **cache_dit_integration.py**: add the field to `CacheDitConfig` (with docstring and default).
4. **cache_dit_integration.py**: add the key to `CACHE_DIT_REQUEST_KNOB_KEYS` (symmetric with taylorseer; automatically joins the request/secondary validation sets and the remount detection `cache_dit_overrides_key`).
5. **denoising.py**: append a knob() entry after taylorseer_order inside `_build_cache_dit_config`:
`enable_dmd=knob("enable_dmd", envs.SGLANG_CACHE_DIT_DMD, envs.SGLANG_CACHE_DIT_SECONDARY_DMD, secondary=secondary)`
6. **Calibrator/special logic**: construct the corresponding cache-dit Config inside `enable_cache_on_transformer` and `enable_cache_on_dual_transformer` (DMD: `if enable_dmd: DMDCalibratorConfig(...) elif enable_taylorseer: ...`), and add the field to the log format string (placeholder and argument counts must stay aligned).
Conflict forecast when porting an upstream PR: `denoising.py` always conflicts (the PR's direct `envs.XXX` → Config construction was refactored away on the target branch) → apply only step 5; `cache_dit_integration.py` conflicts locally (the per-request knob system added on the target branch sits right next to the PR's insertion points) → keep the target branch's frozenset and append the new keys.
## 3. Environment setup (dedicated env)
**Environment principle (mandatory, top priority)**: everything — pip installs/uninstalls, source builds, unit tests and CLI runs — happens inside the dedicated env (e.g. `conda activate sgl`). It is fully isolated from the cache-dit/ffpa dev envs: never install/uninstall packages in other envs, and never let sglang dependency changes (torch, cache-dit, flashinfer, ...) leak into them. If the env does not exist yet, create a dedicated one before continuing — do not reuse any existing env:
```bash
conda create -n sgl python=3.12
conda activate sgl
cd <sglang_dir>
pip install -e ".[diffusion]" --no-build-isolation
```
(Any other virtualenv mechanism works as well; what matters is **dedicated and isolated**. After switching envs, confirm versions with `pip show sglang cache-dit` before testing.)
```bash
conda activate sgl
# Pure Python (no SVDQuant):
pip install cache-dit==<ver> # or cache-dit-cu13==<ver> --no-deps
# SVDQuant (PTQ nvfp4) requires a working CUDA extension. The PyPI wheel can be ABI-incompatible
# with a newer torch (undefined symbol: materialize_cow_storage → the wheel was built against an
# older torch). In that case, build from the local cache-dit source tree
# (no local checkout? git clone https://github.com/vipshop/cache-dit first):
cd <cache_dit_dir>
export CUDA_HOME=<cuda_home>
pip install setuptools-scm # missing → metadata-generation-failed
CACHE_DIT_BUILD_SVDQUANT=1 pip install ".[quantization]" --no-build-isolation
# Extension self-check (pinpoints the load error in one step):
python -c "from cache_dit.quantization.svdquant import svdq_is_available, svdq_get_load_error as e; print(svdq_is_available(), e())"
```
Other environment pitfalls:
- A stale `flashinfer-cubin` that mismatches the flashinfer main package → blocks every sglang import; if the installed flashinfer has no matching cubin release (e.g. 0.6.18), `pip uninstall flashinfer-cubin` directly.
- The dedicated env may lack pytest → the test files use `unittest.main()` style; run `python <test_file.py>` directly.
- `test_cache_dit_integration.py` **replaces the cache_dit module with a stub** (no real install required): after changing top-level imports in `cache_dit_integration.py`, sync the stub's top-level symbols in `_install_cache_dit_stub()` (missing BlockAdapterRegister/Parallelism*/DMDCalibratorConfig once caused 9/9 ImportError).
## 4. CLI test matrix (reference run: PRO 5000, FLUX.1-dev, 1024×1024, 28 steps)
Nine-case design (backend × acceleration feature combos):
| # | backend | feature | driven by |
|---|---------|---------|-----------|
| 1-3 | SGLD (default) | baseline / DBCache / +DMD | env: `SGLANG_CACHE_DIT_ENABLED=true` (+`SGLANG_CACHE_DIT_DMD=true`) |
| 4-6 | diffusers | baseline / DBCache / +DMD | yaml: `cache.yaml` / `cache_dmd.yaml` (<cache_dit_dir>/examples/configs/) |
| 7-9 | diffusers | SVDQ nvfp4 / +compile / +compile+DBCache+DMD | yaml: `blackwell/quantize_svdq.yaml` / `blackwell/cache_dmd_svdq.yaml` + `--enable-torch-compile` |
Common args: `--model-path=$FLUX_DIR --log-level=info --prompt='...' --width=1024 --height=1024 --num-inference-steps=28 --warmup-mode request --warmup-steps 1 --dit-cpu-offload false --text-encoder-cpu-offload false --save-output --output-path .tmp/{task}/outputs`; compile cases add `--warmup-steps 28`; the blackwell yaml is already `svdq_nvfp4_r128_dq` (nvfp4, not int4) — grep to confirm before running.
**run_case script pattern** (saved as `.tmp/{task}/run_matrix.sh`):
- `timeout <1800~3600>` guard + stdout redirect to `logs/{name}.log` + `summary.log` records `PASS/FAIL` (dual criteria: rc + png existence) + a failure does not abort the remaining cases.
- When several cases fail with a common root cause, fix it and **re-run only the failed segment** via a small sub-script, not the whole matrix.
- Submit in the background; then do a **one-shot** health check (`sleep 45-60 && tail -3 logs/<first_case>.log && nvidia-smi -i <gpu_id>`) to confirm the model is loading without CLI errors, then stop and wait for the completion notification — **never poll**.
- Save long verification scripts (PSNR / perf extraction) as `.py` files; avoid long `python -c` one-liners (nested f-strings are error-prone).
**Fixed failure-diagnosis order**:
1. Check `summary.log` and whether the output files were produced (`EOFError` / `worker did not terminate gracefully, forcing` / leaked-semaphore lines in the tail are multiprocess shutdown noise, not the failure itself);
2. `grep -iE 'error|assert|exception|raise' logs/{case}.log | head` to find the **first** traceback (the real cause is often mid-log, e.g. `AssertionError: Quantization backend ... not supported`);
3. Use package-level diagnostics when available (`svdq_is_available()/svdq_get_load_error()`).
### 4.1 Full command reference (battle-tested during the 1.5.1 upgrade; adapt and reuse)
**Unit tests + negatives (mandatory after touching cache modules)**:
```bash
cd <sglang_dir>/.tmp/{task}
conda activate sgl
# Run ALL cache-dit-related tests via glob (unittest style, direct run; the env may lack pytest)
for t in ../../python/sglang/multimodal_gen/test/unit/test_cache_dit*.py; do echo "== $t"; python $t 2>&1 | tail -3; done
# Mutual-exclusion negatives already exist as unit tests
# (test_both_calibrators_raise_on_{,dual_}transformer) — no ad-hoc script needed
```
**CLI matrix driver script** (`.tmp/{task}/run_matrix.sh`; run `bash run_matrix.sh` in the background):
```bash
#!/bin/bash
set -u
BASE=<sglang_dir>/.tmp/{task}
FLUX_DIR=<flux_model_dir> # model actually under test — ask the user if not specified
CFG=<cache_dit_dir>/examples/configs
OUT=$BASE/outputs; LOGS=$BASE/logs; mkdir -p "$OUT" "$LOGS"
export CUDA_VISIBLE_DEVICES=<gpu_id>
PROMPT='A fantasy landscape with mountains and a river, detailed, vibrant colors'
COMMON=(
--model-path="$FLUX_DIR" --log-level=info --prompt="$PROMPT"
--width=1024 --height=1024 --num-inference-steps=28
--warmup-mode request --warmup-steps 1 # old --warmup is gone; verify via --help first
--dit-cpu-offload false --text-encoder-cpu-offload false
--save-output --output-path "$OUT"
)
run_case() { # $1=name $2=timeout_s; remaining args are case-specific
local name=$1; shift; local timeout_s=$1; shift
echo "[$(date '+%H:%M:%S')] START $name" >> "$LOGS/summary.log"
timeout "$timeout_s" sglang generate "${COMMON[@]}" "$@" \
--output-file-name "$name.png" > "$LOGS/$name.log" 2>&1
local rc=$?
[[ $rc -eq 0 && -f "$OUT/$name.png" ]] && st=PASS || st=FAIL
echo "[$(date '+%H:%M:%S')] $st $name (rc=$rc, png=$([[ -f $OUT/$name.png ]] && echo yes || echo no))" >> "$LOGS/summary.log"
}
# SGLD triple (env-driven)
run_case flux_sgld 1800
SGLANG_CACHE_DIT_ENABLED=true run_case flux_cache_sgld 1800
SGLANG_CACHE_DIT_ENABLED=true SGLANG_CACHE_DIT_DMD=true run_case flux_cache_dmd_sgld 1800
# diffusers triple (yaml-driven)
run_case flux_diffusers 1800 --backend diffusers
run_case flux_cache_diffusers 1800 --backend diffusers --cache-dit-config "$CFG/cache.yaml"
run_case flux_cache_dmd_diffusers 1800 --backend diffusers --cache-dit-config "$CFG/cache_dmd.yaml"
# SVDQ nvfp4 triple (requires the svdq extension)
run_case flux_svdq_nvfp4_diffusers 2400 --backend diffusers --cache-dit-config "$CFG/blackwell/quantize_svdq.yaml"
run_case flux_svdq_nvfp4_compile_diffusers 3600 --backend diffusers --warmup-steps 28 \
--enable-torch-compile --cache-dit-config "$CFG/blackwell/quantize_svdq.yaml"
run_case flux_cache_dmd_svdq_nvfp4_compile_diffusers 3600 --backend diffusers --warmup-steps 28 \
--enable-torch-compile --cache-dit-config "$CFG/blackwell/cache_dmd_svdq.yaml"
echo "[$(date '+%H:%M:%S')] MATRIX DONE" >> "$LOGS/summary.log"
```
When several cases fail with a common cause, copy the script keeping only the failed segment and re-run (the run_svdq.sh pattern).
**One-shot health check after startup** (then wait for completion; do not poll):
```bash
sleep 60 && tail -3 logs/flux_sgld.log | cut -c1-160 && nvidia-smi --query-gpu=index,memory.used --format=csv,noheader -i <gpu_id>
# Expect: model loading / inferring + tens of GB of VRAM in use;
# `ambiguous option` / Traceback → kill immediately and fix the args
```
**Feature-activation verification (grep the logs)**:
```bash
grep -E 'DMD=True|Calibrator Config: DMD' logs/flux_cache_dmd_sgld.log # DMD active
grep -E 'Match Blocks|Collected Context Config' logs/flux_cache_sgld.log # DBCache active
grep -E 'SVDQuant.*Type: svdq_nvfp4_r128_dq' logs/flux_svdq_nvfp4_diffusers.log # quantization active
```
**Perf extraction**:
```bash
for f in logs/*.log; do echo "$f: $(grep -oE 'finished in [0-9.]+ seconds' $f | head -1) $(grep -oE '[0-9.]+it/s' $f | tail -1)"; done
```
**Quantitative PSNR/SSIM comparison** — prefer the `cache-dit-metrics` CLI (ships with cache-dit; methodology reference: the cache-dit-model-integration skill's `references/testing.md`):
```bash
# Compare each accelerated result against the same-backend baseline
cache-dit-metrics psnr ssim -i1 outputs/flux_sgld.png -i2 outputs/flux_cache_sgld.png
cache-dit-metrics psnr ssim -i1 outputs/flux_sgld.png -i2 outputs/flux_cache_dmd_sgld.png
cache-dit-metrics psnr ssim -i1 outputs/flux_diffusers.png -i2 outputs/flux_svdq_nvfp4_diffusers.png
cache-dit-metrics psnr ssim -i1 outputs/flux_svdq_nvfp4_compile_diffusers.png \
-i2 outputs/flux_cache_dmd_svdq_nvfp4_compile_diffusers.png
```
Fallback when the CLI is unavailable (save as `.tmp/{task}/psnr.py` and run; avoid long python -c) — PSNR only, no SSIM:
```python
import numpy as np, torch
from PIL import Image
def load(p): return torch.from_numpy(np.array(Image.open(p))).float() / 255.0
def psnr(a, b):
mse = ((a - b) ** 2).mean().item()
return float('inf') if mse == 0 else 10 * np.log10(1.0 / mse)
pairs = [ # (label, baseline, accelerated output) — fill per actual cases
('sgld: cache vs base', 'flux_sgld.png', 'flux_cache_sgld.png'),
('sgld: cache+dmd vs base', 'flux_sgld.png', 'flux_cache_dmd_sgld.png'),
('diff: svdq vs base', 'flux_diffusers.png', 'flux_svdq_nvfp4_diffusers.png'),
('svdq: cache+dmd+compile vs svdq', 'flux_svdq_nvfp4_compile_diffusers.png',
'flux_cache_dmd_svdq_nvfp4_compile_diffusers.png'),
]
import os; os.chdir(os.path.dirname(__file__) + '/outputs')
for name, a, b in pairs: print(f'{name:38s} PSNR = {psnr(load(a), load(b)):6.2f} dB')
```
### 4.2 Command-by-command edition (requirements-doc style, for single-case debugging / manual runs; use the 4.1 script for batch regression)
Environment setup (run once before all commands):
```bash
conda activate sgl
cd <sglang_dir>
export FLUX_DIR=<flux_model_dir> # model actually under test — ask the user if not specified
export CUDA_VISIBLE_DEVICES=<gpu_id>
mkdir -p .tmp/{task}/outputs
# SVDQuant cases require: pip install cache-dit-cu13==<ver> --no-deps;
# fall back to a source build if the wheel is torch-ABI-incompatible (see §3);
# uninstall flashinfer-cubin if a stale copy reports a version mismatch
```
SGLD backend (env-driven):
```bash
# baseline
sglang generate --model-path=$FLUX_DIR --log-level=info \
--prompt='A fantasy landscape with mountains and a river, detailed, vibrant colors' \
--width=1024 --height=1024 --num-inference-steps=28 \
--warmup-mode request --warmup-steps 1 \
--dit-cpu-offload false --text-encoder-cpu-offload false \
--save-output --output-path .tmp/{task}/outputs --output-file-name flux_sgld.png
# DBCache
SGLANG_CACHE_DIT_ENABLED=true \
sglang generate --model-path=$FLUX_DIR --log-level=info \
--prompt='A fantasy landscape with mountains and a river, detailed, vibrant colors' \
--width=1024 --height=1024 --num-inference-steps=28 \
--warmup-mode request --warmup-steps 1 \
--dit-cpu-offload false --text-encoder-cpu-offload false \
--save-output --output-path .tmp/{task}/outputs --output-file-name flux_cache_sgld.png
# DBCache + DMD Calibrator
SGLANG_CACHE_DIT_ENABLED=true SGLANG_CACHE_DIT_DMD=true \
sglang generate --model-path=$FLUX_DIR --log-level=info \
--prompt='A fantasy landscape with mountains and a river, detailed, vibrant colors' \
--width=1024 --height=1024 --num-inference-steps=28 \
--warmup-mode request --warmup-steps 1 \
--dit-cpu-offload false --text-encoder-cpu-offload false \
--save-output --output-path .tmp/{task}/outputs --output-file-name flux_cache_dmd_sgld.png
```
Diffusers backend (yaml-driven, `CFG=<cache_dit_dir>/examples/configs`):
```bash
# baseline
sglang generate --model-path=$FLUX_DIR --backend diffusers --log-level=info \
--prompt='A fantasy landscape with mountains and a river, detailed, vibrant colors' \
--width=1024 --height=1024 --num-inference-steps=28 \
--warmup-mode request --warmup-steps 1 \
--dit-cpu-offload false --text-encoder-cpu-offload false \
--save-output --output-path .tmp/{task}/outputs --output-file-name flux_diffusers.png
# DBCache
sglang generate --model-path=$FLUX_DIR --backend diffusers --log-level=info \
--prompt='A fantasy landscape with mountains and a river, detailed, vibrant colors' \
--width=1024 --height=1024 --num-inference-steps=28 \
--warmup-mode request --warmup-steps 1 \
--dit-cpu-offload false --text-encoder-cpu-offload false \
--cache-dit-config $CFG/cache.yaml \
--save-output --output-path .tmp/{task}/outputs --output-file-name flux_cache_diffusers.png
# DBCache + DMD Calibrator
sglang generate --model-path=$FLUX_DIR --backend diffusers --log-level=info \
--prompt='A fantasy landscape with mountains and a river, detailed, vibrant colors' \
--width=1024 --height=1024 --num-inference-steps=28 \
--warmup-mode request --warmup-steps 1 \
--dit-cpu-offload false --text-encoder-cpu-offload false \
--cache-dit-config $CFG/cache_dmd.yaml \
--save-output --output-path .tmp/{task}/outputs --output-file-name flux_cache_dmd_diffusers.png
```
SVDQuant W4A4 NVFP4 (requires the svdq extension; the blackwell yaml is already nvfp4 — grep to confirm before running):
```bash
# SVDQuant W4A4 NVFP4
sglang generate --model-path=$FLUX_DIR --backend diffusers --log-level=info \
--prompt='A fantasy landscape with mountains and a river, detailed, vibrant colors' \
--width=1024 --height=1024 --num-inference-steps=28 \
--warmup-mode request --warmup-steps 1 \
--dit-cpu-offload false --text-encoder-cpu-offload false \
--cache-dit-config $CFG/blackwell/quantize_svdq.yaml \
--save-output --output-path .tmp/{task}/outputs --output-file-name flux_svdq_nvfp4_diffusers.png
# + compile
sglang generate --model-path=$FLUX_DIR --backend diffusers --log-level=info \
--prompt='A fantasy landscape with mountains and a river, detailed, vibrant colors' \
--width=1024 --height=1024 --num-inference-steps=28 \
--warmup-mode request --warmup-steps 28 \
--dit-cpu-offload false --text-encoder-cpu-offload false \
--enable-torch-compile \
--cache-dit-config $CFG/blackwell/quantize_svdq.yaml \
--save-output --output-path .tmp/{task}/outputs --output-file-name flux_svdq_nvfp4_compile_diffusers.png
# + compile + DBCache + DMD Calibrator
sglang generate --model-path=$FLUX_DIR --backend diffusers --log-level=info \
--prompt='A fantasy landscape with mountains and a river, detailed, vibrant colors' \
--width=1024 --height=1024 --num-inference-steps=28 \
--warmup-mode request --warmup-steps 28 \
--dit-cpu-offload false --text-encoder-cpu-offload false \
--enable-torch-compile \
--cache-dit-config $CFG/blackwell/cache_dmd_svdq.yaml \
--save-output --output-path .tmp/{task}/outputs --output-file-name flux_cache_dmd_svdq_nvfp4_compile_diffusers.png
```
## 5. Precision and feature verification
**Hard log evidence (check this before PSNR)**:
```
Enabling cache-dit ... DMD=True (history=6, rank=0, svd=medium), TaylorSeer=False ..., steps=28
[Cache-DiT] Collected Context Config: DBCache_F1B0_W4I1M0MC3_R0.24_N28_CFG0, Calibrator Config: DMD_H(6, medium)
[Cache-DiT] Match Blocks: CachedBlocks_Pattern_0_1_2, for transformer_blocks ...
[Cache-DiT] SVDQuant Type: svdq_nvfp4_r128_dq, Rank: 128
```
Every accelerated case must show its corresponding line in the log before the feature counts as genuinely active.
**PSNR reference baselines** (FLUX.1-dev on a PRO 5000, vs the same backend without acceleration; sglang's default DBCache R=0.24 is aggressive, so magnitudes differ from the cache-dit-side PSNR>30 standard — don't misjudge):
| comparison | typical PSNR |
|------------|--------------|
| SGLD: cache vs base | ≈23.5 dB |
| SGLD: cache+DMD vs base | ≈24.5 dB (**DMD should slightly beat pure cache**; if worse, investigate) |
| diffusers: cache vs base | ≈31 dB |
| diffusers: cache+DMD vs base | ≈29 dB |
| diffusers: svdq nvfp4 vs fp16 | ≈23.4 dB (normal W4A4 quantization loss) |
| within svdq stack: cache+DMD vs plain svdq | ≈28 dB |
**SSIM matters as much as PSNR** — always compute both (`cache-dit-metrics psnr ssim`). PSNR alone cannot detect structural corruption: a garbled image can still show PSNR > 20 dB, while SSIM collapses (< 0.5). If PSNR looks reasonable but SSIM is low, treat the output as corrupted and investigate — do not accept it.
Visually inspect 2-3 key PNGs first (baseline / DMD / full stack), then compute PSNR. Performance reference (single GPU): fp16 diffusers 17.15s → svdq+compile+DBCache+DMD 4.01s (≈4.3x); extract with grep `'finished in [0-9.]+ seconds'` and `it/s`.
## 6. Wrap-up
- Write the PR commit message in English to `.tmp/{task}/commit_msg.txt` (title `[Diffusion] Cache-DiT x.y.z: ...` + per-file changes + validation data); squash with `git commit -F`.
- Version bumps and cherry-picks keep the upstream commits (attribution); do not push without review.
- Record important findings and pitfalls in the repo-level knowledge base.
## 7. Pitfall quick reference
| pitfall | symptom | fix |
|---------|---------|-----|
| cache-dit wheel ABI incompatibility | `undefined symbol: _ZN3c104impl3cow...` (materialize_cow_storage), svdq_is_available()=False | source build (§3): CUDA_HOME + setuptools-scm |
| stale flashinfer-cubin | every import raises RuntimeError: version mismatch | `pip uninstall flashinfer-cubin` |
| stub tests out of sync | test_cache_dit_integration 9/9 ImportError | add the new top-level symbols to the stub |
| CLI arg drift | `ambiguous option: --warmup` | verify via `--help` first; use `--warmup-mode request --warmup-steps N` |
| running only one test file | missed regressions | run all `test_cache_dit*.py` via glob |
| mistaking shutdown noise for failure | `forcing`/`EOFError` in the log tail | check summary.log + png first; grep the first traceback |
| no pytest in the env | No module named pytest | run `python <test>.py` directly (unittest style) |
## Current Code Areas
| File | Role |
| --- | --- |
| `python/pyproject.toml` | pins the `cache-dit==x.y.z` version in the diffusion extra |
| `python/sglang/multimodal_gen/envs.py` | primary/secondary cache-dit env vars and lazy getters, including `_CACHE_DIT_SECONDARY_CONFIGS` |
| `runtime/cache/cache_dit_integration.py` | `CacheDitConfig`, single/dual-transformer enable paths, custom BlockAdapter, per-request knob validation, calibrator construction and mutual-exclusion guard |
| `runtime/pipelines_core/stages/denoising.py` | `_build_cache_dit_config()` + `_cache_dit_knob()`, the single injection point for new cache knobs |
| `python/sglang/multimodal_gen/test/unit/test_cache_dit_integration.py` | stub-based unit tests (no real cache-dit install needed); calibrator selection and exclusivity |
| `python/sglang/multimodal_gen/test/unit/test_cache_dit_per_request.py` | per-request knob reachability and secondary-inherits-primary coverage |
| `<cache_dit_dir>/examples/configs/` | yaml configs driving the diffusers-backend CLI cases (`cache.yaml`, `cache_dmd.yaml`, `blackwell/quantize_svdq.yaml`, `blackwell/cache_dmd_svdq.yaml`) |
## References
Authoritative usage references for cache in SGLang Diffusion; read them before changing user-facing cache behavior or docs:
- **`references/block_adapter.md`** — **the reference for custom BlockAdapter usage** (ideas only; all sglang adapter code stays in the sglang repo, and PatchFunctor is not recommended): `ForwardPattern` I/O contracts, `BlockAdapter` parameters (`has_separate_cfg`, `check_forward_pattern`), construction templates, third-party (non-diffusers) adapter rules, and cache interception pitfalls. Read it before writing or extending a custom BlockAdapter in `runtime/cache/cache_dit_integration.py`.
- `<cache_dit_dir>/.github/skills/cache-dit-model-integration/references/testing.md` — correctness-verification methodology behind the `cache-dit-metrics` CLI (PSNR+SSIM both mandatory, acceptance criteria, garbled-image red flags)
- `<sglang_dir>/docs/docs/sglang-diffusion/cache_dit.mdx` — Cache-DiT usage guide (env vars, per-request knobs, yaml configs)
- `<sglang_dir>/docs/docs/sglang-diffusion/caching-acceleration.mdx` — caching-acceleration guide (DBCache/DMD/TaylorSeer combinations, acceleration matrix)
- `<sglang_dir>/docs/docs/sglang-diffusion/` — the whole SGLang Diffusion docs directory; auxiliary reference for related topics (quantization, parallelism, installation, performance)
@@ -0,0 +1,359 @@
# BlockAdapter Reference
When to read this: read this file when writing or extending a custom `BlockAdapter` for SGLang Diffusion, selecting a `ForwardPattern`, or diagnosing cache interception issues. Return to `../SKILL.md` for the high-level workflow.
> **SGLang scope rules (read first):**
> - **This document is a reference for ideas only.** It documents how cache-dit itself implements Cache adapters; do not copy its registration/build flow verbatim into sglang.
> - **All sglang-diffusion BlockAdapter code strictly lives in the sglang repo** — adapters are constructed directly in `runtime/cache/cache_dit_integration.py` (sglang transformers are third-party, non-diffusers modules). Do NOT register sglang adapters in the cache-dit repo (`BlockAdapterRegister` / `block_adapters/__init__.py`); the registration flow described in cache-dit does not apply to sglang.
> - **PatchFunctor is NOT recommended.** Monkey-patching `transformer.forward()` does not fit the current sglang diffusion design. Treat §1.6 as diagnostic background: if a sglang transformer hits one of those structural pitfalls, fix the call structure on the sglang side instead of wiring a PatchFunctor.
> - The parts that do transfer to sglang: `ForwardPattern` selection (§1.2), `BlockAdapter` parameters such as `check_forward_pattern` / `has_separate_cfg` (§1.3), construction templates (§1.4), and the third-party (non-diffusers) adapter rules (§1.5).
## 1. Cache Integration: BlockAdapter + ForwardPattern
### 1.1 Concept
cache-dit's caching engine works by intercepting the forward pass of DiT transformer blocks. To do this, it needs to know:
1. **Where the blocks are** — which `ModuleList` attribute holds the repeated transformer blocks.
2. **What goes in and out** — the block's `forward()` input/output signature ("forward pattern").
3. **Any model quirks** — separate CFG passes, special patching needs, etc.
All of this is described by a single `BlockAdapter` dataclass instance.
### 1.2 ForwardPattern — The 6 Block I/O Contracts
`ForwardPattern` is an enum in `src/cache_dit/caching/forward_pattern.py`. It captures the hidden-state ordering and forward-signature shape of a family of transformer blocks. Choose the pattern that matches your block's `forward()` signature:
| Pattern | `forward()` inputs | `forward()` returns | Return_H_First | Return_H_Only | Forward_H_only | Typical Models |
| ------------------- | ------------------------------------------ | ------------------------------------------ | -------------- | ------------- | -------------- | ---------------------------------------------------------------------- |
| **Pattern_0** | `(hidden_states, encoder_hidden_states)` | `(hidden_states, encoder_hidden_states)` | `True` | `False` | `False` | Mochi, CogVideoX, CogView4, HunyuanVideo, EasyAnimate |
| **Pattern_1** | `(hidden_states, encoder_hidden_states)` | `(encoder_hidden_states, hidden_states)` | `False` | `False` | `False` | Flux transformer_blocks, QwenImage, SD3, VisualCloze |
| **Pattern_2** | `(hidden_states, encoder_hidden_states)` | `(hidden_states,)` | `False` | `True` | `False` | Wan, Allegro, Cosmos, LTX-1 |
| **Pattern_3** | `(hidden_states,)` | `(hidden_states,)` | `False` | `True` | `True` | Flux single_transformer_blocks, DiT, PixArt, Sana, Lumina2, SkyReelsV2 |
| **Pattern_4** | `(hidden_states,)` | `(hidden_states, encoder_hidden_states)` | `True` | `False` | `True` | (rare) |
| **Pattern_5** | `(hidden_states,)` | `(encoder_hidden_states, hidden_states)` | `False` | `False` | `True` | (rare) |
**How to determine the correct pattern for your model:**
1. Open the block's `forward()` method in diffusers source.
2. Check the parameter list: does it take only `hidden_states`, or also `encoder_hidden_states`? This determines `Forward_H_only`.
3. Check the return statement: does it return one tensor or two? In what order? This determines `Return_H_Only` / `Return_H_First`.
4. Match against the table above. If none fits exactly, open an issue.
### 1.3 BlockAdapter Parameters
Defined in `<cache_dit_dir>/src/cache_dit/caching/block_adapters/block_adapters.py`. Key parameters:
| Parameter | Type | Description |
| ------------------------- | -------------------------------------------------- | --------------------------------------------------------------------------------------------------------------------------------------------------------------------------- |
| `pipe` | `DiffusionPipeline` or `FakeDiffusionPipeline` | The pipeline instance (or a placeholder if no pipeline is available). |
| `transformer` | `nn.Module` or `List[nn.Module]` | The transformer module(s). Single module for most models; list of 2 for dual-transformer models (e.g., Wan 2.2 MoE). |
| `blocks` | `nn.ModuleList` or `List[nn.ModuleList]` | The block collection(s). Single ModuleList for most models; list of 2 for models with dual block types (e.g., Flux:`transformer_blocks` + `single_transformer_blocks`). |
| `forward_pattern` | `ForwardPattern` or `List[ForwardPattern]` | Must match`blocks` count. Single pattern for single block list; list of patterns for multiple block lists. |
| `check_forward_pattern` | `Optional[bool]` | Validate that each block's I/O matches the declared pattern. If left `None` (default), cache-dit **auto-detects**: `True` for `diffusers` transformers, `False` for third-party ones (`maybe_skip_checks()`); it is also forced `False` when the transformer already has an `_hf_hook` / `_diffusers_hook`. Set explicitly for new models. |
| `check_num_outputs` | `bool` | If`True`, cache-dit additionally validates that each block returns the exact number of outputs the pattern declares. Needed for models whose blocks can return a variable tuple (e.g., HiDream, HunyuanVideo 1.0). Default `False`. |
| `has_separate_cfg` | `bool` | Set `True` if the pipeline runs **two separate `transformer.forward()` calls** for the conditional and unconditional passes of Classifier-Free Guidance (CFG). Set `False` if the pipeline concatenates cond+uncond into a single batch and calls `transformer.forward()` **once**. See §1.3.1 for the decision guide and code patterns. |
| `patch_functor` | `PatchFunctor` or `None` | Optional pre-patch logic. Used when the model needs structural modification before caching hooks are installed (e.g., Flux dummy block merging, DiT re-patching). |
| `blocks_name` | `str` or `List[str]` | Override block attribute names (advanced). |
| `dummy_blocks_names` | `List[str]` | Names of blocks that should be treated as dummy/merged (advanced, e.g., Flux single_transformer_blocks when merged into transformer_blocks). |
### 1.3.1 `has_separate_cfg` — Decision Guide & Code Patterns
> **⚠️ This parameter is about the NUMBER of `transformer.forward()` calls per denoising step, NOT about whether CFG is enabled.** A model can use CFG (`guidance_scale > 1`) and still have `has_separate_cfg=False` if the pipeline batches cond+uncond into one forward call.
**Definition:**
| `has_separate_cfg` | Pipeline behavior per denoising step | Number of `transformer.forward()` calls |
|---|---|---|
| `True` | Pipeline calls `transformer(...)` **twice**: once with cond embeddings, once with uncond embeddings. The two outputs are combined by `noise_pred = uncond + scale * (cond - uncond)`. | **2** |
| `False` | Pipeline concatenates `[latents, latents]` into one batch, calls `transformer(...)` **once** with `encoder_hidden_states=[uncond, cond]`, then splits the output via `chunk(2)`. | **1** |
**Why it matters for caching:** cache-dit caches transformer block outputs. When `has_separate_cfg=True`, the cond and uncond passes have **independent cache contexts** (`"cond"` / `"uncond"`) because their inputs differ. When `False`, there is only one forward pass and one cache context. Setting this incorrectly causes the cache to mix cond/uncond states → garbled output.
**How to decide — read the pipeline's `__call__` denoising loop:**
**Pattern A → `has_separate_cfg=True`** (two separate forward calls):
```python
# WanPipeline (diffusers) — TWO calls, one for cond, one for uncond
latent_model_input = latents.to(transformer_dtype) # NOT concatenated
with current_model.cache_context("cond"):
noise_pred = current_model(
hidden_states=latent_model_input,
encoder_hidden_states=prompt_embeds, # cond embeddings
...
)[0]
if self.do_classifier_free_guidance:
with current_model.cache_context("uncond"):
noise_uncond = current_model( # SECOND forward call
hidden_states=latent_model_input, # same latents, NOT batched
encoder_hidden_states=negative_prompt_embeds, # uncond embeddings
...
)[0]
noise_pred = noise_uncond + guidance_scale * (noise_pred - noise_uncond)
```
**Tell-tale signs of Pattern A:**
- `latent_model_input` is NOT concatenated (no `torch.cat([latents] * 2)`).
- There are **two** `transformer(...)` / `current_model(...)` calls inside the loop.
- The second call uses `negative_prompt_embeds` / `negative_pooled_projections`.
- `cache_context("cond")` / `cache_context("uncond")` wrap the two calls.
**Models using Pattern A:** `Wan` (`has_separate_cfg=True`), `Flux` (with `do_true_cfg`), `QwenImage`, `CogView4`, `Cosmos`, `SkyReelsV2`, `Chroma`, `HunyuanImage`, `OvisImage`, `LongCatImage`, `GlmImage`, `Helios`, `ErnieImage`, `Krea2`, `JoyImage`, `BriaFibo`.
---
**Pattern B → `has_separate_cfg=False`** (single batched forward call):
```python
# AnyFlowPipeline (diffusers) — ONE call, cond+uncond batched together
latent_model_input = torch.cat([latents] * 2) if self.do_classifier_free_guidance else latents
# latent_model_input shape: (2*B, ...) — cond and uncond stacked
noise_pred = self.transformer(
hidden_states=latent_model_input, # batched cond+uncond
timestep=timestep,
encoder_hidden_states=prompt_embeds, # [uncond_embeds, cond_embeds] stacked
...
)[0]
if self.do_classifier_free_guidance:
noise_uncond, noise_pred = noise_pred.chunk(2) # split the batched output
noise_pred = noise_uncond + guidance_scale * (noise_pred - noise_uncond)
```
**Tell-tale signs of Pattern B:**
- `torch.cat([latents] * 2)` or `torch.cat([latents, latents])` before the forward call.
- `encoder_hidden_states` is `torch.cat([negative_prompt_embeds, prompt_embeds])` (stacked).
- Only **one** `transformer(...)` call inside the loop.
- Output is split via `noise_pred.chunk(2)` after the forward.
- No `cache_context("cond")` / `cache_context("uncond")` — single context.
**Models using Pattern B:** `AnyFlow` (`has_separate_cfg=False`), `ErnieImage` (`has_separate_cfg=False`), `Krea2` (when `guidance_scale=0`, no CFG at all), distilled models with CFG folded into weights (`guidance_scale=1.0`).
---
**Pattern C → `has_separate_cfg=False`** (no CFG at all, `guidance_scale=1.0`):
```python
# Distilled model — no CFG, single forward, single batch
noise_pred = self.transformer(
hidden_states=latents, # NOT concatenated
encoder_hidden_states=prompt_embeds, # only cond
...
)[0]
# No chunk(2), no noise_uncond, no CFG combination
```
**Tell-tale signs of Pattern C:**
- `guidance_scale=1.0` (or `0.0`).
- No `do_classifier_free_guidance` branch, no `torch.cat([latents]*2)`.
- Only one forward call with no cond/uncond splitting.
**Models using Pattern C:** AnyFlow (default `guidance_scale=1.0`, CFG folded into weights), ErnieImage Turbo, Krea2 Turbo (`guidance_scale=0.0`), ZImage Turbo (`guidance_scale=0.0`).
> **Note:** Pattern B and Pattern C both use `has_separate_cfg=False`. The difference is whether CFG is active (B: `guidance_scale > 1`, batched) or inactive (C: `guidance_scale <= 1`, single). In both cases there is only ONE `transformer.forward()` call, so cache-dit uses a single cache context.
**Quick decision flowchart:**
```
Read the pipeline __call__ denoising loop.
├─ Does it call transformer(...) TWICE per step
│ (once with cond, once with uncond)?
│ └─ YES → has_separate_cfg = True
├─ Does it torch.cat([latents]*2) and call transformer(...) ONCE,
│ then chunk(2) the output?
│ └─ YES → has_separate_cfg = False
└─ Does it call transformer(...) ONCE with no cat/chunk
(guidance_scale <= 1, no CFG)?
└─ YES → has_separate_cfg = False
```
### 1.4 Reference Templates: Constructing a BlockAdapter
These templates mirror how cache-dit's built-in adapters are written (`<cache_dit_dir>/src/cache_dit/caching/block_adapters/adapters.py`, read-only reference). The same `BlockAdapter(...)` construction applies in sglang's `runtime/cache/cache_dit_integration.py` — as plain construction code, **without** the `@BlockAdapterRegister.register(...)` decorator cache-dit uses internally.
#### Template A: Single block list (most common)
```python
adapter = BlockAdapter(
pipe=pipe,
transformer=pipe.transformer,
blocks=pipe.transformer.transformer_blocks,
forward_pattern=ForwardPattern.Pattern_0, # adjust to your model
check_forward_pattern=True,
)
```
#### Template B: Dual block lists (like Flux)
```python
# Standard Flux: both block types use Pattern_1.
# For Flux2 / Nunchaku variants: single_transformer_blocks use Pattern_3 instead.
adapter = BlockAdapter(
pipe=pipe,
transformer=pipe.transformer,
blocks=[
pipe.transformer.transformer_blocks,
pipe.transformer.single_transformer_blocks,
],
forward_pattern=[
ForwardPattern.Pattern_1,
ForwardPattern.Pattern_1,
],
check_forward_pattern=True,
)
```
#### Template C: Dual transformers (like Wan 2.2 MoE)
```python
adapter = BlockAdapter(
pipe=pipe,
transformer=[
pipe.transformer,
pipe.transformer_2, # second transformer (MoE)
],
blocks=[
pipe.transformer.blocks,
pipe.transformer_2.blocks,
],
forward_pattern=[
ForwardPattern.Pattern_2,
ForwardPattern.Pattern_2,
],
check_forward_pattern=True,
has_separate_cfg=True,
)
```
### 1.5 Third-Party (Non-Diffusers) Models
If your model does **not** come from the official `diffusers` library (e.g., it is defined in `sglang` or another third-party package), follow these rules:
**Do NOT hardcode `from diffusers import ...`.** Instead, use `_safe_import` with name-based matching, or simply skip the diffusers-specific import entirely.
**`_relaxed_assert` is NOT mandatory.** The function (`<cache_dit_dir>/src/cache_dit/caching/block_adapters/adapters.py`) checks `transformer.__module__` — if it does not start with `"diffusers"`, the function logs a warning and skips the strict type check automatically. For third-party models, you can:
- Omit `_relaxed_assert` entirely, or
- Call it with `allow_classes=None` to rely on the automatic skip behavior.
**Example — third-party BlockAdapter without `_relaxed_assert`:**
```python
# No `from diffusers import ...` — the transformer type is resolved at runtime.
adapter = BlockAdapter(
pipe=pipe,
transformer=pipe.transformer,
blocks=pipe.transformer.transformer_blocks,
forward_pattern=ForwardPattern.Pattern_0,
check_forward_pattern=True,
)
```
The same principle applies to cache-dit's distributed planners (CP, TP, TE-P, VAE-P), for reference: they never hardcode diffusers class names for third-party models either — dispatch matches on a registered descriptive name instead. SGLang does not add planners to cache-dit; anything parallelism-related is wired on the sglang side.
### 1.6 Interception Pitfalls — PatchFunctor Background (NOT recommended for sglang)
> ⚠️ **Always check for these pitfalls before declaring the cache integration "done."** A BlockAdapter that looks correct on paper can silently produce wrong results if the `transformer.forward()` has any of the structural issues below. When in doubt, run a full inference with caching enabled and compare PSNR/SSIM against the uncached baseline.
>
> **SGLang:** read this section for diagnosis only. cache-dit's remedy for these pitfalls is a `PatchFunctor` (a monkey-patch of `transformer.forward()`), which does not fit the current sglang diffusion design. If a sglang transformer hits one of these pitfalls, fix the call structure in sglang code instead.
The `BlockAdapter` works by intercepting the block-loop inside `transformer.forward()`. It replaces the original `ModuleList` (e.g., `self.transformer_blocks`) with `UnifiedBlocks` — a wrapper that injects cache look-up/save logic around each block call. This interception is mechanical: it relies on `inspect.signature` to bind arguments and on the assumption that the for-loop body contains **nothing but a single block call**. When the model's `forward()` violates these assumptions, the cache produces wrong results silently (no crash, just corrupted output).
A **`PatchFunctor`** is cache-dit's escape hatch for these cases: a monkey-patch that rewrites `transformer.forward()` *before* the `BlockAdapter` is applied. The two pitfall categories below explain **why** a structural fix is needed; the fix itself should land in sglang code, not a PatchFunctor.
#### Pitfall A: Block call argument mismatch (keyword vs positional)
**Problem**: `transformer.forward()` calls blocks with **keyword arguments** (e.g., `block(hidden_states=x, encoder_hidden_states=e, temb=t)`), but the block's `forward()` signature defines those parameters as **positional**. When cache-dit's `UnifiedBlocks` wrapper intercepts the call, it uses `inspect.signature.bind()` to match arguments — keyword-to-positional mismatches cause `bind()` to fail or bind to the wrong parameters.
**Symptom**: `TypeError` from `inspect.signature.bind()`, or the cache silently feeds wrong tensors to the block.
**Fix idea (reference)**: cache-dit's `LTX2PatchFunctor` rewrites the call site so positional parameters are passed positionally (matching the block's actual signature), keeping only truly keyword-only parameters as keyword args. In sglang, apply the same idea directly in the sglang-side call site.
**Canonical example — `LTX2PatchFunctor`** (`<cache_dit_dir>/src/cache_dit/caching/patch_functors/functor_ltx2.py`):
The original diffusers code for LTX-2.0 passes all block arguments as keywords:
```python
# Original (diffusers) — ALL keyword args:
hidden_states, audio_hidden_states = block(
hidden_states=hidden_states,
audio_hidden_states=audio_hidden_states,
encoder_hidden_states=encoder_hidden_states,
audio_encoder_hidden_states=audio_encoder_hidden_states,
temb=temb,
temb_audio=temb_audio,
...
)
```
The patched version converts the first four positional parameters to positional form, keeping the rest as keyword:
```python
# Patched — positional args match the block's forward(hidden_states, audio_hidden_states, ...):
hidden_states, audio_hidden_states = block(
hidden_states,
audio_hidden_states,
encoder_hidden_states,
audio_encoder_hidden_states,
temb=temb,
temb_audio=temb_audio,
...
)
```
**How to detect this pitfall**: Read the block's `forward()` signature in the diffusers source. Count how many parameters are positional (before any `*` or `*args`). Then check how `transformer.forward()` invokes the block — if it passes any of those positional params as keyword args, the call site needs the fix above.
#### Pitfall B: For-loop body has extra operations
**Problem**: The `for block in self.blocks:` loop in `transformer.forward()` contains operations *other than* the block call itself — such as `temb` reassignment, conditional checks, or tensor reshaping. After `CacheAdapter.apply()` replaces `self.blocks` with `UnifiedBlocks`, the caching wrapper **takes over the iteration** and only executes the block call; all extra operations inside the original loop body are **silently skipped**.
**Symptom**: Cache-enabled output is corrupted (low PSNR/SSIM, visual artifacts) because modulation parameters or intermediate tensors are stale or missing.
**Fix idea (reference)**: cache-dit's `ErnieImagePatchFunctor` moves the extra operations **outside** (before or after) the for-loop, so the loop body contains only the block call. In sglang, restructure the sglang-side forward the same way.
**Canonical example — `ErnieImagePatchFunctor`** (`<cache_dit_dir>/src/cache_dit/caching/patch_functors/functor_ernie_image.py`):
The original diffusers code reconstructs `temb` inside the loop body:
```python
# Original (diffusers) — temb reassigned INSIDE the for-loop:
for layer in self.layers:
temb = [shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp]
x = layer(x, rotary_pos_emb, temb, attention_mask=attention_mask)
```
After `CacheAdapter.apply()` replaces `self.layers` with `UnifiedBlocks`, the `temb = [...]` line is never executed — each block receives a stale or undefined `temb`. The patched version moves `temb` construction **before** the loop:
```python
# Patched — temb constructed ONCE before the loop:
temb = [shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp]
for layer in self.layers:
x = layer(x, rotary_pos_emb, temb, attention_mask=attention_mask)
```
**How to detect this pitfall**: Inspect the for-loop body in `transformer.forward()`. If *any* line between `for ... in self.XXX:` and the actual block call does something other than a trivial `if torch.is_grad_enabled()` guard, the loop body needs restructuring.
#### Beyond the two canonical pitfalls
Pitfalls A and B are the two simplest cases (fix at the call site, or hoist one line out of the loop). Real models often need heavier structural fixes. cache-dit ships 13+ `PatchFunctor`s under `<cache_dit_dir>/src/cache_dit/caching/patch_functors/` — browse them as a **source of fix ideas only**. Recurring patterns include:
- **Per-block `forward()` replacement + block-id injection** — when the loop body has *per-block* extra operations that cannot simply be hoisted (they depend on the block index). The functor patches `transformer.forward()` **and** each block's `forward()`, and injects a `_block_id` / `_layer_id` onto every block so the patched block can look up per-block data (skip-connection lists, per-block encoder states, control hints). Examples: `HiDreamPatchFunctor`, `HunyuanDiTPatchFunctor`, `WanVACEPatchFunctor`, `ChromaPatchFunctor`, `GlmImagePatchFunctor`, `BriaFiboPatchFunctor`.
- **Block signature modification** — rewriting a block's `forward()` signature so the caching wrapper can bind it (e.g. `FluxPatchFunctor` adds an `encoder_hidden_states` parameter to `FluxSingleTransformerBlock` in older diffusers).
- **Block-list merge / dummy blocks** — structurally merging two `ModuleList`s into one for unified caching (e.g. `FluxPatchFunctor` merging `transformer_blocks` + `single_transformer_blocks` when `dummy_blocks_names` is set).
For sglang, whatever fix pattern you borrow, the resulting code must keep the exact same signature and produce identical output with caching disabled (verify via PSNR/SSIM) — and it lands in the sglang repo, not as a cache-dit PatchFunctor.
## More references
We recommend reading the following files for additional context:
- cache related source code: `<cache_dit_dir>/src/cache_dit/caching/`
+38
View File
@@ -68,6 +68,11 @@ if TYPE_CHECKING:
SGLANG_CACHE_DIT_MC: int = 3
SGLANG_CACHE_DIT_TAYLORSEER: bool = False
SGLANG_CACHE_DIT_TS_ORDER: int = 1
SGLANG_CACHE_DIT_DMD: bool = False
SGLANG_CACHE_DIT_DMD_HISTORY: int = 6
SGLANG_CACHE_DIT_DMD_RANK: int = 0
SGLANG_CACHE_DIT_DMD_RIDGE: float = 1e-8
SGLANG_CACHE_DIT_DMD_SVD_PRECISION: str = "medium"
SGLANG_CACHE_DIT_SCM_PRESET: str = "none"
SGLANG_CACHE_DIT_SCM_COMPUTE_BINS: str | None = None
SGLANG_CACHE_DIT_SCM_CACHE_BINS: str | None = None
@@ -80,6 +85,11 @@ if TYPE_CHECKING:
SGLANG_CACHE_DIT_SECONDARY_MC: int = 3
SGLANG_CACHE_DIT_SECONDARY_TAYLORSEER: bool = False
SGLANG_CACHE_DIT_SECONDARY_TS_ORDER: int = 1
SGLANG_CACHE_DIT_SECONDARY_DMD: bool = False
SGLANG_CACHE_DIT_SECONDARY_DMD_HISTORY: int = 6
SGLANG_CACHE_DIT_SECONDARY_DMD_RANK: int = 0
SGLANG_CACHE_DIT_SECONDARY_DMD_RIDGE: float = 1e-8
SGLANG_CACHE_DIT_SECONDARY_DMD_SVD_PRECISION: str = "medium"
# model loading
SGLANG_USE_RUNAI_MODEL_STREAMER: bool = True
SGLANG_LINGBOT_ENABLE_INTERACTIVE_KV_WINDOW: bool = False
@@ -383,6 +393,19 @@ environment_variables: dict[str, Callable[[], Any]] = {
"SGLANG_CACHE_DIT_TAYLORSEER": _lazy_bool("SGLANG_CACHE_DIT_TAYLORSEER", "false"),
# TaylorSeer order (1 or 2)
"SGLANG_CACHE_DIT_TS_ORDER": _lazy_int("SGLANG_CACHE_DIT_TS_ORDER", 1),
# Enable DMD (Dynamic Mode Decomposition) calibrator (mutually exclusive
# with TaylorSeer). DMD forecasts Bn residuals via an exponential basis.
"SGLANG_CACHE_DIT_DMD": _lazy_bool("SGLANG_CACHE_DIT_DMD", "false"),
# DMD snapshot window length (>= 4 uniformly spaced snapshots to engage)
"SGLANG_CACHE_DIT_DMD_HISTORY": _lazy_int("SGLANG_CACHE_DIT_DMD_HISTORY", 6),
# DMD SVD truncation rank (0 = automatic)
"SGLANG_CACHE_DIT_DMD_RANK": _lazy_int("SGLANG_CACHE_DIT_DMD_RANK", 0),
# DMD Tikhonov regularisation term added to inverted singular values
"SGLANG_CACHE_DIT_DMD_RIDGE": _lazy_float("SGLANG_CACHE_DIT_DMD_RIDGE", 1e-8),
# DMD SVD precision mode: low, medium, high
"SGLANG_CACHE_DIT_DMD_SVD_PRECISION": _lazy_str(
"SGLANG_CACHE_DIT_DMD_SVD_PRECISION", "medium"
),
# SCM preset: none, slow, medium, fast, ultra
"SGLANG_CACHE_DIT_SCM_PRESET": _lazy_str("SGLANG_CACHE_DIT_SCM_PRESET", "none"),
# SCM custom compute bins (e.g., "8,3,3,2,2")
@@ -459,6 +482,10 @@ _CACHE_DIT_SECONDARY_CONFIGS = [
("RDT", float, "0.24"),
("MC", int, "3"),
("TS_ORDER", int, "1"),
("DMD_HISTORY", int, "6"),
("DMD_RANK", int, "0"),
("DMD_RIDGE", float, "1e-8"),
("DMD_SVD_PRECISION", str, "medium"),
]
@@ -493,6 +520,17 @@ environment_variables["SGLANG_CACHE_DIT_SECONDARY_TAYLORSEER"] = (
)
# Special handling for boolean secondary var (DMD)
def _secondary_dmd_getter():
return get_bool_env_var(
"SGLANG_CACHE_DIT_SECONDARY_DMD",
default=os.getenv("SGLANG_CACHE_DIT_DMD", "false"),
)
environment_variables["SGLANG_CACHE_DIT_SECONDARY_DMD"] = _secondary_dmd_getter
# end-env-vars-definition
def __getattr__(name: str):
# lazy evaluation of environment variables
@@ -22,16 +22,36 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__)
import cache_dit
from cache_dit import (
BlockAdapter,
DBCacheConfig,
ForwardPattern,
ParamsModifier,
TaylorSeerCalibratorConfig,
steps_mask,
)
from cache_dit.caching.block_adapters import BlockAdapterRegister
from cache_dit.parallelism import ParallelismBackend, ParallelismConfig
try:
from cache_dit import (
BlockAdapter,
BlockAdapterRegister,
DBCacheConfig,
DMDCalibratorConfig,
ForwardPattern,
ParallelismBackend,
ParallelismConfig,
ParamsModifier,
TaylorSeerCalibratorConfig,
steps_mask,
)
except ImportError:
# cache-dit < 1.5.0 exports BlockAdapterRegister only via submodules;
# DMDCalibratorConfig does not exist at all before 1.5.0
from cache_dit import (
BlockAdapter,
DBCacheConfig,
ForwardPattern,
ParamsModifier,
TaylorSeerCalibratorConfig,
steps_mask,
)
from cache_dit.caching.block_adapters import BlockAdapterRegister
from cache_dit.parallelism import ParallelismBackend, ParallelismConfig
# DMD calibrator requires cache-dit >= 1.5.0; guarded at enable time
DMDCalibratorConfig = None
from sglang.multimodal_gen.runtime.distributed.parallel_state import get_dit_group
@@ -61,7 +81,15 @@ def _patch_cache_dit_similarity():
_original_similarity = cache_manager.CachedContextManager.similarity
def patched_similarity(self, t1, t2, *, threshold, parallelized=False, prefix="Fn"):
def patched_similarity(
self: cache_manager.CachedContextManager,
t1: torch.Tensor,
t2: torch.Tensor,
*,
threshold: float,
parallelized: bool = False,
prefix: str = "Fn",
) -> bool:
if not parallelized:
return _original_similarity(
self,
@@ -205,6 +233,11 @@ CACHE_DIT_REQUEST_KNOB_KEYS = frozenset(
"max_continuous_cached_steps",
"enable_taylorseer",
"taylorseer_order",
"enable_dmd",
"dmd_history",
"dmd_rank",
"dmd_ridge",
"dmd_svd_precision",
}
)
CACHE_DIT_REQUEST_SCM_KEYS = frozenset(
@@ -276,6 +309,17 @@ class CacheDitConfig:
max_continuous_cached_steps: Maximum consecutive cached steps (DBCache MC).
enable_taylorseer: Whether to enable TaylorSeer calibrator.
taylorseer_order: Order of Taylor expansion (1 or 2).
enable_dmd: Whether to enable DMD (Dynamic Mode Decomposition) calibrator.
Mutually exclusive with enable_taylorseer; enabling both raises
ValueError at cache-enable time. DMD forecasts Bn residuals via an
exponential basis.
dmd_history: DMD snapshot window length (>= 4 uniformly spaced snapshots
are needed before the exponential fit engages; below the floor DMD
falls back to Taylor expansion automatically).
dmd_rank: SVD truncation rank of the DMD snapshot matrix; 0 selects it
from the spectrum automatically.
dmd_ridge: Tikhonov term added to the inverted singular values.
dmd_svd_precision: SVD precision mode for DMD ("low", "medium", "high").
num_inference_steps: Total number of inference steps (required for transformer-only mode).
steps_computation_mask: Binary mask for step-level caching (1=compute, 0=cache).
Generated by get_scm_mask() (wrapper around cache_dit.steps_mask()).
@@ -301,6 +345,13 @@ class CacheDitConfig:
# Diffusion Transformers, https://arxiv.org/pdf/2508.16211
enable_taylorseer: bool = False
taylorseer_order: int = 1
# DMD calibrator (mutually exclusive with TaylorSeer). DMD forecasts the Bn
# residual with an exponential basis. See cache_dit.DMDCalibratorConfig.
enable_dmd: bool = False
dmd_history: int = 6
dmd_rank: int = 0
dmd_ridge: float = 1e-8
dmd_svd_precision: str = "medium"
num_inference_steps: Optional[int] = None
# SCM fields (generated by _maybe_enable_cache_dit from env configuration)
steps_computation_mask: Optional[List[int]] = None
@@ -390,6 +441,54 @@ def _build_custom_block_adapter(
)
def _assert_calibrator_exclusive(config: CacheDitConfig, label: str = "transformer"):
"""Ensure at most one calibrator is enabled on a CacheDitConfig.
cache-dit accepts a single ``calibrator_config`` per transformer, so DMD and
TaylorSeer cannot run at the same time. ``enabled=False`` configs are skipped
by the callers before reaching here, so this only guards real enable paths.
Args:
config: The CacheDitConfig to validate.
label: Human-readable label (e.g. "primary"/"secondary") for the error
message when both calibrators are enabled.
Raises:
ValueError: If both ``enable_dmd`` and ``enable_taylorseer`` are True.
"""
if config.enable_dmd and config.enable_taylorseer:
raise ValueError(
f"DMD and TaylorSeer calibrators are mutually exclusive on "
f"{label}, but both are enabled. Please set only one of "
f"SGLANG_CACHE_DIT_DMD / SGLANG_CACHE_DIT_TAYLORSEER (or their "
f"SECONDARY_ variants) to true."
)
def _assert_dmd_supported(config: CacheDitConfig, label: str = "transformer"):
"""Reject DMD requests when the installed cache-dit predates DMD support.
DMDCalibratorConfig only exists since cache-dit 1.5.0; environments
pinning older cache-dit (e.g. the CI base jobs with 1.3.0) import this
module via the fallback import path where it is None.
Args:
config: The CacheDitConfig to validate.
label: Human-readable label for the error message.
Raises:
ValueError: If ``enable_dmd`` is True but cache-dit < 1.5.0 is installed.
"""
if config.enable_dmd and DMDCalibratorConfig is None:
raise ValueError(
f"DMD calibrator requires cache-dit >= 1.5.0, but cache-dit "
f"{getattr(cache_dit, '__version__', 'unknown')} is installed on "
f"{label}. Please upgrade cache-dit (e.g. pip install "
f"'cache-dit>=1.5.1') or disable DMD (SGLANG_CACHE_DIT_DMD / "
f"enable_dmd knob)."
)
def enable_cache_on_transformer(
transformer: torch.nn.Module,
config: CacheDitConfig,
@@ -415,6 +514,9 @@ def enable_cache_on_transformer(
if not config.enabled:
return transformer
_assert_calibrator_exclusive(config, label=model_name)
_assert_dmd_supported(config, label=model_name)
if config.num_inference_steps is None:
raise ValueError(
"num_inference_steps is required for transformer-only mode. "
@@ -452,9 +554,17 @@ def enable_cache_on_transformer(
steps_computation_policy=config.steps_computation_policy,
)
# Build calibrator config if TaylorSeer is enabled
# Build calibrator config. DMD and TaylorSeer are mutually exclusive
# (validated above); DMD takes the calibrator slot when enabled.
calibrator_config = None
if config.enable_taylorseer:
if config.enable_dmd:
calibrator_config = DMDCalibratorConfig(
dmd_history=config.dmd_history,
dmd_rank=config.dmd_rank,
dmd_ridge=config.dmd_ridge,
dmd_svd_precision=config.dmd_svd_precision,
)
elif config.enable_taylorseer:
calibrator_config = TaylorSeerCalibratorConfig(
taylorseer_order=config.taylorseer_order,
)
@@ -462,13 +572,17 @@ def enable_cache_on_transformer(
# Enable cache-dit on the transformer
logger.info(
"Enabling cache-dit on %s with config: Fn=%d, Bn=%d, W=%d, R=%.2f, MC=%d, "
"TaylorSeer=%s (order=%d), steps=%d",
"DMD=%s (history=%d, rank=%d, svd=%s), TaylorSeer=%s (order=%d), steps=%d",
model_name,
config.Fn_compute_blocks,
config.Bn_compute_blocks,
config.max_warmup_steps,
config.residual_diff_threshold,
config.max_continuous_cached_steps,
config.enable_dmd,
config.dmd_history,
config.dmd_rank,
config.dmd_svd_precision,
config.enable_taylorseer,
config.taylorseer_order,
config.num_inference_steps,
@@ -559,6 +673,11 @@ def enable_cache_on_dual_transformer(
if not primary_config.enabled:
return transformer, transformer_2
_assert_calibrator_exclusive(primary_config, label="primary")
_assert_calibrator_exclusive(secondary_config, label="secondary")
_assert_dmd_supported(primary_config, label="primary")
_assert_dmd_supported(secondary_config, label="secondary")
if primary_config.num_inference_steps is None:
raise ValueError(
"num_inference_steps is required for dual-transformer mode. "
@@ -589,15 +708,30 @@ def enable_cache_on_dual_transformer(
steps_computation_policy=secondary_config.steps_computation_policy,
)
# Build calibrator configs if TaylorSeer is enabled
# Build calibrator configs. DMD and TaylorSeer are mutually exclusive
# (validated above); DMD takes the calibrator slot when enabled.
primary_calibrator = None
if primary_config.enable_taylorseer:
if primary_config.enable_dmd:
primary_calibrator = DMDCalibratorConfig(
dmd_history=primary_config.dmd_history,
dmd_rank=primary_config.dmd_rank,
dmd_ridge=primary_config.dmd_ridge,
dmd_svd_precision=primary_config.dmd_svd_precision,
)
elif primary_config.enable_taylorseer:
primary_calibrator = TaylorSeerCalibratorConfig(
taylorseer_order=primary_config.taylorseer_order,
)
secondary_calibrator = None
if secondary_config.enable_taylorseer:
if secondary_config.enable_dmd:
secondary_calibrator = DMDCalibratorConfig(
dmd_history=secondary_config.dmd_history,
dmd_rank=secondary_config.dmd_rank,
dmd_ridge=secondary_config.dmd_ridge,
dmd_svd_precision=secondary_config.dmd_svd_precision,
)
elif secondary_config.enable_taylorseer:
secondary_calibrator = TaylorSeerCalibratorConfig(
taylorseer_order=secondary_config.taylorseer_order,
)
@@ -618,21 +752,31 @@ def enable_cache_on_dual_transformer(
model_name,
)
logger.info(
" Primary (transformer): Fn=%d, Bn=%d, W=%d, R=%.2f, MC=%d, TaylorSeer=%s",
" Primary (transformer): Fn=%d, Bn=%d, W=%d, R=%.2f, MC=%d, "
"DMD=%s (history=%d, rank=%d, svd=%s), TaylorSeer=%s",
primary_config.Fn_compute_blocks,
primary_config.Bn_compute_blocks,
primary_config.max_warmup_steps,
primary_config.residual_diff_threshold,
primary_config.max_continuous_cached_steps,
primary_config.enable_dmd,
primary_config.dmd_history,
primary_config.dmd_rank,
primary_config.dmd_svd_precision,
primary_config.enable_taylorseer,
)
logger.info(
" Secondary transformer: Fn=%d, Bn=%d, W=%d, R=%.2f, MC=%d, TaylorSeer=%s",
" Secondary (transformer_2): Fn=%d, Bn=%d, W=%d, R=%.2f, MC=%d, "
"DMD=%s (history=%d, rank=%d, svd=%s), TaylorSeer=%s",
secondary_config.Fn_compute_blocks,
secondary_config.Bn_compute_blocks,
secondary_config.max_warmup_steps,
secondary_config.residual_diff_threshold,
secondary_config.max_continuous_cached_steps,
secondary_config.enable_dmd,
secondary_config.dmd_history,
secondary_config.dmd_rank,
secondary_config.dmd_svd_precision,
secondary_config.enable_taylorseer,
)
@@ -975,6 +975,36 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
envs.SGLANG_CACHE_DIT_SECONDARY_TS_ORDER,
secondary=secondary,
),
enable_dmd=knob(
"enable_dmd",
envs.SGLANG_CACHE_DIT_DMD,
envs.SGLANG_CACHE_DIT_SECONDARY_DMD,
secondary=secondary,
),
dmd_history=knob(
"dmd_history",
envs.SGLANG_CACHE_DIT_DMD_HISTORY,
envs.SGLANG_CACHE_DIT_SECONDARY_DMD_HISTORY,
secondary=secondary,
),
dmd_rank=knob(
"dmd_rank",
envs.SGLANG_CACHE_DIT_DMD_RANK,
envs.SGLANG_CACHE_DIT_SECONDARY_DMD_RANK,
secondary=secondary,
),
dmd_ridge=knob(
"dmd_ridge",
envs.SGLANG_CACHE_DIT_DMD_RIDGE,
envs.SGLANG_CACHE_DIT_SECONDARY_DMD_RIDGE,
secondary=secondary,
),
dmd_svd_precision=knob(
"dmd_svd_precision",
envs.SGLANG_CACHE_DIT_DMD_SVD_PRECISION,
envs.SGLANG_CACHE_DIT_SECONDARY_DMD_SVD_PRECISION,
secondary=secondary,
),
num_inference_steps=num_inference_steps,
steps_computation_mask=steps_computation_mask,
steps_computation_policy=scm_policy,
@@ -1,5 +1,6 @@
import importlib
import importlib.util
import subprocess
import sys
import types
import unittest
@@ -22,6 +23,11 @@ class _FakeForwardPattern:
Pattern_3 = "Pattern_3"
class _FakeTaylorSeerCalibratorConfig:
def __init__(self, **kwargs):
self.kwargs = kwargs
def _install_cache_dit_stub():
cache_dit = types.ModuleType("cache_dit")
cache_dit.enable_calls = []
@@ -58,7 +64,10 @@ def _install_cache_dit_stub():
cache_dit.DBCacheConfig = _FakeDBCacheConfig
cache_dit.ForwardPattern = _FakeForwardPattern
cache_dit.ParamsModifier = object
cache_dit.TaylorSeerCalibratorConfig = object
# Not bare `object`: enable paths construct it with kwargs
# (TaylorSeerCalibratorConfig(taylorseer_order=...)) to fill the
# calibrator slot, and tests assert on the constructed instance.
cache_dit.TaylorSeerCalibratorConfig = _FakeTaylorSeerCalibratorConfig
block_adapters = types.ModuleType("cache_dit.caching.block_adapters")
@@ -70,10 +79,19 @@ def _install_cache_dit_stub():
return cls.supported
block_adapters.BlockAdapterRegister = _FakeBlockAdapterRegister
cache_dit.BlockAdapterRegister = _FakeBlockAdapterRegister
class _FakeDMDCalibratorConfig:
def __init__(self, **kwargs):
self.kwargs = kwargs
cache_dit.DMDCalibratorConfig = _FakeDMDCalibratorConfig
parallelism = types.ModuleType("cache_dit.parallelism")
parallelism.ParallelismBackend = object
parallelism.ParallelismConfig = object
cache_dit.ParallelismBackend = parallelism.ParallelismBackend
cache_dit.ParallelismConfig = parallelism.ParallelismConfig
return {
"cache_dit": cache_dit,
@@ -167,6 +185,39 @@ def _import_module_with_stub():
return module
def _import_module_with_legacy_stub():
# Full stub minus the symbols that only became top-level exports in
# cache-dit 1.5.0: forces cache_dit_integration onto its fallback path.
# All pre-1.5.0 releases pinned by sglang (1.3.0, 1.3.5) share the same
# import surface, so "1.3.0" represents every cache-dit < 1.5.0.
stub_modules = _install_cache_dit_stub()
stub_modules.update(_install_sglang_dependency_stubs())
stub_modules.update(_install_torch_stub())
cache_dit = stub_modules["cache_dit"]
cache_dit.__version__ = "1.3.0"
for name in (
"BlockAdapterRegister",
"DMDCalibratorConfig",
"ParallelismBackend",
"ParallelismConfig",
):
delattr(cache_dit, name)
module_path = (
Path(__file__).resolve().parents[2]
/ "runtime"
/ "cache"
/ "cache_dit_integration.py"
)
with patch.dict(sys.modules, stub_modules):
spec = importlib.util.spec_from_file_location(
"test_cache_dit_integration_legacy_target", module_path
)
module = importlib.util.module_from_spec(spec)
assert spec.loader is not None
spec.loader.exec_module(module)
return module, stub_modules
class TestCacheDitRefreshContext(unittest.TestCase):
def test_refresh_context_without_scm_preset_skips_steps_mask(self):
module = _import_module_with_stub()
@@ -318,11 +369,144 @@ class TestBuildCustomBlockAdapter(unittest.TestCase):
self.assertIs(returned, transformer)
adapter = transformer._sglang_cache_dit_adapter
self.assertIs(module.cache_dit.enable_calls[0]["target"], adapter)
self.assertIs(module.disable_cache_on_transformer(transformer), transformer)
self.assertEqual(module.cache_dit.disable_calls, [adapter])
self.assertFalse(hasattr(transformer, "_sglang_cache_dit_adapter"))
class TestCalibratorSelection(unittest.TestCase):
def _config(self, module, **kwargs):
return module.CacheDitConfig(enabled=True, num_inference_steps=28, **kwargs)
def test_dmd_takes_calibrator_slot(self):
module = _import_module_with_stub()
transformer = _make_transformer("AnyModel")
config = self._config(
module,
enable_dmd=True,
dmd_history=8,
dmd_rank=4,
dmd_ridge=1e-6,
dmd_svd_precision="high",
)
module.enable_cache_on_transformer(transformer, config)
calibrator = module.cache_dit.enable_calls[0]["calibrator_config"]
self.assertIsInstance(calibrator, module.DMDCalibratorConfig)
self.assertEqual(
calibrator.kwargs,
{
"dmd_history": 8,
"dmd_rank": 4,
"dmd_ridge": 1e-6,
"dmd_svd_precision": "high",
},
)
def test_both_calibrators_raise_on_transformer(self):
module = _import_module_with_stub()
transformer = _make_transformer("AnyModel")
config = self._config(module, enable_dmd=True, enable_taylorseer=True)
with self.assertRaisesRegex(ValueError, "mutually exclusive"):
module.enable_cache_on_transformer(transformer, config)
self.assertEqual(module.cache_dit.enable_calls, [])
def test_both_calibrators_raise_on_dual_transformer(self):
module = _import_module_with_stub()
transformer = _make_transformer("AnyModel")
transformer.blocks = ["block_0"]
transformer_2 = _make_transformer("AnyModel")
transformer_2.blocks = ["block_0"]
primary = self._config(module, enable_dmd=True, enable_taylorseer=True)
secondary = self._config(module)
with self.assertRaisesRegex(ValueError, "mutually exclusive"):
module.enable_cache_on_dual_transformer(
transformer,
transformer_2,
primary,
secondary,
model_name="wan2.2",
)
self.assertEqual(module.cache_dit.enable_calls, [])
class TestCacheDitLegacyFallback(unittest.TestCase):
"""cache-dit < 1.5.0 (CI base jobs ship 1.3.0): fallback import + DMD guard."""
def test_fallback_import_binds_registry_and_nulls_dmd(self):
module, stubs = _import_module_with_legacy_stub()
self.assertIsNone(module.DMDCalibratorConfig)
self.assertIs(
module.BlockAdapterRegister,
stubs["cache_dit.caching.block_adapters"].BlockAdapterRegister,
)
def test_enable_dmd_raises_clear_error(self):
module, _ = _import_module_with_legacy_stub()
config = module.CacheDitConfig(
enabled=True, enable_dmd=True, num_inference_steps=4
)
with self.assertRaisesRegex(ValueError, "cache-dit >= 1.5.0"):
module.enable_cache_on_transformer(_make_transformer("AnyModel"), config)
self.assertEqual(module.cache_dit.enable_calls, [])
def test_taylorseer_path_still_enables_cache(self):
module, _ = _import_module_with_legacy_stub()
transformer = _make_transformer("AnyModel")
config = module.CacheDitConfig(
enabled=True,
enable_dmd=False,
enable_taylorseer=True,
num_inference_steps=4,
)
result = module.enable_cache_on_transformer(transformer, config)
self.assertIs(result, transformer)
self.assertIsInstance(
module.cache_dit.enable_calls[0]["calibrator_config"],
_FakeTaylorSeerCalibratorConfig,
)
class TestCacheDitRealPackageBoundary(unittest.TestCase):
"""Real installed cache-dit, no stubs: catches stub/package drift — the
CI base jobs import this chain with cache-dit 1.3.0 while the diffusion
unit lane installs 1.5.1."""
@unittest.skipIf(
importlib.util.find_spec("cache_dit") is None, "cache_dit is not installed"
)
def test_import_chain_matches_installed_package(self):
script = (
"import cache_dit\n"
"import sglang.multimodal_gen.runtime.cache as cache_pkg\n"
"from sglang.multimodal_gen.runtime.cache import cache_dit_integration\n"
"print(hasattr(cache_dit, 'DMDCalibratorConfig'))\n"
"print(cache_dit_integration.DMDCalibratorConfig is not None)\n"
"print(cache_pkg.CacheDitConfig is cache_dit_integration.CacheDitConfig)\n"
)
result = subprocess.run(
[sys.executable, "-c", script], capture_output=True, text=True, timeout=300
)
self.assertEqual(
result.returncode, 0, f"real import chain failed:\n{result.stderr}"
)
has_top_dmd, bound_not_none, reexport_ok = result.stdout.strip().splitlines()[
-3:
]
self.assertEqual(
has_top_dmd,
str(bound_not_none == "True"),
"DMDCalibratorConfig binding mismatch vs installed package",
)
self.assertEqual(reexport_ok, "True")
if __name__ == "__main__":
unittest.main()
@@ -193,6 +193,39 @@ class TestPerRequestCacheDitTransitions(unittest.TestCase):
self.assertEqual(config.steps_computation_policy, "static")
self.assertEqual(config.num_inference_steps, 8)
def test_request_dmd_knobs_reach_cache_dit_config(self):
self.stage._maybe_enable_cache_dit(
8,
_batch(
enable_cache_dit=True,
cache_dit_params={
"enable_dmd": True,
"dmd_history": 8,
"dmd_svd_precision": "high",
},
),
)
(config,) = self.enable_calls
self.assertTrue(config.enable_dmd)
self.assertEqual(config.dmd_history, 8)
self.assertEqual(config.dmd_svd_precision, "high")
# untouched knobs keep their env defaults
self.assertEqual(config.dmd_rank, 0)
def test_secondary_inherits_request_primary_dmd(self):
self.stage._cache_dit_request_overrides = resolve_cache_dit_request_overrides(
{"enable_dmd": True, "secondary": {"dmd_rank": 4}}
)
primary = self.stage._build_cache_dit_config(
10, steps_computation_mask=None, scm_policy="dynamic"
)
secondary = self.stage._build_cache_dit_config(
10, steps_computation_mask=None, scm_policy="dynamic", secondary=True
)
self.assertTrue(primary.enable_dmd)
self.assertTrue(secondary.enable_dmd) # inherited from primary
self.assertEqual(secondary.dmd_rank, 4)
def test_invalid_request_params_raise(self):
with self.assertRaisesRegex(ValueError, "Unknown cache_dit_params keys"):
self.stage._maybe_enable_cache_dit(