diff --git a/docs/cookbook/diffusion/MiniMax/MiniMax-H3.mdx b/docs/cookbook/diffusion/MiniMax/MiniMax-H3.mdx index 278f569c6..9b7b032e8 100644 --- a/docs/cookbook/diffusion/MiniMax/MiniMax-H3.mdx +++ b/docs/cookbook/diffusion/MiniMax/MiniMax-H3.mdx @@ -314,6 +314,12 @@ honors each layer's activation mode; omit `--quantization`. Mixed exports use the same command: SGLang dispatches each marked layer to its serialized W4A4 or INT8 ConvRot kernel instead of applying one global method. +Comfy NVFP4 DiTs reuse SGLang's ModelOpt NVFP4 backend, which requires CUDA +compute capability 10.0 or newer. Pass a pruned FL2VA / Ref2VA file such as +`Abiray/Minimax-H3-nvfp4-INT4-INT8-Convrot/MiniMax_H3_FL2VA_pruned_nvfp4.safetensors` +to `--transformer-weights-path` and omit `--quantization`. SGLang infers the +packed group size and Comfy scale layout from the checkpoint; FSDP is rejected. + ### Advanced: precomputed AdaLN cache The [model card](https://huggingface.co/MiniMaxAI/MiniMax-H3) notes that about diff --git a/docs/docs/sglang-diffusion/quantization.mdx b/docs/docs/sglang-diffusion/quantization.mdx index 136dc9d69..8a0faae09 100644 --- a/docs/docs/sglang-diffusion/quantization.mdx +++ b/docs/docs/sglang-diffusion/quantization.mdx @@ -154,11 +154,11 @@ backend. modelopt-nvfp4 - Mixed transformer directory/repo with config.json, raw NVFP4 safetensors export/repo, or full ModelOpt Diffusers repo + Mixed transformer directory/repo with config.json, raw or Comfy layer-marked NVFP4 safetensors, or full ModelOpt Diffusers repo --transformer-path for mixed overrides; --transformer-weights-path for raw exports; --model-path for full repos - FLUX.1, FLUX.2, Wan2.2, Qwen Image, Qwen Image 2512, Qwen Image Edit, Qwen Image Edit 2511 + FLUX.1, FLUX.2, Wan2.2, Qwen Image, Qwen Image 2512, Qwen Image Edit, Qwen Image Edit 2511, MiniMax-H3 None - Mixed override repos keep the base model separate; full Qwen Image exports can be loaded directly as --model-path; raw exports such as black-forest-labs/FLUX.2-dev-NVFP4 still use the weights-path flow + Mixed override repos keep the base model separate; full Qwen Image exports can be loaded directly as --model-path; raw exports such as black-forest-labs/FLUX.2-dev-NVFP4 use the weights-path flow. Comfy markers select their checkpoint layout automatically; omit --quantization. gguf diff --git a/python/sglang/multimodal_gen/runtime/layers/quantization/configs/base_config.py b/python/sglang/multimodal_gen/runtime/layers/quantization/configs/base_config.py index cdbca71fc..a2adaeb6e 100644 --- a/python/sglang/multimodal_gen/runtime/layers/quantization/configs/base_config.py +++ b/python/sglang/multimodal_gen/runtime/layers/quantization/configs/base_config.py @@ -36,6 +36,7 @@ class QuantizationConfig(SRTQuantizationConfig): # for quantization frameworks with a separate quantized model provided, e.g. Nunchaku quantized_model_path: str | None = None checkpoint_uses_native_qkv_layout: bool = False + checkpoint_uses_comfy_quantization: bool = False supports_srt_linear_layers: bool = False def get_scaled_act_names(self) -> list[str]: 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 a6c7b80a8..b10c5aa12 100755 --- a/python/sglang/multimodal_gen/runtime/layers/quantization/modelopt_quant.py +++ b/python/sglang/multimodal_gen/runtime/layers/quantization/modelopt_quant.py @@ -236,6 +236,7 @@ class ModelOptFp4Config(ModelOptQuantConfig): checkpoint_uses_packed_qkv: bool = False, swap_weight_nibbles: bool = False, checkpoint_weight_scale_layout: str = "linear", + checkpoint_uses_comfy_quantization: bool = False, ) -> None: super().__init__(exclude_modules, packed_modules_mapping) self.is_checkpoint_nvfp4_serialized = is_checkpoint_nvfp4_serialized @@ -248,6 +249,7 @@ class ModelOptFp4Config(ModelOptQuantConfig): self.checkpoint_uses_packed_qkv = checkpoint_uses_packed_qkv self.swap_weight_nibbles = swap_weight_nibbles self.checkpoint_weight_scale_layout = checkpoint_weight_scale_layout + self.checkpoint_uses_comfy_quantization = checkpoint_uses_comfy_quantization @classmethod def get_name(cls) -> str: @@ -348,6 +350,9 @@ class ModelOptFp4Config(ModelOptQuantConfig): checkpoint_weight_scale_layout=config.get( "checkpoint_weight_scale_layout", "linear" ), + checkpoint_uses_comfy_quantization=config.get( + "checkpoint_uses_comfy_quantization", False + ), ) def get_quant_method(self, layer: torch.nn.Module, prefix: str): diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/transformer_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/transformer_loader.py index 251dc6bad..34d0eda5d 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/transformer_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/transformer_loader.py @@ -248,7 +248,10 @@ class TransformerLoader(ComponentLoader): safetensors_list ) checkpoint_quant_config = resolve_minimax_h3_checkpoint_quantization( - layer_markers + layer_markers, + safetensors_list, + dit_config.arch_config.param_names_mapping, + dit_config.arch_config.reverse_param_names_mapping, ) if adaln_curve_shape is not None: ( diff --git a/python/sglang/multimodal_gen/runtime/loader/minimax_h3_weights.py b/python/sglang/multimodal_gen/runtime/loader/minimax_h3_weights.py index 7afda25a9..ad9059c05 100644 --- a/python/sglang/multimodal_gen/runtime/loader/minimax_h3_weights.py +++ b/python/sglang/multimodal_gen/runtime/loader/minimax_h3_weights.py @@ -9,6 +9,7 @@ from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config impor QuantizationConfig, ) from sglang.multimodal_gen.runtime.utils.quantization_utils import ( + build_nvfp4_config_from_safetensors_list, inspect_comfy_quant_markers, resolve_comfy_checkpoint_quantization, ) @@ -47,7 +48,26 @@ def inspect_minimax_h3_safetensors( def resolve_minimax_h3_checkpoint_quantization( layer_markers: dict[str, dict[str, Any]], + safetensors_list: list[str] | None = None, + param_names_mapping: dict | None = None, + reverse_param_names_mapping: dict | None = None, ) -> QuantizationConfig | None: + formats = {str(marker.get("format")) for marker in layer_markers.values()} + if formats == {"nvfp4"}: + if safetensors_list is None: + raise ValueError("MiniMax-H3 NVFP4 metadata requires checkpoint files") + config = build_nvfp4_config_from_safetensors_list( + safetensors_list, + param_names_mapping, + reverse_param_names_mapping, + ) + if config is None: + raise ValueError("Could not resolve MiniMax-H3 NVFP4 checkpoint layout") + config.checkpoint_uses_comfy_quantization = True + config.checkpoint_uses_native_qkv_layout = True + config.checkpoint_weight_scale_layout = "swizzled" + config.swap_weight_nibbles = True + return config return resolve_comfy_checkpoint_quantization(layer_markers) diff --git a/python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py b/python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py index f98a89a26..3216ce3bc 100644 --- a/python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py +++ b/python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py @@ -141,6 +141,14 @@ def _merge_modelopt_fp4_configs( ) if getattr(inferred_config, "group_size", None) is None: inferred_config.group_size = getattr(existing_config, "group_size", None) + inferred_config.checkpoint_uses_comfy_quantization = ( + inferred_config.checkpoint_uses_comfy_quantization + or existing_config.checkpoint_uses_comfy_quantization + ) + inferred_config.checkpoint_uses_native_qkv_layout = ( + inferred_config.checkpoint_uses_native_qkv_layout + or existing_config.checkpoint_uses_native_qkv_layout + ) return inferred_config @@ -194,6 +202,10 @@ class TransformerQuantLoadSpec: or self.is_serialized_kitchen_int8 or self.is_serialized_kitchen_w4a4 or self.is_serialized_kitchen_w4a8 + or ( + self.quant_config is not None + and self.quant_config.checkpoint_uses_comfy_quantization + ) or ( _get_quant_config_name(self.quant_config) == "mxfp8" and self.quant_config.layer_markers is not None diff --git a/python/sglang/multimodal_gen/runtime/utils/quantization_utils.py b/python/sglang/multimodal_gen/runtime/utils/quantization_utils.py index dd91f4a20..62e11954b 100644 --- a/python/sglang/multimodal_gen/runtime/utils/quantization_utils.py +++ b/python/sglang/multimodal_gen/runtime/utils/quantization_utils.py @@ -674,11 +674,12 @@ def _build_nvfp4_config_from_safetensors_files( and "layers" in quant_config_dict ): layers = quant_config_dict.get("layers", {}) - file_quantized_modules.update( + metadata_nvfp4_modules = { layer_name for layer_name, layer_cfg in layers.items() if isinstance(layer_cfg, dict) and layer_cfg.get("format") == "nvfp4" - ) + } + file_quantized_modules.update(metadata_nvfp4_modules) tensor_metadata = _read_safetensors_tensor_metadata(file_path) with safe_open(file_path, framework="pt", device="cpu") as f: @@ -815,6 +816,7 @@ def _build_nvfp4_config_from_safetensors_files( "swizzled" if checkpoint_uses_swizzled_scales else "linear" ), "swap_weight_nibbles": checkpoint_uses_swizzled_scales, + "checkpoint_uses_comfy_quantization": checkpoint_uses_comfy_quant, } ) logger.info( 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 55801716c..ebf124ae3 100644 --- a/python/sglang/multimodal_gen/test/unit/test_transformer_quant.py +++ b/python/sglang/multimodal_gen/test/unit/test_transformer_quant.py @@ -1105,6 +1105,16 @@ class TestTransformerQuantHelpers(unittest.TestCase): self.assertTrue(config.load_in_4bit) def test_nvfp4_safetensors_inference_ignores_fp8_fallback_scales(self): + metadata = { + "_quantization_metadata": json.dumps( + { + "format_version": "1.0", + "layers": { + "layers.0.attention.qkv": {"format": "nvfp4"}, + }, + } + ) + } with tempfile.NamedTemporaryFile(suffix=".safetensors") as f: save_file( { @@ -1127,6 +1137,7 @@ class TestTransformerQuantHelpers(unittest.TestCase): ), }, f.name, + metadata=metadata, ) config = build_nvfp4_config_from_safetensors_list([f.name]) @@ -1137,6 +1148,7 @@ class TestTransformerQuantHelpers(unittest.TestCase): self.assertNotIn("layers.0.attention.qkv", config.exclude_modules) self.assertEqual(config.checkpoint_weight_scale_layout, "linear") self.assertFalse(config.swap_weight_nibbles) + self.assertFalse(config.checkpoint_uses_comfy_quantization) def test_nvfp4_safetensors_inference_uses_comfy_checkpoint_layout(self): with tempfile.NamedTemporaryFile(suffix=".safetensors") as f: @@ -1179,6 +1191,62 @@ class TestTransformerQuantHelpers(unittest.TestCase): self.assertNotIn("layers.0.attention.qkv", config.exclude_modules) self.assertEqual(config.checkpoint_weight_scale_layout, "swizzled") self.assertTrue(config.swap_weight_nibbles) + self.assertTrue(config.checkpoint_uses_comfy_quantization) + self.assertFalse(config.checkpoint_uses_native_qkv_layout) + spec = TransformerQuantLoadSpec( + safetensors_list=[f.name], + quant_config=config, + nunchaku_config=None, + param_dtype=None, + ) + self.assertTrue(spec.uses_comfy_layer_markers) + + def test_minimax_h3_comfy_nvfp4_resolves_modelopt_backend(self): + metadata = { + "_quantization_metadata": json.dumps( + { + "format_version": "1.0", + "layers": { + "blocks.0.attn.qkv_proj": {"format": "nvfp4"}, + }, + } + ) + } + with ( + tempfile.NamedTemporaryFile(suffix=".safetensors") as quantized, + tempfile.NamedTemporaryFile(suffix=".safetensors") as fallback, + ): + save_file( + { + "blocks.0.attn.qkv_proj.weight": torch.zeros( + (32, 8), dtype=torch.uint8 + ), + "blocks.0.attn.qkv_proj.weight_scale": torch.ones( + (32, 1), dtype=torch.float8_e4m3fn + ), + "blocks.0.attn.qkv_proj.weight_scale_2": torch.tensor(1.0), + }, + quantized.name, + metadata=metadata, + ) + save_file( + {"blocks.0.mlp.fc1.weight": torch.ones((2, 2))}, + fallback.name, + ) + checkpoint_files = [quantized.name, fallback.name] + _, markers = inspect_minimax_h3_safetensors(checkpoint_files) + config = resolve_minimax_h3_checkpoint_quantization( + markers, + checkpoint_files, + ) + + self.assertIsInstance(config, ModelOptFp4Config) + self.assertEqual(config.group_size, 16) + self.assertIn("blocks.0.mlp.fc1", config.exclude_modules) + self.assertTrue(config.checkpoint_uses_comfy_quantization) + self.assertTrue(config.checkpoint_uses_native_qkv_layout) + self.assertEqual(config.checkpoint_weight_scale_layout, "swizzled") + self.assertTrue(config.swap_weight_nibbles) def test_builder_adds_diffusers_quant_type_for_nvfp4(self): updated = _updated_quant_config(