fix: make Cache-DiT actually cache on MiniMax-H3 (#33827)
Signed-off-by: YZLi <yuanli@nvidia.com> Signed-off-by: yunch <yunch@nvidia.com> Co-authored-by: YZLi <yuanli@nvidia.com> Co-authored-by: Mick <mickjagger19@icloud.com> Co-authored-by: Xiaoyu Zhang <1182563586@qq.com>
This commit is contained in:
co-authored by
YZLi
Mick
Xiaoyu Zhang
parent
c255fbc4fe
commit
69a31ce342
@@ -53,6 +53,7 @@ def _indexed_gate_bf16_kernel(
|
|||||||
other_ptr,
|
other_ptr,
|
||||||
indices_ptr,
|
indices_ptr,
|
||||||
hidden_size,
|
hidden_size,
|
||||||
|
stride_output_row,
|
||||||
stride_x_row,
|
stride_x_row,
|
||||||
stride_gate_row,
|
stride_gate_row,
|
||||||
stride_other_row,
|
stride_other_row,
|
||||||
@@ -76,7 +77,7 @@ def _indexed_gate_bf16_kernel(
|
|||||||
|
|
||||||
gated = round_bf16_to_fp32(gate * other)
|
gated = round_bf16_to_fp32(gate * other)
|
||||||
tl.store(
|
tl.store(
|
||||||
output_ptr + row * stride_x_row + columns,
|
output_ptr + row * stride_output_row + columns,
|
||||||
x + gated,
|
x + gated,
|
||||||
mask=mask,
|
mask=mask,
|
||||||
)
|
)
|
||||||
@@ -109,7 +110,8 @@ def indexed_scale_shift_bf16_(
|
|||||||
return x
|
return x
|
||||||
|
|
||||||
|
|
||||||
def indexed_gate_bf16_(
|
def _indexed_gate_bf16(
|
||||||
|
output: torch.Tensor,
|
||||||
x: torch.Tensor,
|
x: torch.Tensor,
|
||||||
gate: torch.Tensor,
|
gate: torch.Tensor,
|
||||||
other: torch.Tensor,
|
other: torch.Tensor,
|
||||||
@@ -117,15 +119,16 @@ def indexed_gate_bf16_(
|
|||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
rows, hidden_size = x.shape
|
rows, hidden_size = x.shape
|
||||||
if rows == 0:
|
if rows == 0:
|
||||||
return x
|
return output
|
||||||
block_n = triton.next_power_of_2(hidden_size)
|
block_n = triton.next_power_of_2(hidden_size)
|
||||||
_indexed_gate_bf16_kernel[(rows,)](
|
_indexed_gate_bf16_kernel[(rows,)](
|
||||||
x,
|
output,
|
||||||
x,
|
x,
|
||||||
gate,
|
gate,
|
||||||
other,
|
other,
|
||||||
indices,
|
indices,
|
||||||
hidden_size,
|
hidden_size,
|
||||||
|
output.stride(0),
|
||||||
x.stride(0),
|
x.stride(0),
|
||||||
gate.stride(0),
|
gate.stride(0),
|
||||||
other.stride(0),
|
other.stride(0),
|
||||||
@@ -133,4 +136,22 @@ def indexed_gate_bf16_(
|
|||||||
BLOCK_N=block_n,
|
BLOCK_N=block_n,
|
||||||
num_warps=8,
|
num_warps=8,
|
||||||
)
|
)
|
||||||
return x
|
return output
|
||||||
|
|
||||||
|
|
||||||
|
def indexed_gate_bf16_(
|
||||||
|
x: torch.Tensor,
|
||||||
|
gate: torch.Tensor,
|
||||||
|
other: torch.Tensor,
|
||||||
|
indices: torch.Tensor,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
return _indexed_gate_bf16(x, x, gate, other, indices)
|
||||||
|
|
||||||
|
|
||||||
|
def indexed_gate_bf16(
|
||||||
|
x: torch.Tensor,
|
||||||
|
gate: torch.Tensor,
|
||||||
|
other: torch.Tensor,
|
||||||
|
indices: torch.Tensor,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
return _indexed_gate_bf16(torch.empty_like(x), x, gate, other, indices)
|
||||||
|
|||||||
@@ -273,16 +273,26 @@ DUAL_TRANSFORMER_BLOCK_ADAPTER_SPECS: dict[str, DualTransformerBlockAdapterSpec]
|
|||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
# Custom BlockAdapter for DiT models absent from cache-dit's BlockAdapterRegister.
|
@dataclass(frozen=True)
|
||||||
# Value: (blocks attr, forward_pattern). forward_pattern must
|
class CustomBlockAdapterSpec:
|
||||||
# match the block's forward signature (see cache_dit.ForwardPattern; e.g., ERNIE
|
blocks_attr: str
|
||||||
# uses Pattern_3). has_separate_cfg follows the run (passed by
|
forward_pattern: ForwardPattern
|
||||||
# enable_cache_on_transformer); cache-dit auto-resolves the remaining
|
|
||||||
# fields.
|
|
||||||
_CUSTOM_BLOCK_ADAPTER_SPECS: dict[str, tuple[str, ForwardPattern]] = {
|
# Custom BlockAdapter metadata for models absent from cache-dit's registry.
|
||||||
"ErnieImageTransformer2DModel": ("layers", ForwardPattern.Pattern_3),
|
_CUSTOM_BLOCK_ADAPTER_SPECS: dict[str, CustomBlockAdapterSpec] = {
|
||||||
"Krea2Transformer2DModel": ("transformer_blocks", ForwardPattern.Pattern_3),
|
"ErnieImageTransformer2DModel": CustomBlockAdapterSpec(
|
||||||
"MiniMaxH3DiTModel": ("blocks", ForwardPattern.Pattern_3),
|
blocks_attr="layers",
|
||||||
|
forward_pattern=ForwardPattern.Pattern_3,
|
||||||
|
),
|
||||||
|
"Krea2Transformer2DModel": CustomBlockAdapterSpec(
|
||||||
|
blocks_attr="transformer_blocks",
|
||||||
|
forward_pattern=ForwardPattern.Pattern_3,
|
||||||
|
),
|
||||||
|
"MiniMaxH3DiTModel": CustomBlockAdapterSpec(
|
||||||
|
blocks_attr="blocks",
|
||||||
|
forward_pattern=ForwardPattern.Pattern_3,
|
||||||
|
),
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
@@ -295,17 +305,16 @@ def _build_custom_block_adapter(
|
|||||||
spec = _CUSTOM_BLOCK_ADAPTER_SPECS.get(transformer.__class__.__name__)
|
spec = _CUSTOM_BLOCK_ADAPTER_SPECS.get(transformer.__class__.__name__)
|
||||||
if spec is None:
|
if spec is None:
|
||||||
return None
|
return None
|
||||||
blocks_attr, forward_pattern = spec
|
blocks = getattr(transformer, spec.blocks_attr, None)
|
||||||
blocks = getattr(transformer, blocks_attr, None)
|
|
||||||
if blocks is None:
|
if blocks is None:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Transformer {transformer.__class__.__name__} has no attribute "
|
f"Transformer {transformer.__class__.__name__} has no attribute "
|
||||||
f"{blocks_attr!r} for cache-dit blocks."
|
f"{spec.blocks_attr!r} for cache-dit blocks."
|
||||||
)
|
)
|
||||||
return BlockAdapter(
|
return BlockAdapter(
|
||||||
transformer=transformer,
|
transformer=transformer,
|
||||||
blocks=blocks,
|
blocks=blocks,
|
||||||
forward_pattern=forward_pattern,
|
forward_pattern=spec.forward_pattern,
|
||||||
has_separate_cfg=has_separate_cfg,
|
has_separate_cfg=has_separate_cfg,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -21,6 +21,7 @@ from sglang.kernels.ops.diffusion.qknorm_rope import (
|
|||||||
fused_inplace_qknorm_rope,
|
fused_inplace_qknorm_rope,
|
||||||
)
|
)
|
||||||
from sglang.kernels.ops.diffusion.triton.indexed_modulation import (
|
from sglang.kernels.ops.diffusion.triton.indexed_modulation import (
|
||||||
|
indexed_gate_bf16,
|
||||||
indexed_gate_bf16_,
|
indexed_gate_bf16_,
|
||||||
indexed_scale_shift_bf16_,
|
indexed_scale_shift_bf16_,
|
||||||
)
|
)
|
||||||
@@ -240,8 +241,9 @@ def _modulate_gate(
|
|||||||
indices: torch.Tensor,
|
indices: torch.Tensor,
|
||||||
*,
|
*,
|
||||||
dtype: torch.dtype,
|
dtype: torch.dtype,
|
||||||
|
allow_inplace: bool = True,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
"""Apply indexed gated residual, reusing disposable CUDA BF16 input."""
|
"""Apply an indexed gated residual, optionally reusing the input buffer."""
|
||||||
# Apply the per-index gated residual: x + gate[idx] * other.
|
# Apply the per-index gated residual: x + gate[idx] * other.
|
||||||
if (
|
if (
|
||||||
x.is_cuda
|
x.is_cuda
|
||||||
@@ -252,7 +254,9 @@ def _modulate_gate(
|
|||||||
and x.is_contiguous()
|
and x.is_contiguous()
|
||||||
and other.is_contiguous()
|
and other.is_contiguous()
|
||||||
):
|
):
|
||||||
return indexed_gate_bf16_(x, gate, other, indices)
|
if allow_inplace:
|
||||||
|
return indexed_gate_bf16_(x, gate, other, indices)
|
||||||
|
return indexed_gate_bf16(x, gate, other, indices)
|
||||||
return (x + gate.index_select(0, indices) * other).to(dtype)
|
return (x + gate.index_select(0, indices) * other).to(dtype)
|
||||||
|
|
||||||
|
|
||||||
@@ -904,6 +908,7 @@ class MiniMaxH3DiTBlock(nn.Module):
|
|||||||
expand_ratio=6,
|
expand_ratio=6,
|
||||||
modality_num=MINIMAX_H3_ADALN_MODALITY_NUM,
|
modality_num=MINIMAX_H3_ADALN_MODALITY_NUM,
|
||||||
)
|
)
|
||||||
|
self.preserve_input_for_cache_dit = False
|
||||||
|
|
||||||
def forward(
|
def forward(
|
||||||
self,
|
self,
|
||||||
@@ -929,7 +934,9 @@ class MiniMaxH3DiTBlock(nn.Module):
|
|||||||
if adaln_params is None:
|
if adaln_params is None:
|
||||||
adaln_params = self.adaln_proj(adaln_input)
|
adaln_params = self.adaln_proj(adaln_input)
|
||||||
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = adaln_params
|
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = adaln_params
|
||||||
|
# Cache-DiT retains the inputs to its Fn and Mn block ranges. Only the
|
||||||
|
# first gated residual writes to that tensor; the second one operates on
|
||||||
|
# a block-local buffer.
|
||||||
residual = x
|
residual = x
|
||||||
h = self.norm1(x)
|
h = self.norm1(x)
|
||||||
h = _modulate_scale_shift(
|
h = _modulate_scale_shift(
|
||||||
@@ -944,7 +951,14 @@ class MiniMaxH3DiTBlock(nn.Module):
|
|||||||
ulysses_active=ulysses_active,
|
ulysses_active=ulysses_active,
|
||||||
ring_active=ring_active,
|
ring_active=ring_active,
|
||||||
)
|
)
|
||||||
x = _modulate_gate(residual, gate_msa, h, combined_indices, dtype=_BF16_DTYPE)
|
x = _modulate_gate(
|
||||||
|
residual,
|
||||||
|
gate_msa,
|
||||||
|
h,
|
||||||
|
combined_indices,
|
||||||
|
dtype=_BF16_DTYPE,
|
||||||
|
allow_inplace=not self.preserve_input_for_cache_dit,
|
||||||
|
)
|
||||||
|
|
||||||
residual = x
|
residual = x
|
||||||
h = self.norm2(x)
|
h = self.norm2(x)
|
||||||
@@ -952,8 +966,14 @@ class MiniMaxH3DiTBlock(nn.Module):
|
|||||||
h, shift_mlp, scale_mlp, combined_indices, dtype=_BF16_DTYPE
|
h, shift_mlp, scale_mlp, combined_indices, dtype=_BF16_DTYPE
|
||||||
)
|
)
|
||||||
h = self.mlp(h)
|
h = self.mlp(h)
|
||||||
|
# `residual` is block-local here (see above), so this stays in-place
|
||||||
|
# even while Cache-DiT is attached.
|
||||||
return _modulate_gate(
|
return _modulate_gate(
|
||||||
residual, gate_mlp, h, combined_indices, dtype=_BF16_DTYPE
|
residual,
|
||||||
|
gate_mlp,
|
||||||
|
h,
|
||||||
|
combined_indices,
|
||||||
|
dtype=_BF16_DTYPE,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -1190,6 +1210,22 @@ class MiniMaxH3DiTModel(BaseDiT, LayerwiseOffloadableModuleMixin):
|
|||||||
self._resolved_attention_backend: AttentionBackendEnum | None = None
|
self._resolved_attention_backend: AttentionBackendEnum | None = None
|
||||||
self._mark_missing_params_required()
|
self._mark_missing_params_required()
|
||||||
|
|
||||||
|
def set_cache_dit_input_preservation(self, enabled: bool) -> None:
|
||||||
|
"""Stop the blocks from overwriting the input Cache-DiT holds by reference.
|
||||||
|
|
||||||
|
Cache-DiT snapshots the block-stack input to measure its residuals, so a
|
||||||
|
block that rewrites its own input in place makes that residual read as
|
||||||
|
zero. Only the first gated residual of a block writes the block input;
|
||||||
|
the second one operates on a buffer this block just allocated, so it is
|
||||||
|
left on the in-place fused path either way.
|
||||||
|
|
||||||
|
The caller owns the lifecycle. It has to be on before Cache-DiT mounts,
|
||||||
|
because mounting replaces `blocks` with a wrapper and the real blocks
|
||||||
|
stop being reachable by iterating it.
|
||||||
|
"""
|
||||||
|
for block in self.blocks:
|
||||||
|
block.preserve_input_for_cache_dit = enabled
|
||||||
|
|
||||||
def _resolve_attention_backend_once(self) -> None:
|
def _resolve_attention_backend_once(self) -> None:
|
||||||
if self._resolved_attention_backend is not None:
|
if self._resolved_attention_backend is not None:
|
||||||
return
|
return
|
||||||
|
|||||||
+79
-1
@@ -417,16 +417,94 @@ class MiniMaxH3DenoisingStage(DenoisingStage):
|
|||||||
# a time. Combined with `quality` in the dynamic-batch signature, this
|
# a time. Combined with `quality` in the dynamic-batch signature, this
|
||||||
# makes the process-wide hook transition safe at this batch boundary.
|
# makes the process-wide hook transition safe at this batch boundary.
|
||||||
if self._cache_dit_enabled and current_mode != desired_mode:
|
if self._cache_dit_enabled and current_mode != desired_mode:
|
||||||
|
# Unmount first: the blocks must stay preserved for as long as
|
||||||
|
# Cache-DiT still holds references to their inputs. Settle the state
|
||||||
|
# fields before restoring the in-place path, so a failure there
|
||||||
|
# costs throughput rather than leaving the stage inconsistent.
|
||||||
self.transformer = disable_cache_on_transformer(self.transformer)
|
self.transformer = disable_cache_on_transformer(self.transformer)
|
||||||
self._cache_dit_enabled = False
|
self._cache_dit_enabled = False
|
||||||
self._cached_num_steps = None
|
self._cached_num_steps = None
|
||||||
self._minimax_h3_cache_mode = None
|
self._minimax_h3_cache_mode = None
|
||||||
|
self._set_cache_dit_input_preservation(False)
|
||||||
|
|
||||||
if desired_mode is None:
|
if desired_mode is None:
|
||||||
return
|
return
|
||||||
super()._maybe_enable_cache_dit(num_inference_steps, batch)
|
|
||||||
|
# Arm before delegating whenever this H3 stage requests caching,
|
||||||
|
# without predicting whether the parent will accept the mount.
|
||||||
|
# cache_dit.enable_cache swaps `blocks` for a single CachedBlocks
|
||||||
|
# wrapper, so the real blocks are only reachable beforehand -- and a
|
||||||
|
# wrong prediction would hand Cache-DiT unpreserved blocks, whose
|
||||||
|
# residuals read as zero, which is silent. Arming is one boolean per
|
||||||
|
# block and nothing runs before the parent decides, so guessing
|
||||||
|
# conservatively costs approximately nothing.
|
||||||
|
was_enabled = self._cache_dit_enabled
|
||||||
|
if not was_enabled:
|
||||||
|
self._set_cache_dit_input_preservation(True)
|
||||||
|
try:
|
||||||
|
super()._maybe_enable_cache_dit(num_inference_steps, batch)
|
||||||
|
except Exception:
|
||||||
|
if not was_enabled:
|
||||||
|
self._disarm_after_failed_mount()
|
||||||
|
raise
|
||||||
|
|
||||||
if self._cache_dit_enabled:
|
if self._cache_dit_enabled:
|
||||||
self._minimax_h3_cache_mode = desired_mode
|
self._minimax_h3_cache_mode = desired_mode
|
||||||
|
elif not was_enabled:
|
||||||
|
# The parent declined to mount, for example because breakable
|
||||||
|
# CUDA graphs are enabled or this is an ordinary warmup. Nothing
|
||||||
|
# holds the block inputs, so go back to the in-place path.
|
||||||
|
self._set_cache_dit_input_preservation(False)
|
||||||
|
|
||||||
|
def _disarm_after_failed_mount(self) -> None:
|
||||||
|
"""Restore the in-place path only once Cache-DiT is confirmed gone.
|
||||||
|
|
||||||
|
cache_dit.enable_cache swaps the block list before the rest of the mount
|
||||||
|
runs, so a later failure can leave caching attached. Disarming then
|
||||||
|
would hand Cache-DiT unpreserved blocks and silently reproduce the
|
||||||
|
zero-residual bug, so if the unmount does not succeed we stay armed and
|
||||||
|
pay throughput instead.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
self.transformer = disable_cache_on_transformer(self.transformer)
|
||||||
|
except Exception:
|
||||||
|
logger.warning(
|
||||||
|
"Could not unmount Cache-DiT after a failed mount; leaving "
|
||||||
|
"MiniMax-H3 input preservation on",
|
||||||
|
exc_info=True,
|
||||||
|
)
|
||||||
|
return
|
||||||
|
# The parent may have flipped these before failing; leaving them set
|
||||||
|
# would send the next request down the refresh path with nothing
|
||||||
|
# mounted.
|
||||||
|
self._cache_dit_enabled = False
|
||||||
|
self._cached_num_steps = None
|
||||||
|
self._minimax_h3_cache_mode = None
|
||||||
|
self._set_cache_dit_input_preservation(False)
|
||||||
|
|
||||||
|
def _set_cache_dit_input_preservation(self, enabled: bool) -> None:
|
||||||
|
"""Flip input preservation on the H3 model, failing closed.
|
||||||
|
|
||||||
|
Skipping this silently would put us back where this stage started: cache
|
||||||
|
mounted, residuals reading as zero, no hits and nothing logged. If the
|
||||||
|
model cannot be reached, that is a wiring bug and should surface here.
|
||||||
|
"""
|
||||||
|
model = self.transformer
|
||||||
|
for _ in range(4): # unwrap compile/parallel wrappers, if any
|
||||||
|
if hasattr(model, "set_cache_dit_input_preservation"):
|
||||||
|
break
|
||||||
|
inner = getattr(model, "_orig_mod", None) or getattr(model, "module", None)
|
||||||
|
if inner is None:
|
||||||
|
break
|
||||||
|
model = inner
|
||||||
|
setter = getattr(model, "set_cache_dit_input_preservation", None)
|
||||||
|
if not callable(setter):
|
||||||
|
raise TypeError(
|
||||||
|
"MiniMax-H3 Cache-DiT requires set_cache_dit_input_preservation() "
|
||||||
|
f"on the transformer, but {type(self.transformer).__name__} does "
|
||||||
|
"not expose it"
|
||||||
|
)
|
||||||
|
setter(enabled)
|
||||||
|
|
||||||
def _cache_dit_scm_masks(
|
def _cache_dit_scm_masks(
|
||||||
self, primary_num_steps: int, secondary_num_steps: int | None = None
|
self, primary_num_steps: int, secondary_num_steps: int | None = None
|
||||||
|
|||||||
@@ -25,8 +25,10 @@ from sglang.multimodal_gen.runtime.loader.utils import get_param_names_mapping
|
|||||||
from sglang.multimodal_gen.runtime.models.dits.minimax_h3 import (
|
from sglang.multimodal_gen.runtime.models.dits.minimax_h3 import (
|
||||||
MINIMAX_H3_FP32_BUFFER_NAMES,
|
MINIMAX_H3_FP32_BUFFER_NAMES,
|
||||||
MINIMAX_H3_FP32_PARAM_NAMES,
|
MINIMAX_H3_FP32_PARAM_NAMES,
|
||||||
|
MiniMaxH3DiTBlock,
|
||||||
MiniMaxH3DiTModel,
|
MiniMaxH3DiTModel,
|
||||||
_copy_grouped_qkv_tp_shard,
|
_copy_grouped_qkv_tp_shard,
|
||||||
|
_modulate_gate,
|
||||||
_reorder_grouped_qkv_to_qkv,
|
_reorder_grouped_qkv_to_qkv,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.test.single_test_file.component_accuracy.utils import (
|
from sglang.multimodal_gen.test.single_test_file.component_accuracy.utils import (
|
||||||
@@ -127,6 +129,97 @@ def test_native_weight_names_and_grouped_qkv_reorder():
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class _KwargIdentity(torch.nn.Module):
|
||||||
|
def forward(self, x, **_kwargs):
|
||||||
|
return x
|
||||||
|
|
||||||
|
|
||||||
|
def test_cache_dit_preservation_only_makes_first_gate_out_of_place():
|
||||||
|
block = MiniMaxH3DiTBlock.__new__(MiniMaxH3DiTBlock)
|
||||||
|
torch.nn.Module.__init__(block)
|
||||||
|
block.norm1 = torch.nn.Identity()
|
||||||
|
block.norm2 = torch.nn.Identity()
|
||||||
|
block.attn = _KwargIdentity()
|
||||||
|
block.mlp = torch.nn.Identity()
|
||||||
|
gate_modes = []
|
||||||
|
|
||||||
|
def fake_gate(residual, _gate, _other, _indices, *, dtype, allow_inplace=True):
|
||||||
|
gate_modes.append(allow_inplace)
|
||||||
|
return residual.to(dtype)
|
||||||
|
|
||||||
|
def run(preserve):
|
||||||
|
block.preserve_input_for_cache_dit = preserve
|
||||||
|
gate_modes.clear()
|
||||||
|
block(
|
||||||
|
torch.zeros(2, 4),
|
||||||
|
adaln_input=torch.zeros(1, 4),
|
||||||
|
combined_indices=torch.zeros(2, dtype=torch.long),
|
||||||
|
rope_cache=None,
|
||||||
|
cu_seqlens=torch.tensor([0, 2], dtype=torch.int32),
|
||||||
|
max_seqlen=2,
|
||||||
|
adaln_params=tuple(torch.zeros(1, 4) for _ in range(6)),
|
||||||
|
)
|
||||||
|
return list(gate_modes)
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch(
|
||||||
|
"sglang.multimodal_gen.runtime.models.dits.minimax_h3._modulate_scale_shift",
|
||||||
|
side_effect=lambda value, *_args, **_kwargs: value,
|
||||||
|
),
|
||||||
|
patch(
|
||||||
|
"sglang.multimodal_gen.runtime.models.dits.minimax_h3._modulate_gate",
|
||||||
|
side_effect=fake_gate,
|
||||||
|
),
|
||||||
|
):
|
||||||
|
assert run(preserve=False) == [True, True]
|
||||||
|
# Only the first gated residual can alias the block input Cache-DiT
|
||||||
|
# holds by reference, so only it goes out-of-place. The second works on
|
||||||
|
# a block-local buffer and keeps the fused in-place kernel.
|
||||||
|
assert run(preserve=True) == [False, True]
|
||||||
|
|
||||||
|
|
||||||
|
def test_cache_dit_input_preservation_toggles_every_block():
|
||||||
|
model = MiniMaxH3DiTModel.__new__(MiniMaxH3DiTModel)
|
||||||
|
torch.nn.Module.__init__(model)
|
||||||
|
model.blocks = torch.nn.ModuleList([torch.nn.Identity() for _ in range(5)])
|
||||||
|
|
||||||
|
model.set_cache_dit_input_preservation(True)
|
||||||
|
assert all(block.preserve_input_for_cache_dit for block in model.blocks)
|
||||||
|
|
||||||
|
model.set_cache_dit_input_preservation(False)
|
||||||
|
assert not any(block.preserve_input_for_cache_dit for block in model.blocks)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA")
|
||||||
|
def test_cache_dit_out_of_place_gate_preserves_cuda_input():
|
||||||
|
x = torch.randn(4, 16, device="cuda", dtype=torch.bfloat16)
|
||||||
|
original = x.clone()
|
||||||
|
gate = torch.randn(2, 16, device="cuda", dtype=torch.bfloat16)
|
||||||
|
other = torch.randn_like(x)
|
||||||
|
indices = torch.tensor([0, 1, 0, 1], device="cuda", dtype=torch.long)
|
||||||
|
expected = _modulate_gate(
|
||||||
|
x.clone(),
|
||||||
|
gate,
|
||||||
|
other,
|
||||||
|
indices,
|
||||||
|
dtype=torch.bfloat16,
|
||||||
|
allow_inplace=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
output = _modulate_gate(
|
||||||
|
x,
|
||||||
|
gate,
|
||||||
|
other,
|
||||||
|
indices,
|
||||||
|
dtype=torch.bfloat16,
|
||||||
|
allow_inplace=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert output.data_ptr() != x.data_ptr()
|
||||||
|
torch.testing.assert_close(x, original, rtol=0, atol=0)
|
||||||
|
torch.testing.assert_close(output, expected, rtol=0, atol=0)
|
||||||
|
|
||||||
|
|
||||||
def test_tp_and_ulysses_admission_uses_tp_local_shapes():
|
def test_tp_and_ulysses_admission_uses_tp_local_shapes():
|
||||||
arch = MiniMaxH3DiTArchConfig()
|
arch = MiniMaxH3DiTArchConfig()
|
||||||
model = MiniMaxH3DiTModel.__new__(MiniMaxH3DiTModel)
|
model = MiniMaxH3DiTModel.__new__(MiniMaxH3DiTModel)
|
||||||
|
|||||||
Reference in New Issue
Block a user