diff --git a/.gitignore b/.gitignore index 581f2c091ddd..fbf422663bc3 100644 --- a/.gitignore +++ b/.gitignore @@ -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/ diff --git a/docs/docs/sglang-diffusion/cache_dit.mdx b/docs/docs/sglang-diffusion/cache_dit.mdx index 1cffcd6f9a84..d26ac88b9ffa 100644 --- a/docs/docs/sglang-diffusion/cache_dit.mdx +++ b/docs/docs/sglang-diffusion/cache_dit.mdx @@ -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 5–6 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== --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: +### 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: + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + +
ParameterEnv VariableDefaultDescription
Enable`SGLANG_CACHE_DIT_DMD`falseEnable the DMD calibrator
History`SGLANG_CACHE_DIT_DMD_HISTORY`6Snapshot window length; 5-6 typical. Needs >= 4 uniformly spaced snapshots, otherwise DMD falls back to TaylorSeer
Rank`SGLANG_CACHE_DIT_DMD_RANK`0SVD truncation rank; 0 = automatic (drop modes below 1e-4 of the leading singular value)
Ridge`SGLANG_CACHE_DIT_DMD_RIDGE`1e-8Tikhonov regularization added to the inverted singular values
SVD Precision`SGLANG_CACHE_DIT_DMD_SVD_PRECISION`mediumSVD precision: "low", "medium" or "high"
+ +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) diff --git a/docs/docs/sglang-diffusion/environment_variables.mdx b/docs/docs/sglang-diffusion/environment_variables.mdx index 89ece541946d..fbf899a8b9b3 100644 --- a/docs/docs/sglang-diffusion/environment_variables.mdx +++ b/docs/docs/sglang-diffusion/environment_variables.mdx @@ -253,6 +253,31 @@ See [cache-dit documentation](./cache_dit) for detailed configuration. 1 TaylorSeer order (1 or 2) + + `SGLANG_CACHE_DIT_DMD` + false + Enable the DMD (Dynamic Mode Decomposition) calibrator (mutually exclusive with TaylorSeer) + + + `SGLANG_CACHE_DIT_DMD_HISTORY` + 6 + DMD snapshot window length (5-6 typical; needs >= 4 uniformly spaced snapshots, otherwise DMD falls back to TaylorSeer) + + + `SGLANG_CACHE_DIT_DMD_RANK` + 0 + DMD SVD truncation rank (0 = automatic) + + + `SGLANG_CACHE_DIT_DMD_RIDGE` + 1e-8 + DMD Tikhonov regularization added to the inverted singular values + + + `SGLANG_CACHE_DIT_DMD_SVD_PRECISION` + medium + DMD SVD precision (low/medium/high) + `SGLANG_CACHE_DIT_SCM_PRESET` none @@ -329,6 +354,31 @@ For dual-transformer models (e.g., Wan2.2 with high/low-noise experts), these va (from primary) TaylorSeer order (1 or 2) + + SGLANG_CACHE_DIT_SECONDARY_DMD + (from primary) + Enable the DMD calibrator (mutually exclusive with TaylorSeer) + + + SGLANG_CACHE_DIT_SECONDARY_DMD_HISTORY + (from primary) + DMD snapshot window length + + + SGLANG_CACHE_DIT_SECONDARY_DMD_RANK + (from primary) + DMD SVD truncation rank (0 = automatic) + + + SGLANG_CACHE_DIT_SECONDARY_DMD_RIDGE + (from primary) + DMD Tikhonov regularization term + + + SGLANG_CACHE_DIT_SECONDARY_DMD_SVD_PRECISION + (from primary) + DMD SVD precision (low/medium/high) + diff --git a/python/pyproject.toml b/python/pyproject.toml index 0527b548f366..f0ba0e642321 100755 --- a/python/pyproject.toml +++ b/python/pyproject.toml @@ -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", diff --git a/python/sglang/multimodal_gen/.claude/skills/sglang-diffusion-cache-dit/SKILL.md b/python/sglang/multimodal_gen/.claude/skills/sglang-diffusion-cache-dit/SKILL.md new file mode 100644 index 000000000000..f127d44b193c --- /dev/null +++ b/python/sglang/multimodal_gen/.claude/skills/sglang-diffusion-cache-dit/SKILL.md @@ -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 repo root, `` = local cache-dit repo root (if not present locally, clone it: `git clone https://github.com/vipshop/cache-dit`), `` = 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 toolkit root (typically `/usr/local/cuda`), `` = 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=); + other GPUs may be busy with other jobs. + 3. The test artifact directory /.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 -- ` + (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 `` exist? The yaml configs in §4 come from + `/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 +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== # or cache-dit-cu13== --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 +export 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 ` 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` (/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/.log && nvidia-smi -i `) 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 /.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=/.tmp/{task} +FLUX_DIR= # model actually under test — ask the user if not specified +CFG=/examples/configs +OUT=$BASE/outputs; LOGS=$BASE/logs; mkdir -p "$OUT" "$LOGS" +export CUDA_VISIBLE_DEVICES= +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 +# 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 +export FLUX_DIR= # model actually under test — ask the user if not specified +export CUDA_VISIBLE_DEVICES= +mkdir -p .tmp/{task}/outputs +# SVDQuant cases require: pip install cache-dit-cu13== --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=/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 .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 | +| `/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`. +- `/.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) +- `/docs/docs/sglang-diffusion/cache_dit.mdx` — Cache-DiT usage guide (env vars, per-request knobs, yaml configs) +- `/docs/docs/sglang-diffusion/caching-acceleration.mdx` — caching-acceleration guide (DBCache/DMD/TaylorSeer combinations, acceleration matrix) +- `/docs/docs/sglang-diffusion/` — the whole SGLang Diffusion docs directory; auxiliary reference for related topics (quantization, parallelism, installation, performance) diff --git a/python/sglang/multimodal_gen/.claude/skills/sglang-diffusion-cache-dit/references/block_adapter.md b/python/sglang/multimodal_gen/.claude/skills/sglang-diffusion-cache-dit/references/block_adapter.md new file mode 100644 index 000000000000..61bc9ff39a33 --- /dev/null +++ b/python/sglang/multimodal_gen/.claude/skills/sglang-diffusion-cache-dit/references/block_adapter.md @@ -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 `/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 (`/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 (`/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`** (`/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`** (`/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 `/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: `/src/cache_dit/caching/` diff --git a/python/sglang/multimodal_gen/envs.py b/python/sglang/multimodal_gen/envs.py index 1295e03c85e2..93fea8168540 100644 --- a/python/sglang/multimodal_gen/envs.py +++ b/python/sglang/multimodal_gen/envs.py @@ -69,6 +69,11 @@ 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 @@ -81,6 +86,11 @@ 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 @@ -374,6 +384,19 @@ def _getter(): "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") @@ -450,6 +473,10 @@ def _getter(): ("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"), ] @@ -484,6 +511,17 @@ def _secondary_taylorseer_getter(): ) +# 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 diff --git a/python/sglang/multimodal_gen/runtime/cache/cache_dit_integration.py b/python/sglang/multimodal_gen/runtime/cache/cache_dit_integration.py index afbbf995e391..72b8ed739d02 100644 --- a/python/sglang/multimodal_gen/runtime/cache/cache_dit_integration.py +++ b/python/sglang/multimodal_gen/runtime/cache/cache_dit_integration.py @@ -22,16 +22,36 @@ 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 @@ def get_scm_mask( "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, ) diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py index 4fa862cc904e..b0eec2d9d7b9 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py @@ -978,6 +978,36 @@ def _build_cache_dit_config( 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, diff --git a/python/sglang/multimodal_gen/test/unit/test_cache_dit_integration.py b/python/sglang/multimodal_gen/test/unit/test_cache_dit_integration.py index eeabcec06c86..b5ae4b24a97d 100644 --- a/python/sglang/multimodal_gen/test/unit/test_cache_dit_integration.py +++ b/python/sglang/multimodal_gen/test/unit/test_cache_dit_integration.py @@ -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 steps_mask(*, mask_policy, total_steps): 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 is_supported(cls, _transformer): 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 @@ def test_custom_adapter_is_retained_until_disable(self): 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() diff --git a/python/sglang/multimodal_gen/test/unit/test_cache_dit_per_request.py b/python/sglang/multimodal_gen/test/unit/test_cache_dit_per_request.py index 86e8c03b861d..d2d985799b7d 100644 --- a/python/sglang/multimodal_gen/test/unit/test_cache_dit_per_request.py +++ b/python/sglang/multimodal_gen/test/unit/test_cache_dit_per_request.py @@ -193,6 +193,39 @@ def test_request_knobs_reach_cache_dit_config(self): 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(