[diffusion] fix: cache dit with parallel (#15163)

Co-authored-by: Mick <mickjagger19@icloud.com>
This commit is contained in:
Xiaoyu Zhang
2025-12-15 19:15:51 +08:00
committed by GitHub
parent bf6438142a
commit 92c29d43ac
3 changed files with 177 additions and 7 deletions

View File

@@ -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)",

View File

@@ -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:

View File

@@ -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