diff --git a/python/sglang/multimodal_gen/runtime/layers/quantization/modelopt_quant.py b/python/sglang/multimodal_gen/runtime/layers/quantization/modelopt_quant.py index 9049ed076..a6c7b80a8 100755 --- a/python/sglang/multimodal_gen/runtime/layers/quantization/modelopt_quant.py +++ b/python/sglang/multimodal_gen/runtime/layers/quantization/modelopt_quant.py @@ -462,6 +462,20 @@ class ModelOptFp4LinearMethod(LinearMethodBase): """NVFP4 linear method using the selected FP4 GEMM backend.""" def __init__(self, quant_config: ModelOptFp4Config): + # the FlashInfer FP4 kernels this method dispatches to are Blackwell-only. + # without this the load succeeds and the failure surfaces much later as + # an opaque CUDA error inside the kernel. an undetectable capability is + # left alone rather than rejected, so this only ever converts a crash + # into a message and never blocks a card that would have worked. + capability = current_platform.get_device_capability() + min_capability = quant_config.get_min_capability() + if capability is not None and capability.to_int() < min_capability: + raise RuntimeError( + f"NVFP4 checkpoints need compute capability " + f"{min_capability // 10}.{min_capability % 10} or newer " + f"(Blackwell); this GPU is {capability.as_version_str()}. " + f"Load an FP8 or BF16 checkpoint instead." + ) self.quant_config = quant_config def create_weights( diff --git a/python/sglang/multimodal_gen/runtime/models/vaes/minimax_h3_video_vae/klvae.py b/python/sglang/multimodal_gen/runtime/models/vaes/minimax_h3_video_vae/klvae.py index 9aafb6b3a..74a64733d 100644 --- a/python/sglang/multimodal_gen/runtime/models/vaes/minimax_h3_video_vae/klvae.py +++ b/python/sglang/multimodal_gen/runtime/models/vaes/minimax_h3_video_vae/klvae.py @@ -501,6 +501,17 @@ class AutoencoderKL(ModelMixin, ConfigMixin, FromOriginalModelMixin): dec = self._assemble_tiles(rows, y_overlap, x_overlap) return dec + def enable_tiling(self) -> None: + """Turn on tiled decode for subsequent decodes. + + `decoder_tiling` is read per decode in `_adaptive_decode`, so setting + it here takes effect on the next call. Models that already tile from + their VAE config (MiniMax-H3) are unaffected; this exists so the + runtime `--vae-tiling` switch reaches this VAE at all instead of + raising into the caller's guard. + """ + self.decoder_tiling = True + def _adaptive_encode(self, x): if self.encoder_tiling: return self.tiled_encode(x) diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/decoding.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/decoding.py index 5534530af..14590065b 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/decoding.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/decoding.py @@ -241,12 +241,18 @@ class DecodingStage(PipelineStage): # Decode latents with autocast_context(vae_dtype, server_args.disable_autocast): - try: - # TODO: make it more specific - if server_args.pipeline_config.vae_tiling: + # not every VAE supports toggling tiling at runtime; say so instead + # of dropping the request, since the OOM advice below points here + if server_args.pipeline_config.vae_tiling: + try: self.vae.enable_tiling() - except Exception: - pass + except AttributeError: + logger.warning( + "--vae-tiling has no effect: %s does not support " + "enabling tiling at runtime. Whether it tiles is fixed " + "by its VAE config.", + type(self.vae).__name__, + ) should_cast_vae = not vae_autocast_enabled if not vae_autocast_enabled: latents = latents.to(vae_dtype) @@ -257,16 +263,24 @@ class DecodingStage(PipelineStage): decode_output = self._get_vae_decode_fn(vae, server_args)(latents) except Exception as error: if "out of memory" in str(error).lower(): + # decode runs after denoising, so the DiT and encoders + # are idle but may still hold VRAM; freeing them is the + # lever here. --vae-cpu-offload is not: it moves VAE + # weights, not the activations that overflow. if not server_args.pipeline_config.vae_tiling: logger.warning( - "OOM detected during VAE decoding. Please enable " - "--vae-tiling to reduce peak memory usage." + "OOM detected during VAE decoding. Enable " + "--vae-tiling to bound the decode working set, " + "and free the components that finished earlier " + "with --cpu-offload-components dit,text_encoder." ) else: logger.warning( "OOM detected during VAE decoding with tiling enabled. " - "Please reduce the resolution or enable " - "--vae-cpu-offload." + "Free the components that finished earlier with " + "--cpu-offload-components dit,text_encoder, then " + "lower the tile size in the model's VAE config, " + "then reduce resolution or frame count." ) raise image = _ensure_tensor_decode_output(decode_output) diff --git a/python/sglang/multimodal_gen/test/unit/test_ideogram4.py b/python/sglang/multimodal_gen/test/unit/test_ideogram4.py index 5607a4007..e4a7f4e8f 100644 --- a/python/sglang/multimodal_gen/test/unit/test_ideogram4.py +++ b/python/sglang/multimodal_gen/test/unit/test_ideogram4.py @@ -101,6 +101,7 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.text_encoding import ( TextEncodingStage, ) from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum +from sglang.multimodal_gen.runtime.platforms.interface import DeviceCapability from sglang.multimodal_gen.runtime.server_args import set_global_server_args @@ -996,9 +997,15 @@ class TestIdeogram4(unittest.TestCase): sp_split_auto=False, ) ) - with patch( - "sglang.multimodal_gen.runtime.layers.attention.layer.get_ring_parallel_world_size", - return_value=1, + with ( + patch( + "sglang.multimodal_gen.runtime.layers.quantization.modelopt_quant.current_platform.get_device_capability", + return_value=DeviceCapability(10, 0), + ), + patch( + "sglang.multimodal_gen.runtime.layers.attention.layer.get_ring_parallel_world_size", + return_value=1, + ), ): with torch.device("meta"): model = Ideogram4Transformer2DModel( @@ -1058,6 +1065,10 @@ class TestIdeogram4(unittest.TestCase): ) ) with ( + patch( + "sglang.multimodal_gen.runtime.layers.quantization.modelopt_quant.current_platform.get_device_capability", + return_value=DeviceCapability(10, 0), + ), patch( "sglang.multimodal_gen.runtime.models.dits.ideogram.model_parallel_is_initialized", return_value=True, diff --git a/python/sglang/multimodal_gen/test/unit/test_modelopt_fp4_capability_gate.py b/python/sglang/multimodal_gen/test/unit/test_modelopt_fp4_capability_gate.py new file mode 100644 index 000000000..77a9eb08e --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_modelopt_fp4_capability_gate.py @@ -0,0 +1,54 @@ +"""NVFP4 rejects pre-Blackwell GPUs at load instead of inside the CUDA kernel.""" + +import unittest +from unittest.mock import patch + +from sglang.multimodal_gen.runtime.layers.quantization.modelopt_quant import ( + ModelOptFp4Config, + ModelOptFp4LinearMethod, +) +from sglang.multimodal_gen.runtime.platforms.interface import DeviceCapability + +CAPABILITY_PATH = ( + "sglang.multimodal_gen.runtime.layers.quantization." + "modelopt_quant.current_platform.get_device_capability" +) + + +def _config() -> ModelOptFp4Config: + return ModelOptFp4Config(is_checkpoint_nvfp4_serialized=True, group_size=16) + + +class TestModelOptFp4CapabilityGate(unittest.TestCase): + def test_rejects_pre_blackwell_with_an_actionable_message(self): + for major, minor, name in [(8, 6, "3090"), (8, 9, "4090"), (9, 0, "H100")]: + with self.subTest(gpu=name): + with patch( + CAPABILITY_PATH, return_value=DeviceCapability(major, minor) + ): + with self.assertRaises(RuntimeError) as caught: + ModelOptFp4LinearMethod(_config()) + message = str(caught.exception) + self.assertIn(f"{major}.{minor}", message) + self.assertIn("Blackwell", message) + + def test_allows_blackwell(self): + # 10.0 is B200/B300, 12.0 is the consumer Blackwell line; both carry the + # FlashInfer FP4 kernels, so neither may be rejected + for major, minor, name in [(10, 0, "B200"), (12, 0, "RTX 5090")]: + with self.subTest(gpu=name): + with patch( + CAPABILITY_PATH, return_value=DeviceCapability(major, minor) + ): + ModelOptFp4LinearMethod(_config()) + + def test_undetectable_capability_is_left_alone(self): + # the gate exists to turn a kernel crash into a message, so a GPU whose + # capability cannot be read keeps the previous behavior rather than + # gaining a new way to fail + with patch(CAPABILITY_PATH, return_value=None): + ModelOptFp4LinearMethod(_config()) + + +if __name__ == "__main__": + unittest.main() diff --git a/python/sglang/multimodal_gen/test/unit/test_transformer_quant.py b/python/sglang/multimodal_gen/test/unit/test_transformer_quant.py index b30eef561..d8d1131af 100644 --- a/python/sglang/multimodal_gen/test/unit/test_transformer_quant.py +++ b/python/sglang/multimodal_gen/test/unit/test_transformer_quant.py @@ -73,6 +73,7 @@ from sglang.multimodal_gen.runtime.loader.transformer_load_utils import ( from sglang.multimodal_gen.runtime.loader.weight_load_plan import WeightLoadPlan from sglang.multimodal_gen.runtime.models.dits.flux import FluxSingleTransformerBlock from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum +from sglang.multimodal_gen.runtime.platforms.interface import DeviceCapability from sglang.multimodal_gen.runtime.utils.quantization_utils import ( build_nvfp4_config_from_safetensors_list, get_quant_config, @@ -651,14 +652,18 @@ class TestTransformerQuantHelpers(unittest.TestCase): ], ) - block = FluxSingleTransformerBlock( - dim=64, - num_attention_heads=4, - attention_head_dim=16, - mlp_ratio=2.0, - quant_config=quant_config, - prefix="single_transformer_blocks.0", - ) + with patch( + "sglang.multimodal_gen.runtime.layers.quantization.modelopt_quant.current_platform.get_device_capability", + return_value=DeviceCapability(10, 0), + ): + block = FluxSingleTransformerBlock( + dim=64, + num_attention_heads=4, + attention_head_dim=16, + mlp_ratio=2.0, + quant_config=quant_config, + prefix="single_transformer_blocks.0", + ) self.assertEqual(block.proj_mlp.prefix, "single_transformer_blocks.0.proj_mlp") self.assertEqual(block.proj_out.prefix, "single_transformer_blocks.0.proj_out")