[diffusion] fix: keep cosmos3 T=1 fusion on blackwell only (#35612)
This commit is contained in:
@@ -55,6 +55,7 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload im
|
|||||||
LayerwiseOffloadableModuleMixin,
|
LayerwiseOffloadableModuleMixin,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.models.dits.base import CachableDiT
|
from sglang.multimodal_gen.runtime.models.dits.base import CachableDiT
|
||||||
|
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
from sglang.srt.utils import add_prefix
|
from sglang.srt.utils import add_prefix
|
||||||
|
|
||||||
@@ -1600,6 +1601,14 @@ class Cosmos3OmniTransformer(CachableDiT, LayerwiseOffloadableModuleMixin):
|
|||||||
|
|
||||||
self._ensure_cache_dicts()
|
self._ensure_cache_dicts()
|
||||||
|
|
||||||
|
# The T=1 fused path is faster on Blackwell, but regresses the Hopper
|
||||||
|
# Cosmos3-Super two-GPU workload. Keep Hopper on the original split path.
|
||||||
|
enable_t1_fused_qk_norm_rope = (
|
||||||
|
T == 1
|
||||||
|
and current_platform.is_blackwell()
|
||||||
|
and not self._gen_layers_torch_compiled
|
||||||
|
)
|
||||||
|
|
||||||
# Compute UND K/V cache for this cache_key if not already cached
|
# Compute UND K/V cache for this cache_key if not already cached
|
||||||
# This allows reusing the cache across denoising steps for the same text
|
# This allows reusing the cache across denoising steps for the same text
|
||||||
if (
|
if (
|
||||||
@@ -1636,7 +1645,7 @@ class Cosmos3OmniTransformer(CachableDiT, LayerwiseOffloadableModuleMixin):
|
|||||||
vis_pos_ids, cache_dtype=hidden_gen.dtype
|
vis_pos_ids, cache_dtype=hidden_gen.dtype
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
if T == 1 and not self._gen_layers_torch_compiled:
|
if enable_t1_fused_qk_norm_rope:
|
||||||
# build_rope_cache_inputs already rounds through cache_dtype
|
# build_rope_cache_inputs already rounds through cache_dtype
|
||||||
# before returning FP32 storage. Keep that rounded cache in the
|
# before returning FP32 storage. Keep that rounded cache in the
|
||||||
# activation dtype so the exact fused QKNorm+RoPE kernel can
|
# activation dtype so the exact fused QKNorm+RoPE kernel can
|
||||||
@@ -1654,11 +1663,11 @@ class Cosmos3OmniTransformer(CachableDiT, LayerwiseOffloadableModuleMixin):
|
|||||||
# fused add+rmsnorm path instead of separate add + norm kernels.
|
# fused add+rmsnorm path instead of separate add + norm kernels.
|
||||||
cached_kv_for_key = self.cached_kv[cache_key]
|
cached_kv_for_key = self.cached_kv[cache_key]
|
||||||
residual: torch.Tensor | None = None
|
residual: torch.Tensor | None = None
|
||||||
round_norm_before_rope = T == 1
|
round_norm_before_rope = enable_t1_fused_qk_norm_rope
|
||||||
use_fused_qk_norm_rope = T > 1 or (
|
use_fused_qk_norm_rope = T > 1 or (
|
||||||
hidden_gen.device.type == "cuda"
|
enable_t1_fused_qk_norm_rope
|
||||||
|
and hidden_gen.device.type == "cuda"
|
||||||
and not torch.compiler.is_compiling()
|
and not torch.compiler.is_compiling()
|
||||||
and not self._gen_layers_torch_compiled
|
|
||||||
and get_sp_world_size() == 1
|
and get_sp_world_size() == 1
|
||||||
and can_use_fused_inplace_qknorm_rope(
|
and can_use_fused_inplace_qknorm_rope(
|
||||||
self.head_dim,
|
self.head_dim,
|
||||||
|
|||||||
Reference in New Issue
Block a user