[diffusion] chore: expose architecture config at the dit runtime boundary (#34248)
This commit is contained in:
+5
-4
@@ -331,12 +331,12 @@ class SpectrumMixin:
|
|||||||
- ``spectrum_record_features()`` — after a real forward, store block outputs.
|
- ``spectrum_record_features()`` — after a real forward, store block outputs.
|
||||||
- ``spectrum_predict_features()`` — on skipped steps, return forecasted outputs.
|
- ``spectrum_predict_features()`` — on skipped steps, return forecasted outputs.
|
||||||
|
|
||||||
Models with separate CFG branches (Wan, Hunyuan, SD3) list their config prefix
|
Models with separate CFG branches (Wan, Hunyuan, SD3) list their model prefix
|
||||||
in ``_CFG_SUPPORTED_PREFIXES`` so cond/uncond maintain independent counters and
|
in ``_CFG_SUPPORTED_PREFIXES`` so cond/uncond maintain independent counters and
|
||||||
forecasters. All other ``CachableDiT`` subclasses share one counter.
|
forecasters. All other ``CachableDiT`` subclasses share one counter.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
# DiT config prefixes that run separate cond/uncond forwards (see TeaCache).
|
# DiT model prefixes that run separate cond/uncond forwards (see TeaCache).
|
||||||
_CFG_SUPPORTED_PREFIXES: set[str] = {"wan", "hunyuan", "sd3"}
|
_CFG_SUPPORTED_PREFIXES: set[str] = {"wan", "hunyuan", "sd3"}
|
||||||
|
|
||||||
def _init_spectrum_state(self) -> None:
|
def _init_spectrum_state(self) -> None:
|
||||||
@@ -366,8 +366,9 @@ class SpectrumMixin:
|
|||||||
# Runtime branch tracking
|
# Runtime branch tracking
|
||||||
self.spectrum_is_cfg_negative = False
|
self.spectrum_is_cfg_negative = False
|
||||||
self._spectrum_ctx: Optional[SpectrumContext] = None
|
self._spectrum_ctx: Optional[SpectrumContext] = None
|
||||||
prefix = getattr(self.config, "prefix", "").lower()
|
self._spectrum_supports_cfg_cache = (
|
||||||
self._spectrum_supports_cfg_cache = prefix in self._CFG_SUPPORTED_PREFIXES
|
self.prefix.lower() in self._CFG_SUPPORTED_PREFIXES
|
||||||
|
)
|
||||||
|
|
||||||
def reset_spectrum_state(self, spectrum_params: SpectrumParams) -> None:
|
def reset_spectrum_state(self, spectrum_params: SpectrumParams) -> None:
|
||||||
self.spectrum_cnt = 0
|
self.spectrum_cnt = 0
|
||||||
|
|||||||
+2
-6
@@ -23,8 +23,6 @@ from typing import TYPE_CHECKING, Any
|
|||||||
import numpy as np
|
import numpy as np
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.multimodal_gen.configs.models import DiTConfig
|
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.multimodal_gen.configs.sample.teacache import TeaCacheParams
|
from sglang.multimodal_gen.configs.sample.teacache import TeaCacheParams
|
||||||
|
|
||||||
@@ -127,7 +125,7 @@ class TeaCacheMixin:
|
|||||||
# Models that support CFG cache separation (wan/hunyuan/zimage)
|
# Models that support CFG cache separation (wan/hunyuan/zimage)
|
||||||
# Models not in this set (flux/qwen) auto-disable TeaCache when CFG is enabled
|
# Models not in this set (flux/qwen) auto-disable TeaCache when CFG is enabled
|
||||||
_CFG_SUPPORTED_PREFIXES: set[str] = {"wan", "hunyuan", "zimage"}
|
_CFG_SUPPORTED_PREFIXES: set[str] = {"wan", "hunyuan", "zimage"}
|
||||||
config: DiTConfig
|
prefix: str
|
||||||
|
|
||||||
def _init_teacache_state(self) -> None:
|
def _init_teacache_state(self) -> None:
|
||||||
"""Initialize TeaCache state. Call this in subclass __init__."""
|
"""Initialize TeaCache state. Call this in subclass __init__."""
|
||||||
@@ -135,9 +133,7 @@ class TeaCacheMixin:
|
|||||||
self.cnt = 0
|
self.cnt = 0
|
||||||
self.enable_teacache = True
|
self.enable_teacache = True
|
||||||
# Flag indicating if this model supports CFG cache separation
|
# Flag indicating if this model supports CFG cache separation
|
||||||
self._supports_cfg_cache = (
|
self._supports_cfg_cache = self.prefix.lower() in self._CFG_SUPPORTED_PREFIXES
|
||||||
self.config.prefix.lower() in self._CFG_SUPPORTED_PREFIXES
|
|
||||||
)
|
|
||||||
|
|
||||||
# Always initialize positive cache fields (used in all modes)
|
# Always initialize positive cache fields (used in all modes)
|
||||||
self.previous_modulated_input: torch.Tensor | None = None
|
self.previous_modulated_input: torch.Tensor | None = None
|
||||||
|
|||||||
@@ -452,10 +452,13 @@ class FP32LayerNorm(CustomOp, nn.LayerNorm):
|
|||||||
)
|
)
|
||||||
self._forward_method = self.dispatch_forward()
|
self._forward_method = self.dispatch_forward()
|
||||||
|
|
||||||
|
if _is_npu:
|
||||||
try:
|
try:
|
||||||
import attentions # noqa: F401
|
import attentions # noqa: F401
|
||||||
except ImportError:
|
except ImportError:
|
||||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
from sglang.multimodal_gen.runtime.utils.logging_utils import (
|
||||||
|
init_logger,
|
||||||
|
)
|
||||||
|
|
||||||
logger = init_logger(__name__) # pylint: disable=invalid-name
|
logger = init_logger(__name__) # pylint: disable=invalid-name
|
||||||
logger.warning(
|
logger.warning(
|
||||||
|
|||||||
@@ -7,7 +7,7 @@ from typing import Any
|
|||||||
import torch
|
import torch
|
||||||
from torch import nn
|
from torch import nn
|
||||||
|
|
||||||
from sglang.multimodal_gen.configs.models import DiTConfig
|
from sglang.multimodal_gen.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||||
|
|
||||||
# NOTE: SpectrumMixin lives in runtime.cache.spectrum
|
# NOTE: SpectrumMixin lives in runtime.cache.spectrum
|
||||||
from sglang.multimodal_gen.runtime.cache.spectrum import SpectrumMixin
|
from sglang.multimodal_gen.runtime.cache.spectrum import SpectrumMixin
|
||||||
@@ -49,7 +49,10 @@ class BaseDiT(nn.Module, ABC):
|
|||||||
|
|
||||||
def __init__(self, config: DiTConfig, hf_config: dict[str, Any], **kwargs) -> None:
|
def __init__(self, config: DiTConfig, hf_config: dict[str, Any], **kwargs) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.config = config
|
# runtime models expose checkpoint architecture through `config`; load
|
||||||
|
# settings such as the model prefix stay separate
|
||||||
|
self.config: DiTArchConfig = config.arch_config
|
||||||
|
self.prefix = config.prefix
|
||||||
self.hf_config = hf_config
|
self.hf_config = hf_config
|
||||||
if not self.supported_attention_backends:
|
if not self.supported_attention_backends:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
|
|||||||
@@ -520,7 +520,7 @@ class CausalWanTransformer3DModel(BaseDiT, LayerwiseOffloadableModuleMixin):
|
|||||||
|
|
||||||
# Causal-specific
|
# Causal-specific
|
||||||
self.block_mask = None
|
self.block_mask = None
|
||||||
self.num_frame_per_block = config.arch_config.num_frames_per_block
|
self.num_frame_per_block = self.config.num_frames_per_block
|
||||||
# Block size is bounded only by the causal block-mask construction, which
|
# Block size is bounded only by the causal block-mask construction, which
|
||||||
# supports any positive value.
|
# supports any positive value.
|
||||||
assert self.num_frame_per_block >= 1
|
assert self.num_frame_per_block >= 1
|
||||||
|
|||||||
@@ -933,7 +933,7 @@ class Cosmos3OmniTransformer(CachableDiT, LayerwiseOffloadableModuleMixin):
|
|||||||
) -> None:
|
) -> None:
|
||||||
super().__init__(config=config, hf_config=hf_config)
|
super().__init__(config=config, hf_config=hf_config)
|
||||||
|
|
||||||
arch = config.arch_config
|
arch = self.config
|
||||||
self.hidden_size = arch.hidden_size
|
self.hidden_size = arch.hidden_size
|
||||||
self.num_hidden_layers = arch.num_hidden_layers
|
self.num_hidden_layers = arch.num_hidden_layers
|
||||||
self.num_attention_heads = arch.num_attention_heads
|
self.num_attention_heads = arch.num_attention_heads
|
||||||
|
|||||||
@@ -427,7 +427,7 @@ class ErnieImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin)
|
|||||||
):
|
):
|
||||||
super().__init__(config=config, hf_config=hf_config)
|
super().__init__(config=config, hf_config=hf_config)
|
||||||
|
|
||||||
arch = config.arch_config
|
arch = self.config
|
||||||
self.hidden_size = arch.hidden_size
|
self.hidden_size = arch.hidden_size
|
||||||
self.num_attention_heads = arch.num_attention_heads
|
self.num_attention_heads = arch.num_attention_heads
|
||||||
self.num_channels_latents = arch.out_channels
|
self.num_channels_latents = arch.out_channels
|
||||||
|
|||||||
@@ -1202,7 +1202,6 @@ class FluxTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
|
|||||||
quant_config: Optional[QuantizationConfig] = None,
|
quant_config: Optional[QuantizationConfig] = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
super().__init__(config=config, hf_config=hf_config)
|
super().__init__(config=config, hf_config=hf_config)
|
||||||
self.config = config.arch_config
|
|
||||||
|
|
||||||
self.out_channels = (
|
self.out_channels = (
|
||||||
getattr(self.config, "out_channels", None) or self.config.in_channels
|
getattr(self.config, "out_channels", None) or self.config.in_channels
|
||||||
|
|||||||
@@ -945,8 +945,7 @@ class GlmImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
|
|||||||
):
|
):
|
||||||
super().__init__(config=config, hf_config=hf_config)
|
super().__init__(config=config, hf_config=hf_config)
|
||||||
|
|
||||||
self.config_data = config # Store config
|
arch_config = self.config
|
||||||
arch_config = config.arch_config
|
|
||||||
|
|
||||||
self.in_channels = arch_config.in_channels
|
self.in_channels = arch_config.in_channels
|
||||||
self.out_channels = arch_config.out_channels
|
self.out_channels = arch_config.out_channels
|
||||||
|
|||||||
@@ -487,7 +487,7 @@ class Hunyuan3D2DiT(CachableDiT, LayerwiseOffloadableModuleMixin):
|
|||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
super().__init__(config=config, hf_config=hf_config or {}, **kwargs)
|
super().__init__(config=config, hf_config=hf_config or {}, **kwargs)
|
||||||
arch = config.arch_config
|
arch = self.config
|
||||||
|
|
||||||
in_channels = arch.in_channels
|
in_channels = arch.in_channels
|
||||||
context_in_dim = arch.context_in_dim
|
context_in_dim = arch.context_in_dim
|
||||||
|
|||||||
@@ -513,7 +513,7 @@ class Ideogram4Transformer2DModel(BaseDiT, LayerwiseOffloadableModuleMixin):
|
|||||||
**kwargs,
|
**kwargs,
|
||||||
) -> None:
|
) -> None:
|
||||||
super().__init__(config, hf_config, **kwargs)
|
super().__init__(config, hf_config, **kwargs)
|
||||||
cfg = config.arch_config
|
cfg = self.config
|
||||||
use_weight_only_fp8_linears = config.use_weight_only_fp8_linears
|
use_weight_only_fp8_linears = config.use_weight_only_fp8_linears
|
||||||
self._supported_attention_backends = cfg._supported_attention_backends
|
self._supported_attention_backends = cfg._supported_attention_backends
|
||||||
hidden_size = cfg.num_attention_heads * cfg.attention_head_dim
|
hidden_size = cfg.num_attention_heads * cfg.attention_head_dim
|
||||||
@@ -586,7 +586,7 @@ class Ideogram4Transformer2DModel(BaseDiT, LayerwiseOffloadableModuleMixin):
|
|||||||
def post_load_weights(self) -> None:
|
def post_load_weights(self) -> None:
|
||||||
if not self.rotary_emb.inv_freq.is_meta:
|
if not self.rotary_emb.inv_freq.is_meta:
|
||||||
return
|
return
|
||||||
cfg = self.config.arch_config
|
cfg = self.config
|
||||||
inv_freq = 1.0 / (
|
inv_freq = 1.0 / (
|
||||||
cfg.rope_theta
|
cfg.rope_theta
|
||||||
** (
|
** (
|
||||||
|
|||||||
@@ -522,7 +522,7 @@ class Krea2Transformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
|
|||||||
quant_config: Optional[Any] = None,
|
quant_config: Optional[Any] = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
super().__init__(config=config, hf_config=hf_config)
|
super().__init__(config=config, hf_config=hf_config)
|
||||||
ac = config.arch_config
|
ac = self.config
|
||||||
self.arch_config = ac
|
self.arch_config = ac
|
||||||
|
|
||||||
self.hidden_size = ac.features
|
self.hidden_size = ac.features
|
||||||
|
|||||||
@@ -1518,7 +1518,7 @@ class LTX2VideoTransformer3DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
|
|||||||
return timestep.amax(dim=tuple(range(1, timestep.ndim)))
|
return timestep.amax(dim=tuple(range(1, timestep.ndim)))
|
||||||
|
|
||||||
def _scale_timestep_for_adaln(self, timestep: torch.Tensor) -> torch.Tensor:
|
def _scale_timestep_for_adaln(self, timestep: torch.Tensor) -> torch.Tensor:
|
||||||
ltx_variant = str(getattr(self.config.arch_config, "ltx_variant", "ltx_2"))
|
ltx_variant = str(getattr(self.config, "ltx_variant", "ltx_2"))
|
||||||
if ltx_variant == "ltx_2_3" and bool(
|
if ltx_variant == "ltx_2_3" and bool(
|
||||||
getattr(self, "_sglang_use_ltx23_hq_timestep_semantics", False)
|
getattr(self, "_sglang_use_ltx23_hq_timestep_semantics", False)
|
||||||
):
|
):
|
||||||
@@ -1583,7 +1583,7 @@ class LTX2VideoTransformer3DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
|
|||||||
) -> None:
|
) -> None:
|
||||||
super().__init__(config=config, hf_config=hf_config)
|
super().__init__(config=config, hf_config=hf_config)
|
||||||
|
|
||||||
arch = config.arch_config
|
arch = self.config
|
||||||
self.hidden_size = arch.hidden_size
|
self.hidden_size = arch.hidden_size
|
||||||
self.num_attention_heads = arch.num_attention_heads
|
self.num_attention_heads = arch.num_attention_heads
|
||||||
self.audio_hidden_size = arch.audio_hidden_size
|
self.audio_hidden_size = arch.audio_hidden_size
|
||||||
@@ -1844,7 +1844,7 @@ class LTX2VideoTransformer3DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
|
|||||||
return video_coords.to(device=hidden_device)
|
return video_coords.to(device=hidden_device)
|
||||||
|
|
||||||
def _get_av_ca_gate_timestep_factor(self) -> float:
|
def _get_av_ca_gate_timestep_factor(self) -> float:
|
||||||
ltx_variant = str(getattr(self.config.arch_config, "ltx_variant", "ltx_2"))
|
ltx_variant = str(getattr(self.config, "ltx_variant", "ltx_2"))
|
||||||
if ltx_variant == "ltx_2_3":
|
if ltx_variant == "ltx_2_3":
|
||||||
return self.av_ca_timestep_scale_multiplier / self.timestep_scale_multiplier
|
return self.av_ca_timestep_scale_multiplier / self.timestep_scale_multiplier
|
||||||
return float(self.av_ca_timestep_scale_multiplier)
|
return float(self.av_ca_timestep_scale_multiplier)
|
||||||
@@ -1856,7 +1856,7 @@ class LTX2VideoTransformer3DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
|
|||||||
prompt_timestep: torch.Tensor | None,
|
prompt_timestep: torch.Tensor | None,
|
||||||
audio_prompt_timestep: torch.Tensor | None,
|
audio_prompt_timestep: torch.Tensor | None,
|
||||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||||
ltx_variant = str(getattr(self.config.arch_config, "ltx_variant", "ltx_2"))
|
ltx_variant = str(getattr(self.config, "ltx_variant", "ltx_2"))
|
||||||
if ltx_variant != "ltx_2_3":
|
if ltx_variant != "ltx_2_3":
|
||||||
return timestep, audio_timestep
|
return timestep, audio_timestep
|
||||||
|
|
||||||
|
|||||||
@@ -1109,7 +1109,7 @@ class MiniMaxH3DiTModel(BaseDiT, LayerwiseOffloadableModuleMixin):
|
|||||||
quant_config: QuantizationConfig | None = None,
|
quant_config: QuantizationConfig | None = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
super().__init__(config=config, hf_config=hf_config)
|
super().__init__(config=config, hf_config=hf_config)
|
||||||
arch = config.arch_config
|
arch = self.config
|
||||||
self.arch = arch
|
self.arch = arch
|
||||||
self.hidden_size = arch.hidden_size
|
self.hidden_size = arch.hidden_size
|
||||||
self.num_attention_heads = arch.num_attention_heads
|
self.num_attention_heads = arch.num_attention_heads
|
||||||
|
|||||||
@@ -1371,23 +1371,24 @@ class QwenImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
|
|||||||
quant_config: Optional[QuantizationConfig] = None,
|
quant_config: Optional[QuantizationConfig] = None,
|
||||||
):
|
):
|
||||||
super().__init__(config=config, hf_config=hf_config)
|
super().__init__(config=config, hf_config=hf_config)
|
||||||
patch_size = config.arch_config.patch_size
|
arch = self.config
|
||||||
in_channels = config.arch_config.in_channels
|
patch_size = arch.patch_size
|
||||||
out_channels = config.arch_config.out_channels
|
in_channels = arch.in_channels
|
||||||
num_layers = config.arch_config.num_layers
|
out_channels = arch.out_channels
|
||||||
attention_head_dim = config.arch_config.attention_head_dim
|
num_layers = arch.num_layers
|
||||||
num_attention_heads = config.arch_config.num_attention_heads
|
attention_head_dim = arch.attention_head_dim
|
||||||
joint_attention_dim = config.arch_config.joint_attention_dim
|
num_attention_heads = arch.num_attention_heads
|
||||||
axes_dims_rope = config.arch_config.axes_dims_rope
|
joint_attention_dim = arch.joint_attention_dim
|
||||||
self.zero_cond_t = getattr(config.arch_config, "zero_cond_t", False)
|
axes_dims_rope = arch.axes_dims_rope
|
||||||
|
self.zero_cond_t = getattr(arch, "zero_cond_t", False)
|
||||||
self.out_channels = out_channels or in_channels
|
self.out_channels = out_channels or in_channels
|
||||||
self.inner_dim = num_attention_heads * attention_head_dim
|
self.inner_dim = num_attention_heads * attention_head_dim
|
||||||
|
|
||||||
self.use_additional_t_cond: bool = getattr(
|
self.use_additional_t_cond: bool = getattr(
|
||||||
config.arch_config, "use_additional_t_cond", False
|
arch, "use_additional_t_cond", False
|
||||||
) # For qwen-image-layered now
|
) # For qwen-image-layered now
|
||||||
self.use_layer3d_rope: bool = getattr(
|
self.use_layer3d_rope: bool = getattr(
|
||||||
config.arch_config, "use_layer3d_rope", False
|
arch, "use_layer3d_rope", False
|
||||||
) # For qwen-image-layered now
|
) # For qwen-image-layered now
|
||||||
|
|
||||||
if not self.use_layer3d_rope:
|
if not self.use_layer3d_rope:
|
||||||
|
|||||||
@@ -401,7 +401,7 @@ class SanaTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
|
|||||||
def __init__(self, config: SanaConfig, hf_config=None, **kwargs):
|
def __init__(self, config: SanaConfig, hf_config=None, **kwargs):
|
||||||
super().__init__(config, hf_config=hf_config or {}, **kwargs)
|
super().__init__(config, hf_config=hf_config or {}, **kwargs)
|
||||||
|
|
||||||
arch = config.arch_config
|
arch = self.config
|
||||||
self.out_channels = arch.out_channels
|
self.out_channels = arch.out_channels
|
||||||
self.patch_size = arch.patch_size
|
self.patch_size = arch.patch_size
|
||||||
self.inner_dim = arch.num_attention_heads * arch.attention_head_dim
|
self.inner_dim = arch.num_attention_heads * arch.attention_head_dim
|
||||||
|
|||||||
@@ -353,7 +353,7 @@ class SanaWMTransformer3DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
|
|||||||
|
|
||||||
def __init__(self, config: SanaWMConfig, hf_config=None, **kwargs) -> None:
|
def __init__(self, config: SanaWMConfig, hf_config=None, **kwargs) -> None:
|
||||||
super().__init__(config, hf_config=hf_config or {}, **kwargs)
|
super().__init__(config, hf_config=hf_config or {}, **kwargs)
|
||||||
arch = config.arch_config
|
arch = self.config
|
||||||
|
|
||||||
self.patch_size = (arch.patch_size_t, arch.patch_size, arch.patch_size)
|
self.patch_size = (arch.patch_size_t, arch.patch_size, arch.patch_size)
|
||||||
self.inner_dim = arch.num_attention_heads * arch.attention_head_dim
|
self.inner_dim = arch.num_attention_heads * arch.attention_head_dim
|
||||||
|
|||||||
@@ -226,7 +226,7 @@ class SanaWMLTX2VideoRefiner(CachableDiT, LayerwiseOffloadableModuleMixin):
|
|||||||
quant_config: QuantizationConfig | None = None,
|
quant_config: QuantizationConfig | None = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
super().__init__(config, hf_config=hf_config)
|
super().__init__(config, hf_config=hf_config)
|
||||||
arch = config.arch_config
|
arch = self.config
|
||||||
|
|
||||||
self.in_channels = int(arch.in_channels)
|
self.in_channels = int(arch.in_channels)
|
||||||
self.out_channels = int(arch.out_channels)
|
self.out_channels = int(arch.out_channels)
|
||||||
|
|||||||
@@ -40,8 +40,7 @@ class SD3Transformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
|
|||||||
quant_config=None,
|
quant_config=None,
|
||||||
):
|
):
|
||||||
super().__init__(config=config, hf_config=hf_config)
|
super().__init__(config=config, hf_config=hf_config)
|
||||||
self.config = config
|
arch_config = self.config
|
||||||
arch_config = config.arch_config
|
|
||||||
sample_size = arch_config.sample_size
|
sample_size = arch_config.sample_size
|
||||||
patch_size = arch_config.patch_size
|
patch_size = arch_config.patch_size
|
||||||
in_channels = arch_config.in_channels
|
in_channels = arch_config.in_channels
|
||||||
|
|||||||
@@ -808,8 +808,7 @@ class ZImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
|
|||||||
) -> None:
|
) -> None:
|
||||||
super().__init__(config=config, hf_config=hf_config)
|
super().__init__(config=config, hf_config=hf_config)
|
||||||
|
|
||||||
self.config_data = config # Store config
|
arch_config = self.config
|
||||||
arch_config = config.arch_config
|
|
||||||
|
|
||||||
self.in_channels = arch_config.in_channels
|
self.in_channels = arch_config.in_channels
|
||||||
self.out_channels = arch_config.out_channels
|
self.out_channels = arch_config.out_channels
|
||||||
|
|||||||
@@ -157,12 +157,10 @@ class CausalDMDDenoisingStage(DenoisingStage):
|
|||||||
self.causal_kv_cache_neg: list | None = None
|
self.causal_kv_cache_neg: list | None = None
|
||||||
self.crossattn_cache_neg: list | None = None
|
self.crossattn_cache_neg: list | None = None
|
||||||
# Model-dependent constants (aligned with causal_inference.py assumptions)
|
# Model-dependent constants (aligned with causal_inference.py assumptions)
|
||||||
self.num_transformer_blocks = self.transformer.config.arch_config.num_layers
|
self.num_transformer_blocks = self.transformer.config.num_layers
|
||||||
self.num_frames_per_block = (
|
self.num_frames_per_block = self.transformer.config.num_frames_per_block
|
||||||
self.transformer.config.arch_config.num_frames_per_block
|
|
||||||
)
|
|
||||||
self.sliding_window_num_frames = (
|
self.sliding_window_num_frames = (
|
||||||
self.transformer.config.arch_config.sliding_window_num_frames
|
self.transformer.config.sliding_window_num_frames
|
||||||
)
|
)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@@ -171,7 +169,7 @@ class CausalDMDDenoisingStage(DenoisingStage):
|
|||||||
) # type: ignore
|
) # type: ignore
|
||||||
except Exception:
|
except Exception:
|
||||||
self.local_attn_size = -1
|
self.local_attn_size = -1
|
||||||
self.sink_size = self.transformer.config.arch_config.sink_size
|
self.sink_size = self.transformer.config.sink_size
|
||||||
|
|
||||||
self._causal_attn_metadata_builder_cls = None
|
self._causal_attn_metadata_builder_cls = None
|
||||||
self._causal_attn_metadata_builder = None
|
self._causal_attn_metadata_builder = None
|
||||||
@@ -190,8 +188,8 @@ class CausalDMDDenoisingStage(DenoisingStage):
|
|||||||
|
|
||||||
def _prepare_frame_seq_length(self, h: int, w: int) -> int:
|
def _prepare_frame_seq_length(self, h: int, w: int) -> int:
|
||||||
patch_ratio = (
|
patch_ratio = (
|
||||||
self.transformer.config.arch_config.patch_size[-1]
|
self.transformer.config.patch_size[-1]
|
||||||
* self.transformer.config.arch_config.patch_size[-2]
|
* self.transformer.config.patch_size[-2]
|
||||||
)
|
)
|
||||||
self.num_token_per_frame = (h * w) // patch_ratio
|
self.num_token_per_frame = (h * w) // patch_ratio
|
||||||
return self.num_token_per_frame
|
return self.num_token_per_frame
|
||||||
|
|||||||
+2
-3
@@ -27,8 +27,7 @@ class RealtimeChunkLatentPreparationStage(LatentPreparationStage):
|
|||||||
server_args: ServerArgs,
|
server_args: ServerArgs,
|
||||||
) -> int:
|
) -> int:
|
||||||
return int(
|
return int(
|
||||||
batch.realtime_chunk_size
|
batch.realtime_chunk_size or self.transformer.config.num_frames_per_block
|
||||||
or self.transformer.config.arch_config.num_frames_per_block
|
|
||||||
)
|
)
|
||||||
|
|
||||||
def get_latent_preparation_spec(
|
def get_latent_preparation_spec(
|
||||||
@@ -47,7 +46,7 @@ class RealtimeChunkLatentPreparationStage(LatentPreparationStage):
|
|||||||
return LatentPreparationSpec(
|
return LatentPreparationSpec(
|
||||||
shape=(
|
shape=(
|
||||||
condition_latent.shape[0],
|
condition_latent.shape[0],
|
||||||
self.transformer.config.arch_config.out_channels,
|
self.transformer.config.out_channels,
|
||||||
num_frames,
|
num_frames,
|
||||||
condition_latent.shape[3],
|
condition_latent.shape[3],
|
||||||
condition_latent.shape[4],
|
condition_latent.shape[4],
|
||||||
|
|||||||
@@ -0,0 +1,25 @@
|
|||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.configs.models.dits.base import DiTConfig
|
||||||
|
from sglang.multimodal_gen.runtime.models.dits.base import CachableDiT
|
||||||
|
|
||||||
|
|
||||||
|
class _TestDiT(CachableDiT):
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
hidden_states: torch.Tensor,
|
||||||
|
encoder_hidden_states: torch.Tensor,
|
||||||
|
timestep: torch.LongTensor,
|
||||||
|
**kwargs,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
return hidden_states
|
||||||
|
|
||||||
|
|
||||||
|
def test_dit_runtime_keeps_architecture_and_component_config_separate():
|
||||||
|
component_config = DiTConfig(prefix="Wan")
|
||||||
|
model = _TestDiT(config=component_config, hf_config={})
|
||||||
|
|
||||||
|
assert model.config is component_config.arch_config
|
||||||
|
assert model.prefix == "Wan"
|
||||||
|
assert model._supports_cfg_cache
|
||||||
|
assert model._spectrum_supports_cfg_cache
|
||||||
Reference in New Issue
Block a user