diff --git a/python/sglang/multimodal_gen/configs/models/dits/minimax_h3.py b/python/sglang/multimodal_gen/configs/models/dits/minimax_h3.py index d773411e0..925f08934 100644 --- a/python/sglang/multimodal_gen/configs/models/dits/minimax_h3.py +++ b/python/sglang/multimodal_gen/configs/models/dits/minimax_h3.py @@ -18,6 +18,7 @@ class MiniMaxH3DiTArchConfig(DiTArchConfig): _supported_attention_backends: set[AttentionBackendEnum] = field( default_factory=lambda: { AttentionBackendEnum.FA, + AttentionBackendEnum.SAGE_ATTN, AttentionBackendEnum.AITER, AttentionBackendEnum.TORCH_SDPA, } diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/minimax_h3.py b/python/sglang/multimodal_gen/configs/pipeline_configs/minimax_h3.py index d3e6e4d03..2fe52d8d2 100644 --- a/python/sglang/multimodal_gen/configs/pipeline_configs/minimax_h3.py +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/minimax_h3.py @@ -169,12 +169,6 @@ class MiniMaxH3PipelineConfig(PipelineConfig): def validate_server_args(self, server_args) -> None: # Reject known-inexact VAE modes before any large component download. self.vae_config.resolved_parallel_decode_mode() - attention_backend = self._server_arg_value(server_args.attention_backend) - if str(attention_backend).strip().lower() == "sage_attn": - raise ValueError( - "MiniMax-H3 does not support SageAttention: the current packed " - "varlen path does not preserve model output" - ) def select_vae_weight_files( self, diff --git a/python/sglang/multimodal_gen/runtime/layers/attention/backends/sage_attn.py b/python/sglang/multimodal_gen/runtime/layers/attention/backends/sage_attn.py index 973608f64..ad785dd4a 100644 --- a/python/sglang/multimodal_gen/runtime/layers/attention/backends/sage_attn.py +++ b/python/sglang/multimodal_gen/runtime/layers/attention/backends/sage_attn.py @@ -17,6 +17,21 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger logger = init_logger(__name__) +def _trailing_padding_used_len( + *, + total_tokens: int, + max_seqlen: int, + bounds: tuple[int, ...], +) -> int | None: + """Return live token count for H3-style [0, used, total] trailing padding.""" + if len(bounds) != 3: + return None + start, used, total = bounds + if start != 0 or used >= total or total != total_tokens or used != max_seqlen: + return None + return used + + class SageAttentionBackend(AttentionBackend): accept_output_buffer: bool = True @@ -72,3 +87,68 @@ class SageAttentionImpl(AttentionImpl): output, softmax_lse = output return output, softmax_lse return output + + def forward_varlen( + self, + query: torch.Tensor, + key: torch.Tensor, + value: torch.Tensor, + *, + cu_seqlens: torch.Tensor, + max_seqlen: int, + cu_seqlens_host: tuple[int, ...] | None = None, + ) -> torch.Tensor: + bounds = ( + cu_seqlens_host + if cu_seqlens_host is not None + else tuple(int(x) for x in cu_seqlens.tolist()) + ) + return self._sage_packed( + query.contiguous(), + key.contiguous(), + value.contiguous(), + bounds=bounds, + max_seqlen=max_seqlen, + ) + + def _sage_packed( + self, + query: torch.Tensor, + key: torch.Tensor, + value: torch.Tensor, + *, + bounds: tuple[int, ...], + max_seqlen: int, + ) -> torch.Tensor: + # MiniMax-H3 packs one live document as bounds=(0, used, total): + # [0, used) are real tokens; [used, total) is 64-aligned tail padding. + used = _trailing_padding_used_len( + total_tokens=query.shape[0], + max_seqlen=max_seqlen, + bounds=bounds, + ) + if used is not None: + live_out = self.forward( + query[:used].unsqueeze(0), + key[:used].unsqueeze(0), + value[:used].unsqueeze(0), + None, + )[0] + if used == query.shape[0]: + return live_out + # Keep padded tail at zero so downstream masked rows stay inactive. + output = torch.zeros_like(query) + output[:used] = live_out + return output + + output = torch.empty_like(query) + for start, stop in zip(bounds[:-1], bounds[1:]): + if start == stop: + continue + output[start:stop] = self.forward( + query[start:stop].unsqueeze(0), + key[start:stop].unsqueeze(0), + value[start:stop].unsqueeze(0), + None, + )[0] + return output diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/minimax_h3/release_metadata.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/minimax_h3/release_metadata.py index 475de96aa..eeba8b353 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/minimax_h3/release_metadata.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/minimax_h3/release_metadata.py @@ -153,12 +153,6 @@ class MiniMaxH3PartitionAdmissionStage(PipelineStage): f"quality must be one of {list(QUALITY_LEVELS)}, got {quality!r}" ) high_quality = quality == "high" - attention_backend = str(server_args.attention_backend or "").strip().lower() - if attention_backend == "sage_attn" and not batch.is_warmup: - raise ValueError( - "MiniMax-H3 does not support SageAttention: the current packed " - "varlen path does not preserve model output" - ) if high_quality and not batch.is_warmup: server_args.pipeline_config.validate_quality_deployment(server_args) plan = minimax_h3_plan_from_batch(batch) diff --git a/python/sglang/multimodal_gen/test/unit/test_minimax_h3_admission.py b/python/sglang/multimodal_gen/test/unit/test_minimax_h3_admission.py index 19b8ee681..181eb496a 100644 --- a/python/sglang/multimodal_gen/test/unit/test_minimax_h3_admission.py +++ b/python/sglang/multimodal_gen/test/unit/test_minimax_h3_admission.py @@ -301,8 +301,7 @@ def test_quality_admission_fails_closed_outside_validated_request(): batch.sampling_params.quality = "lossless" batch.num_inference_steps = 50 server_args.attention_backend = "sage_attn" - with pytest.raises(ValueError, match="does not support SageAttention"): - stage.forward(batch, server_args) + assert stage.forward(batch, server_args) is batch batch.sampling_params.quality = "ultra" server_args.attention_backend = None