Prhub

#29688 [diffusion][cache-dit] support Krea-2 + run-driven `has_separate_cfg`

原始 PR 作者 AgainstEntropy 合并时间 2026-06-30 23:22 文件变更 5 提交数 2 评论 2 代码增减 +176 / -14

执行摘要

支持 Krea-2 Cache-DiT 加速并动态检测 CFG 模式

PR 要解决两个问题:1)为 Krea-2 (Turbo 和 Raw) 提供 Cache-DiT 支持;2)修正之前将所有模型变体强制赋予 has_separate_cfg=True 的硬编码设计。例如 baidu/ERNIE-Image-Turbokrea/Krea-2-Turbo 不使用 CFG,但硬编码为 True 会导致 SGLANG_CACHE_DIT_ENABLED=true 对这些蒸馏模型几乎不生效(实测仅 1.01x 加速)。

推荐阅读此 PR,它展示了如何将一个新模型家族(Krea-2)优雅地接入缓存系统,并通过运行时参数解决模型变体间的 CFG 差异问题。设计决策(动态探测 vs 静态注册)值得借鉴,尤其是解决蒸馏模型缓存无效的通用思路。

讨论亮点

PR 获得两位 reviewer (mickqian, zijiexia) 的快速批准,无公开讨论记录。

实现拆解

  1. 扩展自定义 BlockAdapter 注册表:在 cache_dit_integration.py_CUSTOM_BLOCK_ADAPTER_SPECS 中添加 Krea2Transformer2DModel,指定其 block 属性名为 transformer_blocks,forward 模式为 Pattern_3
  2. 移除静态 has_separate_cfg 字段:将 _CUSTOM_BLOCK_ADAPTER_SPECS 的值从三元组变为二元组,不再硬编码 CFG 模式;改为在 _build_custom_block_adapterenable_cache_on_transformer 函数签名中增加 has_separate_cfg: bool 参数,由调用者传入。
  3. 在去噪阶段注入运行时 CFG 状态:在 denoising.py_maybe_enable_cache_dit 中,调用 enable_cache_on_transformer 时传入 has_separate_cfg=batch.do_classifier_free_guidance,实现每批次动态适配。
  4. 统一 Krea-2 SingleStreamBlock 前向参数名:将 krea2.pySingleStreamBlock.forward 的参数 x 重命名为 hidden_states,与 cache-dit 的 Pattern_3 期望命名一致。
  5. 补充文档与测试:在 Krea-2 的 cookbook 中新增 Cache-DiT 使用说明及性能数据;添加单元测试 test_has_separate_cfg_follows_runtime 验证 has_separate_cfg 正确随运行时参数变化。
文件 模块 状态 重要度
python/sglang/multimodal_gen/runtime/cache/cache_dit_integration.py 缓存层 modified 6.58
python/sglang/multimodal_gen/runtime/models/dits/krea2.py 模型适配 modified 6.18
python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py 去噪阶段 modified 4.96
python/sglang/multimodal_gen/test/unit/test_cache_dit_integration.py 测试 modified 5.35
docs_new/cookbook/diffusion/Krea/Krea-2.mdx 文档 modified 4.01

关键符号

_build_custom_block_adapter enable_cache_on_transformer SingleStreamBlock.forward _maybe_enable_cache_dit

关键源码片段

python/sglang/multimodal_gen/runtime/cache/cache_dit_integration.py core-logic

核心变更:修改自定义 BlockAdapter 注册表格式,移除静态 has_separate_cfg,改为运行时参数;添加 Krea-2 注册;修改 _build_custom_block_adapter 和 enable_cache_on_transformer 签名。

def _build_custom_block_adapter(
    transformer: torch.nn.Module,
    has_separate_cfg: bool = False, # 新增参数,由调用者(去噪阶段)根据 batch 的 CFG 模式传入
) -> Optional[BlockAdapter]:
    """为未注册到 cache-dit 的模型手动构建 BlockAdapter。"""
    spec = _CUSTOM_BLOCK_ADAPTER_SPECS.get(transformer.__class__.__name__)
    if spec is None:
        return None
    blocks_attr, forward_pattern = spec # 原三元组去掉 has_separate_cfg
    blocks = getattr(transformer, blocks_attr, None)
    if blocks is None:
        raise ValueError(...)
    return BlockAdapter(
        transformer=transformer,
        blocks=blocks,
        forward_pattern=forward_pattern,
        has_separate_cfg=has_separate_cfg, # 现在从参数动态传入
    )
python/sglang/multimodal_gen/runtime/models/dits/krea2.py data-contract

Krea-2 模型定义:修改 SingleStreamBlock.forward 的参数名 x 为 hidden_states,与 cache-dit 的 Pattern_3 一致;新增 Krea2Transformer2DModel 类(已存在,但相关适配已完成)。

class SingleStreamBlock(nn.Module):
    def forward(
        self,
        hidden_states: Tensor, # 原参数名为 x,现改为 hidden_states,与 cache-dit 的 Pattern_3 预期匹配
        vec: Tensor,
        freqs: Tensor,
        key_mask: Tensor | None = None,
        mask_meta: dict | None = None,
    ) -> Tensor:
        # ... 内部所有 x 引用改为 hidden_states
        hidden_states = hidden_states + pregate * self.attn(...)
        hidden_states = hidden_states + postgate * self.ff(...)
        return hidden_states
python/sglang/multimodal_gen/test/unit/test_cache_dit_integration.py test-coverage

新增 test_has_separate_cfg_follows_runtime 测试,验证 Krea-2 Raw 和 Turbo 使用不同 has_separate_cfg 值时 adapter 正确设置;同时修改已有测试以适配新函数签名。

def test_has_separate_cfg_follows_runtime(self):
    # 验证 has_separate_cfg 由运行时参数决定,而非模型类硬编码
    module = _import_module_with_stub()
    blocks = ["block_0", "block_1"]
​
    # 模拟 Krea-2 Raw(CFG=True)
    transformer_raw = _make_transformer("Krea2Transformer2DModel")
    transformer_raw.transformer_blocks = blocks
    adapter_raw = module._build_custom_block_adapter(
        transformer_raw, has_separate_cfg=True
    )
    self.assertTrue(adapter_raw.has_separate_cfg)
​
    # 模拟 Krea-2 Turbo(CFG=False)
    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)

评论区精华

没有提炼出高价值讨论线程

当前评论区没有形成足够清晰的争议点或结论,后续有更多讨论时会体现在这里。

风险与影响

核心风险在于 has_separate_cfg 的传递路径:若 batch.do_classifier_free_guidance 在特定调度场景下为空或错误,可能导致缓存行为偏离预期(但 PR 备注指出 mismatch 仅禁用缓存,不会损坏输出)。另外,修改了 _CUSTOM_BLOCK_ADAPTER_SPECS 的合约(从三元组变为二元组),任何外部依赖该字典的代码(如有)需同步更新。最后,文档中新增的大量表格和配置示例可能存在渲染问题,但无功能风险。

用户影响:使用 Krea-2 或 ERNIE-Image 系列(含 Turbo 变体)的用户可直接受益于 Cache-DiT,加速显著(1.4x~2.9x)。系统影响:Cache-DiT 集成层现在支持按批次动态设置 CFG 模式,为未来更多模型族(如 raw/turbo 对)提供了可复用机制。团队影响:降低了维护成本——无需为每个模型变体单独注册 has_separate_cfg

核心路径变更 CFG 模式传递依赖 字典合约变更

关联 Issue

未识别关联 Issue

当前没有检测到明确关联的 Issue 链接,后续同步到相关引用后会出现在这里。

完整报告

参与讨论