From 92c29d43ac5f9a316e3d5470c1465591f4919c2d Mon Sep 17 00:00:00 2001 From: Xiaoyu Zhang <35585791+BBuf@users.noreply.github.com> Date: Mon, 15 Dec 2025 19:15:51 +0800 Subject: [PATCH] [diffusion] fix: cache dit with parallel (#15163) Co-authored-by: Mick --- .../pipelines_core/stages/denoising.py | 39 ++++- .../multimodal_gen/runtime/server_args.py | 10 ++ .../runtime/utils/cache_dit_integration.py | 135 ++++++++++++++++++ 3 files changed, 177 insertions(+), 7 deletions(-) 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 ac1b40bb1..81af951c3 100755 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py @@ -192,7 +192,11 @@ class DenoisingStage(PipelineStage): if not envs.SGLANG_CACHE_DIT_ENABLED: return - from sglang.multimodal_gen.runtime.distributed import get_world_size + from sglang.multimodal_gen.runtime.distributed import ( + get_sp_group, + get_tp_group, + get_world_size, + ) from sglang.multimodal_gen.runtime.utils.cache_dit_integration import ( CacheDitConfig, enable_cache_on_dual_transformer, @@ -200,13 +204,30 @@ class DenoisingStage(PipelineStage): get_scm_mask, ) - if get_world_size() > 1: - logger.warning( - "cache-dit is disabled in distributed environment (world_size=%d). " - "Distributed support will be added in a future version.", - get_world_size(), + world_size = get_world_size() + parallelized = world_size > 1 + + sp_group = None + tp_group = None + if parallelized: + sp_group_candidate = get_sp_group() + tp_group_candidate = get_tp_group() + + sp_world_size = sp_group_candidate.world_size if sp_group_candidate else 1 + tp_world_size = tp_group_candidate.world_size if tp_group_candidate else 1 + + has_sp = sp_world_size > 1 + has_tp = tp_world_size > 1 + + sp_group = sp_group_candidate.device_group if has_sp else None + tp_group = tp_group_candidate.device_group if has_tp else None + + logger.info( + "cache-dit enabled in distributed environment (world_size=%d, has_sp=%s, has_tp=%s)", + world_size, + has_sp, + has_tp, ) - return # === Parse SCM configuration from envs === # SCM is shared between primary and secondary transformers scm_preset = envs.SGLANG_CACHE_DIT_SCM_PRESET @@ -288,6 +309,8 @@ class DenoisingStage(PipelineStage): primary_config, secondary_config, model_name="wan2.2", + sp_group=sp_group, + tp_group=tp_group, ) logger.info( "cache-dit enabled on dual transformers (steps=%d)", @@ -299,6 +322,8 @@ class DenoisingStage(PipelineStage): self.transformer, primary_config, model_name="transformer", + sp_group=sp_group, + tp_group=tp_group, ) logger.info( "cache-dit enabled on transformer (steps=%d, Fn=%d, Bn=%d, rdt=%.3f)", diff --git a/python/sglang/multimodal_gen/runtime/server_args.py b/python/sglang/multimodal_gen/runtime/server_args.py index 985e6564a..df0142ef3 100644 --- a/python/sglang/multimodal_gen/runtime/server_args.py +++ b/python/sglang/multimodal_gen/runtime/server_args.py @@ -8,6 +8,7 @@ import argparse import dataclasses import inspect import json +import os import random import sys import tempfile @@ -956,6 +957,15 @@ class ServerArgs: "CFG Parallelism is enabled via `--enable-cfg-parallel`, while -num-gpus==1" ) + if os.getenv("SGLANG_CACHE_DIT_ENABLED", "").lower() == "true": + has_sp = self.sp_degree > 1 + has_tp = self.tp_size > 1 + if has_sp and has_tp: + raise ValueError( + "cache-dit does not support hybrid parallelism (SP + TP). " + "Please use either sequence parallelism or tensor parallelism, not both." + ) + @dataclasses.dataclass class PortArgs: diff --git a/python/sglang/multimodal_gen/runtime/utils/cache_dit_integration.py b/python/sglang/multimodal_gen/runtime/utils/cache_dit_integration.py index 0f0abd00d..877db9935 100644 --- a/python/sglang/multimodal_gen/runtime/utils/cache_dit_integration.py +++ b/python/sglang/multimodal_gen/runtime/utils/cache_dit_integration.py @@ -10,6 +10,7 @@ from dataclasses import dataclass from typing import List, Optional import torch +import torch.distributed as dist from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger @@ -25,6 +26,102 @@ from cache_dit import ( steps_mask, ) from cache_dit.caching.block_adapters import BlockAdapterRegister +from cache_dit.parallelism import ParallelismBackend, ParallelismConfig + +_original_similarity = None + + +def _patch_cache_dit_similarity(): + from cache_dit.caching.cache_contexts import cache_manager + + global _original_similarity + if _original_similarity is not None: + return + + _original_similarity = cache_manager.CachedContextManager.similarity + + def patched_similarity(self, t1, t2, *, threshold, parallelized=False, prefix="Fn"): + if not parallelized: + return _original_similarity( + self, + t1, + t2, + threshold=threshold, + parallelized=parallelized, + prefix=prefix, + ) + + sp_group = getattr(self, "_sglang_sp_group", None) + tp_group = getattr(self, "_sglang_tp_group", None) + target_group = sp_group or tp_group + + if target_group is None: + return _original_similarity( + self, + t1, + t2, + threshold=threshold, + parallelized=parallelized, + prefix=prefix, + ) + + # Adapted from https://github.com/vipshop/cache-dit/blob/main/src/cache_dit/caching/cache_contexts/cache_manager.py#L495-L523 + condition_thresh = self.get_important_condition_threshold() + if condition_thresh > 0.0: + raw_diff = (t1 - t2).abs() + token_m_df = raw_diff.mean(dim=-1) + token_m_t1 = t1.abs().mean(dim=-1) + token_diff = token_m_df / token_m_t1 + condition = token_diff > condition_thresh + if condition.sum() > 0: + condition = condition.unsqueeze(-1).expand_as(raw_diff) + mean_diff = raw_diff[condition].mean() + mean_t1 = t1[condition].abs().mean() + else: + mean_diff = (t1 - t2).abs().mean() + mean_t1 = t1.abs().mean() + else: + mean_diff = (t1 - t2).abs().mean() + mean_t1 = t1.abs().mean() + + dist.all_reduce(mean_diff, op=dist.ReduceOp.AVG, group=target_group) + dist.all_reduce(mean_t1, op=dist.ReduceOp.AVG, group=target_group) + + diff = (mean_diff / mean_t1).item() + self.add_residual_diff(diff) + return diff < threshold + + cache_manager.CachedContextManager.similarity = patched_similarity + + +def _build_parallelism_config(sp_group, tp_group): + if sp_group is None and tp_group is None: + return None + + ulysses_size = None + ring_size = None + if sp_group is not None: + ulysses_size = getattr(sp_group, "ulysses_world_size", None) + ring_size = getattr(sp_group, "ring_world_size", None) + + tp_size = None + if tp_group is not None: + tp_size = dist.get_world_size(tp_group.device_group) + + return ParallelismConfig( + backend=ParallelismBackend.NATIVE_PYTORCH, + ulysses_size=ulysses_size, + ring_size=ring_size, + tp_size=tp_size, + ) + + +def _mark_transformer_parallelized(transformer, config, sp_group, tp_group): + if config is None: + return + + transformer._is_parallelized = True + transformer._parallelism_config = config def get_scm_mask( @@ -118,6 +215,8 @@ def enable_cache_on_transformer( transformer: torch.nn.Module, config: CacheDitConfig, model_name: str = "transformer", + sp_group: Optional[torch.distributed.ProcessGroup] = None, + tp_group: Optional[torch.distributed.ProcessGroup] = None, ) -> torch.nn.Module: """Enable cache-dit on a transformer module, by wrapping the module with cache-dit @@ -126,6 +225,8 @@ def enable_cache_on_transformer( Args: model_name: Name of the model for logging purposes. + sp_group: Sequence parallel process group (for Ulysses/Ring). + tp_group: Tensor parallel process group. """ if not config.enabled: @@ -194,12 +295,25 @@ def enable_cache_on_transformer( config.steps_computation_policy, ) + parallelism_config = _build_parallelism_config(sp_group, tp_group) + if parallelism_config is not None: + _patch_cache_dit_similarity() + + _mark_transformer_parallelized(transformer, parallelism_config, sp_group, tp_group) + cache_dit.enable_cache( transformer, cache_config=cache_config, calibrator_config=calibrator_config, + parallelism_config=None, ) + if parallelism_config is not None: + context_manager = getattr(transformer, "_context_manager", None) + if context_manager is not None: + context_manager._sglang_sp_group = sp_group + context_manager._sglang_tp_group = tp_group + return transformer @@ -209,6 +323,8 @@ def enable_cache_on_dual_transformer( primary_config: CacheDitConfig, secondary_config: CacheDitConfig, model_name: str = "wan2.2", + sp_group: Optional[torch.distributed.ProcessGroup] = None, + tp_group: Optional[torch.distributed.ProcessGroup] = None, ) -> tuple[torch.nn.Module, torch.nn.Module]: """Enable cache-dit on dual transformers using BlockAdapter. @@ -219,6 +335,8 @@ def enable_cache_on_dual_transformer( Args: primary_config: CacheDitConfig for primary transformer. secondary_config: CacheDitConfig for secondary transformer. + sp_group: Sequence parallel process group (for Ulysses/Ring). + tp_group: Tensor parallel process group. """ _supported_dual_transformer_models = [ "wan2.2", # Currently, only Wan2.2 will run into dual-transformer case @@ -320,6 +438,15 @@ def enable_cache_on_dual_transformer( primary_config.steps_computation_policy, ) + parallelism_config = _build_parallelism_config(sp_group, tp_group) + if parallelism_config is not None: + _patch_cache_dit_similarity() + + _mark_transformer_parallelized(transformer, parallelism_config, sp_group, tp_group) + _mark_transformer_parallelized( + transformer_2, parallelism_config, sp_group, tp_group + ) + # Get blocks attribute - Wan transformers use 'blocks' attribute transformer_blocks = getattr(transformer, "blocks", None) transformer_2_blocks = getattr(transformer_2, "blocks", None) @@ -345,10 +472,18 @@ def enable_cache_on_dual_transformer( params_modifiers=[primary_modifier, secondary_modifier], has_separate_cfg=True, ), + parallelism_config=None, ) else: raise ValueError( f"Dual-transformer is not implemented for model {model_name} yet." ) + if parallelism_config is not None: + for t in [transformer, transformer_2]: + context_manager = getattr(t, "_context_manager", None) + if context_manager is not None: + context_manager._sglang_sp_group = sp_group + context_manager._sglang_tp_group = tp_group + return transformer, transformer_2