[diffusion][cache-dit] support Krea-2 + run-driven has_separate_cfg (#29688)
Co-authored-by: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 4.8
parent
c6a7c98ae4
commit
a531d81c19
@@ -88,6 +88,129 @@ sglang generate --model-path krea/Krea-2-Turbo \
|
|||||||
|
|
||||||
### 4.2 Advanced Usage
|
### 4.2 Advanced Usage
|
||||||
|
|
||||||
|
#### 4.2.1 Cache-DiT Acceleration
|
||||||
|
|
||||||
|
SGLang integrates [Cache-DiT](https://github.com/vipshop/cache-dit), a caching acceleration engine for Diffusion Transformers (DiT), to speed up inference with minimal quality loss. Enable it by setting `SGLANG_CACHE_DIT_ENABLED=true`. For more details, see the SGLang Cache-DiT [documentation](/docs/sglang-diffusion/cache_dit).
|
||||||
|
|
||||||
|
Cache-DiT works for **both** Krea-2 variants with no extra configuration: SGLang tracks each request's classifier-free-guidance mode, so Krea-2-Turbo (no CFG, `guidance_scale = 1.0`) and Krea-2-Raw (CFG, `guidance_scale ≈ 4.5`) both cache correctly and automatically.
|
||||||
|
|
||||||
|
**Basic Usage**
|
||||||
|
|
||||||
|
```bash Command
|
||||||
|
SGLANG_CACHE_DIT_ENABLED=true sglang serve \
|
||||||
|
--model-path krea/Krea-2-Turbo \
|
||||||
|
--num-gpus 1 \
|
||||||
|
--port 30000
|
||||||
|
```
|
||||||
|
|
||||||
|
Measured per-image denoise speedup with the default cache settings (NVIDIA H200, 1024x1024, seed 0):
|
||||||
|
|
||||||
|
| Variant | Inference steps | Denoise (no cache → cache) | Speedup |
|
||||||
|
| :--- | :--- | :--- | :--- |
|
||||||
|
| Krea-2-Turbo (no CFG) | 8 | 1.27s → 0.92s | ~1.4x |
|
||||||
|
| Krea-2-Raw (CFG 4.5) | 50 | 18.0s → 6.3s | ~2.9x |
|
||||||
|
|
||||||
|
Caching has the most headroom on Raw's longer schedule; the 8-step distilled Turbo has only a few cacheable steps after warmup.
|
||||||
|
|
||||||
|
**Advanced Usage**
|
||||||
|
|
||||||
|
- DBCache Parameters: DBCache controls block-level caching behavior:
|
||||||
|
|
||||||
|
<table style={{width: "100%", borderCollapse: "collapse", tableLayout: "fixed"}}>
|
||||||
|
<colgroup>
|
||||||
|
<col style={{width: "25.0%"}} />
|
||||||
|
<col style={{width: "25.0%"}} />
|
||||||
|
<col style={{width: "25.0%"}} />
|
||||||
|
<col style={{width: "25.0%"}} />
|
||||||
|
</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)"}}>Fn</td>
|
||||||
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>`SGLANG_CACHE_DIT_FN`</td>
|
||||||
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>1</td>
|
||||||
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Number of first blocks to always compute</td>
|
||||||
|
</tr>
|
||||||
|
<tr>
|
||||||
|
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>Bn</td>
|
||||||
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>`SGLANG_CACHE_DIT_BN`</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)"}}>Number of last blocks to always compute</td>
|
||||||
|
</tr>
|
||||||
|
<tr>
|
||||||
|
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>W</td>
|
||||||
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>`SGLANG_CACHE_DIT_WARMUP`</td>
|
||||||
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>4</td>
|
||||||
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Warmup steps before caching starts</td>
|
||||||
|
</tr>
|
||||||
|
<tr>
|
||||||
|
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>R</td>
|
||||||
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>`SGLANG_CACHE_DIT_RDT`</td>
|
||||||
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>0.24</td>
|
||||||
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Residual difference threshold</td>
|
||||||
|
</tr>
|
||||||
|
<tr>
|
||||||
|
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>MC</td>
|
||||||
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>`SGLANG_CACHE_DIT_MC`</td>
|
||||||
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>3</td>
|
||||||
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Maximum continuous cached steps</td>
|
||||||
|
</tr>
|
||||||
|
</tbody>
|
||||||
|
</table>
|
||||||
|
- TaylorSeer Configuration: TaylorSeer improves caching accuracy using Taylor expansion (best suited to the longer Raw schedule; not recommended for the 8-step Turbo):
|
||||||
|
|
||||||
|
<table style={{width: "100%", borderCollapse: "collapse", tableLayout: "fixed"}}>
|
||||||
|
<colgroup>
|
||||||
|
<col style={{width: "25.0%"}} />
|
||||||
|
<col style={{width: "25.0%"}} />
|
||||||
|
<col style={{width: "25.0%"}} />
|
||||||
|
<col style={{width: "25.0%"}} />
|
||||||
|
</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_TAYLORSEER`</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 TaylorSeer calibrator</td>
|
||||||
|
</tr>
|
||||||
|
<tr>
|
||||||
|
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>Order</td>
|
||||||
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>`SGLANG_CACHE_DIT_TS_ORDER`</td>
|
||||||
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>1</td>
|
||||||
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Taylor expansion order (1 or 2)</td>
|
||||||
|
</tr>
|
||||||
|
</tbody>
|
||||||
|
</table>
|
||||||
|
|
||||||
|
Combined Configuration Example (Krea-2-Raw, default cache settings shown explicitly):
|
||||||
|
|
||||||
|
```bash Command
|
||||||
|
SGLANG_CACHE_DIT_ENABLED=true \
|
||||||
|
SGLANG_CACHE_DIT_FN=1 \
|
||||||
|
SGLANG_CACHE_DIT_BN=0 \
|
||||||
|
SGLANG_CACHE_DIT_WARMUP=4 \
|
||||||
|
SGLANG_CACHE_DIT_RDT=0.24 \
|
||||||
|
SGLANG_CACHE_DIT_MC=3 \
|
||||||
|
sglang serve --model-path krea/Krea-2-Raw
|
||||||
|
```
|
||||||
|
|
||||||
|
#### 4.2.2 Memory & CPU Offload
|
||||||
|
|
||||||
Krea-2's DiT is ~24 GB in bf16 (the bulk of the model). On memory-constrained GPUs you can keep less of it resident:
|
Krea-2's DiT is ~24 GB in bf16 (the bulk of the model). On memory-constrained GPUs you can keep less of it resident:
|
||||||
|
|
||||||
- `--dit-layerwise-offload`: stream the DiT's transformer blocks layer-by-layer with async host-to-device prefetch overlap, so only a small working set stays on the GPU. This is the primary way to fit Krea-2 on a single consumer / 32 GB-class card, at a modest latency cost. Tune the memory/latency trade-off with `--dit-offload-prefetch-size` (`0.0` prefetches one layer for the lowest memory; larger values prefetch more layers -- faster but more memory).
|
- `--dit-layerwise-offload`: stream the DiT's transformer blocks layer-by-layer with async host-to-device prefetch overlap, so only a small working set stays on the GPU. This is the primary way to fit Krea-2 on a single consumer / 32 GB-class card, at a modest latency cost. Tune the memory/latency trade-off with `--dit-offload-prefetch-size` (`0.0` prefetches one layer for the lowest memory; larger values prefetch more layers -- faster but more memory).
|
||||||
|
|||||||
@@ -223,25 +223,27 @@ class CacheDitConfig:
|
|||||||
|
|
||||||
|
|
||||||
# Custom BlockAdapter for DiT models absent from cache-dit's BlockAdapterRegister.
|
# Custom BlockAdapter for DiT models absent from cache-dit's BlockAdapterRegister.
|
||||||
# Value: (blocks attr, forward_pattern, has_separate_cfg). forward_pattern must
|
# Value: (blocks attr, forward_pattern). forward_pattern must
|
||||||
# match the block's forward signature (see cache_dit.ForwardPattern; e.g., ERNIE
|
# match the block's forward signature (see cache_dit.ForwardPattern; e.g., ERNIE
|
||||||
# uses Pattern_3). has_separate_cfg=True aligns cache-dit's step counter for
|
# uses Pattern_3). has_separate_cfg follows the run (passed by
|
||||||
# sequential CFG (two forwards per step); cache-dit auto-resolves the remaining
|
# enable_cache_on_transformer); cache-dit auto-resolves the remaining
|
||||||
# fields.
|
# fields.
|
||||||
_CUSTOM_BLOCK_ADAPTER_SPECS: dict[str, tuple[str, ForwardPattern, bool]] = {
|
_CUSTOM_BLOCK_ADAPTER_SPECS: dict[str, tuple[str, ForwardPattern]] = {
|
||||||
"ErnieImageTransformer2DModel": ("layers", ForwardPattern.Pattern_3, True),
|
"ErnieImageTransformer2DModel": ("layers", ForwardPattern.Pattern_3),
|
||||||
|
"Krea2Transformer2DModel": ("transformer_blocks", ForwardPattern.Pattern_3),
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
def _build_custom_block_adapter(
|
def _build_custom_block_adapter(
|
||||||
transformer: torch.nn.Module,
|
transformer: torch.nn.Module,
|
||||||
|
has_separate_cfg: bool = False,
|
||||||
) -> Optional[BlockAdapter]:
|
) -> Optional[BlockAdapter]:
|
||||||
"""Build a manual BlockAdapter for a model absent from cache-dit's registry,
|
"""Build a manual BlockAdapter for a model absent from cache-dit's registry,
|
||||||
or None if the class is unknown."""
|
or None if the class is unknown."""
|
||||||
spec = _CUSTOM_BLOCK_ADAPTER_SPECS.get(transformer.__class__.__name__)
|
spec = _CUSTOM_BLOCK_ADAPTER_SPECS.get(transformer.__class__.__name__)
|
||||||
if spec is None:
|
if spec is None:
|
||||||
return None
|
return None
|
||||||
blocks_attr, forward_pattern, has_separate_cfg = spec
|
blocks_attr, forward_pattern = spec
|
||||||
blocks = getattr(transformer, blocks_attr, None)
|
blocks = getattr(transformer, blocks_attr, None)
|
||||||
if blocks is None:
|
if blocks is None:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
@@ -262,6 +264,7 @@ def enable_cache_on_transformer(
|
|||||||
model_name: str = "transformer",
|
model_name: str = "transformer",
|
||||||
sp_group: Optional[torch.distributed.ProcessGroup] = None,
|
sp_group: Optional[torch.distributed.ProcessGroup] = None,
|
||||||
tp_group: Optional[torch.distributed.ProcessGroup] = None,
|
tp_group: Optional[torch.distributed.ProcessGroup] = None,
|
||||||
|
has_separate_cfg: bool = False,
|
||||||
) -> torch.nn.Module:
|
) -> torch.nn.Module:
|
||||||
"""Enable cache-dit on a transformer module, by wrapping the module with cache-dit
|
"""Enable cache-dit on a transformer module, by wrapping the module with cache-dit
|
||||||
|
|
||||||
@@ -272,6 +275,9 @@ def enable_cache_on_transformer(
|
|||||||
model_name: Name of the model for logging purposes.
|
model_name: Name of the model for logging purposes.
|
||||||
sp_group: Sequence parallel process group (for Ulysses/Ring).
|
sp_group: Sequence parallel process group (for Ulysses/Ring).
|
||||||
tp_group: Tensor parallel process group.
|
tp_group: Tensor parallel process group.
|
||||||
|
has_separate_cfg: Whether the run issues separate conditional/unconditional
|
||||||
|
passes per step (CFG). Used by custom adapters (ERNIE, Krea-2); a
|
||||||
|
mismatch only disables caching, never corrupts output.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
if not config.enabled:
|
if not config.enabled:
|
||||||
@@ -288,7 +294,9 @@ def enable_cache_on_transformer(
|
|||||||
# _build_custom_block_adapter).
|
# _build_custom_block_adapter).
|
||||||
custom_adapter = None
|
custom_adapter = None
|
||||||
if not BlockAdapterRegister.is_supported(transformer):
|
if not BlockAdapterRegister.is_supported(transformer):
|
||||||
custom_adapter = _build_custom_block_adapter(transformer)
|
custom_adapter = _build_custom_block_adapter(
|
||||||
|
transformer, has_separate_cfg=has_separate_cfg
|
||||||
|
)
|
||||||
if custom_adapter is None:
|
if custom_adapter is None:
|
||||||
transformer_cls_name = transformer.__class__.__name__
|
transformer_cls_name = transformer.__class__.__name__
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
|
|||||||
@@ -445,7 +445,7 @@ class SingleStreamBlock(nn.Module):
|
|||||||
|
|
||||||
def forward(
|
def forward(
|
||||||
self,
|
self,
|
||||||
x: Tensor,
|
hidden_states: Tensor,
|
||||||
vec: Tensor,
|
vec: Tensor,
|
||||||
freqs: Tensor,
|
freqs: Tensor,
|
||||||
key_mask: Tensor | None = None,
|
key_mask: Tensor | None = None,
|
||||||
@@ -455,20 +455,28 @@ class SingleStreamBlock(nn.Module):
|
|||||||
prescale, preshift, pregate, postscale, postshift, postgate = mod.chunk(
|
prescale, preshift, pregate, postscale, postshift, postgate = mod.chunk(
|
||||||
6, dim=-1
|
6, dim=-1
|
||||||
)
|
)
|
||||||
x = x + pregate * self.attn(
|
hidden_states = hidden_states + pregate * self.attn(
|
||||||
norm_scale_shift(
|
norm_scale_shift(
|
||||||
x, self.norm1.weight + 1, prescale, preshift, self.norm1.eps
|
hidden_states,
|
||||||
|
self.norm1.weight + 1,
|
||||||
|
prescale,
|
||||||
|
preshift,
|
||||||
|
self.norm1.eps,
|
||||||
),
|
),
|
||||||
freqs,
|
freqs,
|
||||||
key_mask,
|
key_mask,
|
||||||
mask_meta,
|
mask_meta,
|
||||||
)
|
)
|
||||||
x = x + postgate * self.ff(
|
hidden_states = hidden_states + postgate * self.ff(
|
||||||
norm_scale_shift(
|
norm_scale_shift(
|
||||||
x, self.norm2.weight + 1, postscale, postshift, self.norm2.eps
|
hidden_states,
|
||||||
|
self.norm2.weight + 1,
|
||||||
|
postscale,
|
||||||
|
postshift,
|
||||||
|
self.norm2.eps,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
return x
|
return hidden_states
|
||||||
|
|
||||||
|
|
||||||
# --------------------------------------------------------------------------- #
|
# --------------------------------------------------------------------------- #
|
||||||
|
|||||||
@@ -545,6 +545,7 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
|
|||||||
model_name="transformer",
|
model_name="transformer",
|
||||||
sp_group=sp_group,
|
sp_group=sp_group,
|
||||||
tp_group=tp_group,
|
tp_group=tp_group,
|
||||||
|
has_separate_cfg=batch.do_classifier_free_guidance,
|
||||||
)
|
)
|
||||||
logger.info(
|
logger.info(
|
||||||
"cache-dit enabled on transformer (steps=%d, Fn=%d, Bn=%d, rdt=%.3f)",
|
"cache-dit enabled on transformer (steps=%d, Fn=%d, Bn=%d, rdt=%.3f)",
|
||||||
|
|||||||
@@ -230,7 +230,7 @@ class TestBuildCustomBlockAdapter(unittest.TestCase):
|
|||||||
blocks = ["block_0", "block_1"]
|
blocks = ["block_0", "block_1"]
|
||||||
transformer = _make_transformer("ErnieImageTransformer2DModel", blocks)
|
transformer = _make_transformer("ErnieImageTransformer2DModel", blocks)
|
||||||
|
|
||||||
adapter = module._build_custom_block_adapter(transformer)
|
adapter = module._build_custom_block_adapter(transformer, has_separate_cfg=True)
|
||||||
|
|
||||||
self.assertIsNotNone(adapter)
|
self.assertIsNotNone(adapter)
|
||||||
self.assertEqual(adapter.blocks, blocks)
|
self.assertEqual(adapter.blocks, blocks)
|
||||||
@@ -250,6 +250,28 @@ class TestBuildCustomBlockAdapter(unittest.TestCase):
|
|||||||
with self.assertRaises(ValueError):
|
with self.assertRaises(ValueError):
|
||||||
module._build_custom_block_adapter(transformer)
|
module._build_custom_block_adapter(transformer)
|
||||||
|
|
||||||
|
def test_has_separate_cfg_follows_runtime(self):
|
||||||
|
# No model pins the mode; has_separate_cfg always follows the run's CFG mode
|
||||||
|
# (Krea-2 Raw -> True, Krea-2 Turbo -> False).
|
||||||
|
module = _import_module_with_stub()
|
||||||
|
blocks = ["block_0", "block_1"]
|
||||||
|
|
||||||
|
transformer_raw = _make_transformer("Krea2Transformer2DModel")
|
||||||
|
transformer_raw.transformer_blocks = blocks
|
||||||
|
adapter_raw = module._build_custom_block_adapter(
|
||||||
|
transformer_raw, has_separate_cfg=True
|
||||||
|
)
|
||||||
|
self.assertEqual(adapter_raw.blocks, blocks)
|
||||||
|
self.assertEqual(adapter_raw.forward_pattern, "Pattern_3")
|
||||||
|
self.assertTrue(adapter_raw.has_separate_cfg)
|
||||||
|
|
||||||
|
transformer_turbo = _make_transformer("Krea2Transformer2DModel")
|
||||||
|
transformer_turbo.transformer_blocks = blocks
|
||||||
|
adapter_turbo = module._build_custom_block_adapter(
|
||||||
|
transformer_turbo, has_separate_cfg=False
|
||||||
|
)
|
||||||
|
self.assertFalse(adapter_turbo.has_separate_cfg)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
Reference in New Issue
Block a user