From d7c4c57ace5d28e775eb8593e43dd4cedcf7f212 Mon Sep 17 00:00:00 2001 From: Mick Date: Fri, 27 Mar 2026 09:58:49 +0800 Subject: [PATCH] [diffusion] refactor: move format-specific weight loading hooks (quant-related) to a dedicated file (#21366) --- .../nunchaku.py} | 85 +++-- .../component_loaders/transformer_loader.py | 156 ++------- .../runtime/loader/transformer_load_utils.py | 306 ++++++++++++++++++ .../runtime/pipelines/flux_2_nvfp4.py | 101 ++++-- .../multimodal_gen/runtime/platforms/cuda.py | 2 +- .../multimodal_gen/runtime/server_args.py | 27 +- 6 files changed, 476 insertions(+), 201 deletions(-) rename python/sglang/multimodal_gen/configs/{quantization.py => quantization/nunchaku.py} (69%) create mode 100644 python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py diff --git a/python/sglang/multimodal_gen/configs/quantization.py b/python/sglang/multimodal_gen/configs/quantization/nunchaku.py similarity index 69% rename from python/sglang/multimodal_gen/configs/quantization.py rename to python/sglang/multimodal_gen/configs/quantization/nunchaku.py index f0bd7f9c8..71fef1b0a 100644 --- a/python/sglang/multimodal_gen/configs/quantization.py +++ b/python/sglang/multimodal_gen/configs/quantization/nunchaku.py @@ -4,12 +4,13 @@ from __future__ import annotations import os import re -from dataclasses import dataclass +from dataclasses import dataclass, replace from typing import Any import torch from sglang.multimodal_gen.runtime.layers.quantization.configs.nunchaku_config import ( + NunchakuConfig, is_nunchaku_available, ) from sglang.multimodal_gen.runtime.platforms import current_platform @@ -19,6 +20,14 @@ from sglang.multimodal_gen.utils import StoreBoolean logger = init_logger(__name__) +@dataclass +class NunchakuArgsResolution: + """Normalized runtime settings derived from Nunchaku CLI-facing args.""" + + transformer_weights_path: str | None = None + nunchaku_config: NunchakuConfig | None = None + + @dataclass class NunchakuSVDQuantArgs: """CLI-facing configuration for Nunchaku (SVDQuant) inference. @@ -33,20 +42,22 @@ class NunchakuSVDQuantArgs: quantization_rank: int | None = None quantization_act_unsigned: bool = False - def _adjust_config(self) -> None: - """infer precision and rank from filename if not provided""" - if self.transformer_weights_path and not self.enable_svdquant: - filename = os.path.basename(self.transformer_weights_path) - if re.search(r"svdq-(int4|fp4)_r(\d+)", filename): - self.enable_svdquant = True - - if not self.enable_svdquant or not self.transformer_weights_path: - return - + def _infer_from_weights_path(self) -> tuple[bool, str | None, int | None]: + """Infer whether SVDQuant is enabled and parse precision/rank from filename.""" inferred_precision = None inferred_rank = None + enable_svdquant = self.enable_svdquant + + if not self.transformer_weights_path: + return enable_svdquant, inferred_precision, inferred_rank filename = os.path.basename(self.transformer_weights_path) + if not enable_svdquant and re.search(r"svdq-(int4|fp4)_r(\d+)", filename): + enable_svdquant = True + + if not enable_svdquant: + return enable_svdquant, inferred_precision, inferred_rank + # Expected pattern: svdq-{precision}_r{rank}-... # e.g., svdq-int4_r32-qwen-image.safetensors match = re.search(r"svdq-(int4|fp4)_r(\d+)", filename) @@ -56,26 +67,39 @@ class NunchakuSVDQuantArgs: inferred_precision = "nvfp4" if p_str == "fp4" else "int4" inferred_rank = int(r_str) - if self.quantization_precision is None: - self.quantization_precision = inferred_precision or "int4" + return enable_svdquant, inferred_precision, inferred_rank + + def _normalized(self) -> "NunchakuSVDQuantArgs": + enable_svdquant, inferred_precision, inferred_rank = ( + self._infer_from_weights_path() + ) + normalized = replace( + self, + enable_svdquant=enable_svdquant, + quantization_precision=( + self.quantization_precision or inferred_precision or "int4" + ), + quantization_rank=self.quantization_rank or inferred_rank or 32, + ) + + if self.quantization_precision is None and inferred_precision: if inferred_precision: logger.info( - f"inferred --quantization-precision: {self.quantization_precision} " + f"inferred --quantization-precision: {normalized.quantization_precision} " f"from --transformer-weights-path: {self.transformer_weights_path}" ) - if self.quantization_rank is None: - self.quantization_rank = inferred_rank or 32 + if self.quantization_rank is None and inferred_rank: if inferred_rank: logger.info( - f"inferred --quantization-rank: {self.quantization_rank} " + f"inferred --quantization-rank: {normalized.quantization_rank} " f"from --transformer-weights-path: {self.transformer_weights_path}" ) - def validate(self) -> None: - # TODO: warn if the served model doesn't support nunchaku - self._adjust_config() + return normalized + def _validate(self) -> None: + # TODO: warn if the served model doesn't support nunchaku if not self.enable_svdquant: return @@ -98,7 +122,6 @@ class NunchakuSVDQuantArgs: if unsupported: raise ValueError( "Nunchaku SVDQuant is currently only supported on Ampere (SM8x) or SM12x GPUs; " - "Hopper (SM90) is not supported. " f"Unsupported devices: {', '.join(unsupported)}. " "Disable it with --enable-svdquant false." ) @@ -124,6 +147,26 @@ class NunchakuSVDQuantArgs: f"Invalid --quantization-rank: {self.quantization_rank}. Must be > 0" ) + def resolve_runtime_config(self) -> NunchakuArgsResolution: + normalized = self._normalized() + normalized._validate() + + if not normalized.enable_svdquant or not normalized.transformer_weights_path: + return NunchakuArgsResolution( + transformer_weights_path=normalized.transformer_weights_path, + nunchaku_config=None, + ) + + return NunchakuArgsResolution( + transformer_weights_path=normalized.transformer_weights_path, + nunchaku_config=NunchakuConfig( + precision=normalized.quantization_precision, + rank=normalized.quantization_rank, + act_unsigned=normalized.quantization_act_unsigned, + transformer_weights_path=normalized.transformer_weights_path, + ), + ) + @staticmethod def add_cli_args(parser) -> None: parser.add_argument( 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 16a2951b3..7c0b6b788 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 @@ -1,38 +1,24 @@ -import json import logging -import os -from typing import Any, Dict, List, Optional +from typing import Any import torch from sglang.multimodal_gen.runtime.distributed import get_local_torch_device -from sglang.multimodal_gen.runtime.layers.quantization.configs.nunchaku_config import ( - NunchakuConfig, - _patch_nunchaku_scales, -) from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import ( ComponentLoader, ) from sglang.multimodal_gen.runtime.loader.fsdp_load import maybe_load_fsdp_model -from sglang.multimodal_gen.runtime.loader.utils import ( - _list_safetensors_files, - _normalize_component_type, +from sglang.multimodal_gen.runtime.loader.transformer_load_utils import ( + resolve_transformer_quant_load_spec, + resolve_transformer_safetensors_to_load, ) +from sglang.multimodal_gen.runtime.loader.utils import _normalize_component_type from sglang.multimodal_gen.runtime.models.registry import ModelRegistry from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import ( get_diffusers_component_config, - maybe_download_model, ) from sglang.multimodal_gen.runtime.utils.logging_utils import get_log_level, init_logger -from sglang.multimodal_gen.runtime.utils.quantization_utils import ( - build_nvfp4_config_from_safetensors_list, - get_metadata_from_safetensors_file, - get_quant_config, - get_quant_config_from_safetensors_metadata, -) -from sglang.multimodal_gen.utils import PRECISION_TO_TYPE -from sglang.srt.layers.quantization import QuantizationConfig from sglang.srt.utils import is_npu _is_npu = is_npu() @@ -46,97 +32,6 @@ class TransformerLoader(ComponentLoader): component_names = ["transformer", "audio_dit", "video_dit"] expected_library = "diffusers" - def get_list_of_safetensors_to_load( - self, server_args: ServerArgs, component_model_path: str - ) -> list[str]: - """ - get list of safetensors to load. - - If --transformer-weights-path is provided, load weights from that path - instead of the base model's component directory. - """ - quantized_path = server_args.transformer_weights_path - - if quantized_path: - quantized_path = maybe_download_model(quantized_path) - logger.info("using quantized transformer weights from: %s", quantized_path) - if os.path.isfile(quantized_path) and quantized_path.endswith( - ".safetensors" - ): - safetensors_list = [quantized_path] - else: - safetensors_list = _list_safetensors_files(quantized_path) - else: - safetensors_list = _list_safetensors_files(component_model_path) - - if not safetensors_list: - raise ValueError( - f"no safetensors files found in " - f"{quantized_path or component_model_path}" - ) - - return safetensors_list - - def _resolve_quant_config( - self, - hf_config: Dict[str, List[str]], - server_args: ServerArgs, - safetensors_list: list[str], - component_model_path: str, - ) -> Optional[QuantizationConfig]: - # priority: model config.json → safetensors metadata → quantization config (nvfp4, nunchaku, ...) - quant_config = get_quant_config(hf_config, component_model_path) - if quant_config is None and server_args.transformer_weights_path: - # try to read quantization_config from the safetensors metadata header - for safetensors_file in safetensors_list: - quant_config = get_quant_config_from_safetensors_metadata( - safetensors_file - ) - if quant_config: - return quant_config - - # fallback: handle nvfp4 per-layer format metadata - # ({"format_version": ..., "layers": {"name": {"format": "nvfp4"}, ...}}) - param_names_mapping_dict = ( - server_args.pipeline_config.dit_config.arch_config.param_names_mapping - ) - quant_config = build_nvfp4_config_from_safetensors_list( - safetensors_list, param_names_mapping_dict - ) - if quant_config: - return quant_config - return quant_config - - def _resolve_target_param_dtype( - self, - quant_config: Optional[dict], - nunchaku_config: Optional[NunchakuConfig], - model_cls, - server_args: ServerArgs, - ) -> Optional[torch.dtype]: - if quant_config is not None or nunchaku_config is not None: - # TODO: improve the condition - # respect dtype from checkpoint - param_dtype = None - else: - param_dtype = PRECISION_TO_TYPE[server_args.pipeline_config.dit_precision] - - if nunchaku_config is not None: - nunchaku_config.model_cls = model_cls - # verify that the nunchaku checkpoint matches the selected model class - original_dit_cls_name = json.loads( - get_metadata_from_safetensors_file( - nunchaku_config.transformer_weights_path - ).get("config") - )["_class_name"] - specified_dit_cls_name = str(model_cls.__name__) - if original_dit_cls_name != specified_dit_cls_name: - raise Exception( - f"Class name of DiT specified in nunchaku transformer_weights_path: {original_dit_cls_name} does not match that of specified DiT name: {specified_dit_cls_name}" - ) - - return param_dtype - def load_customized( self, component_model_path: str, server_args: ServerArgs, component_name: str ): @@ -144,16 +39,11 @@ class TransformerLoader(ComponentLoader): # 1. hf config config = get_diffusers_component_config(component_path=component_model_path) - # 2. quant config - safetensors_list = self.get_list_of_safetensors_to_load( + safetensors_list = resolve_transformer_safetensors_to_load( server_args, component_model_path ) - quant_config = self._resolve_quant_config( - config, server_args, safetensors_list, component_model_path - ) - - # 3. dit config + # 2. dit config # Config from Diffusers supersedes sgl_diffusion's model config component_name = _normalize_component_type(component_name) server_args.model_paths[component_name] = component_model_path @@ -169,9 +59,13 @@ class TransformerLoader(ComponentLoader): cls_name = config.pop("_class_name") model_cls, _ = ModelRegistry.resolve_model_cls(cls_name) - nunchaku_config = server_args.nunchaku_config - param_dtype = self._resolve_target_param_dtype( - quant_config, nunchaku_config, model_cls, server_args + quant_spec = resolve_transformer_quant_load_spec( + hf_config=config, + server_args=server_args, + safetensors_list=safetensors_list, + component_model_path=component_model_path, + model_cls=model_cls, + cls_name=cls_name, ) logger.info( @@ -179,14 +73,13 @@ class TransformerLoader(ComponentLoader): cls_name, len(safetensors_list), f": {safetensors_list}" if get_log_level() == logging.DEBUG else "", - param_dtype, + quant_spec.param_dtype, ) - # prepare init_param init_params: dict[str, Any] = { "config": dit_config, "hf_config": config, - "quant_config": (quant_config if quant_config else nunchaku_config), + "quant_config": quant_spec.runtime_quant_config, } if ( init_params["quant_config"] is None @@ -209,23 +102,28 @@ class TransformerLoader(ComponentLoader): cpu_offload=server_args.dit_cpu_offload, pin_cpu_memory=server_args.pin_cpu_memory, fsdp_inference=server_args.use_fsdp_inference, - # TODO(will): make these configurable - param_dtype=param_dtype, + param_dtype=quant_spec.param_dtype, reduce_dtype=torch.float32, output_dtype=None, strict=False, ) - if nunchaku_config is not None: - _patch_nunchaku_scales(model, safetensors_list) + # post-hooks (e.g., patch scales (nunchaku)) + for post_load_hook in quant_spec.post_load_hooks: + post_load_hook(model) total_params = sum(p.numel() for p in model.parameters()) logger.info("Loaded model with %.2fB parameters", total_params / 1e9) # considering the existent of mixed-precision models (e.g., nunchaku) - if next(model.parameters()).dtype != param_dtype and param_dtype: + if ( + next(model.parameters()).dtype != quant_spec.param_dtype + and quant_spec.param_dtype + ): logger.warning( - f"Model dtype does not match expected param dtype, {next(model.parameters()).dtype} vs {param_dtype}" + "Model dtype does not match expected param dtype, %s vs %s", + next(model.parameters()).dtype, + quant_spec.param_dtype, ) return model diff --git a/python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py b/python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py new file mode 100644 index 000000000..23c60043e --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py @@ -0,0 +1,306 @@ +"""Helpers and adapters for transformer quantized checkpoint loading. + +This module keeps format-specific loading quirks out of `TransformerLoader`. +The loader should stay focused on the generic load flow, while special cases +such as Nunchaku validation, NVFP4 fallback adjustments, and post-load patching +are handled here behind a small helper/adapter layer. +""" + +import json +import os +from dataclasses import dataclass, field +from functools import partial +from typing import Callable, Optional + +import torch +from torch import nn + +from sglang.multimodal_gen.runtime.layers.quantization.configs.nunchaku_config import ( + NunchakuConfig, + _patch_nunchaku_scales, +) +from sglang.multimodal_gen.runtime.loader.utils import _list_safetensors_files +from sglang.multimodal_gen.runtime.platforms import current_platform +from sglang.multimodal_gen.runtime.server_args import ServerArgs +from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import maybe_download_model +from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger +from sglang.multimodal_gen.runtime.utils.quantization_utils import ( + build_nvfp4_config_from_safetensors_list, + get_metadata_from_safetensors_file, + get_quant_config, + get_quant_config_from_safetensors_metadata, +) +from sglang.multimodal_gen.utils import PRECISION_TO_TYPE +from sglang.srt.layers.quantization import QuantizationConfig + +logger = init_logger(__name__) + +PostLoadHook = Callable[[nn.Module], None] + + +@dataclass +class TransformerQuantLoadSpec: + """Resolved loading plan for a transformer checkpoint.""" + + safetensors_list: list[str] + quant_config: Optional[QuantizationConfig] + nunchaku_config: Optional[NunchakuConfig] + param_dtype: Optional[torch.dtype] + post_load_hooks: list[PostLoadHook] = field(default_factory=list) + + @property + def runtime_quant_config(self) -> Optional[object]: + if self.quant_config is not None: + return self.quant_config + return self.nunchaku_config + + +class _TransformerQuantAdapter: + def prepare(self) -> None: + """initialize""" + pass + + def get_post_load_hooks(self) -> list[PostLoadHook]: + """post - fsdp load - hook""" + return [] + + +class _NunchakuQuantAdapter(_TransformerQuantAdapter): + """Adapter for Nunchaku checkpoints""" + + def __init__( + self, + *, + nunchaku_config: NunchakuConfig, + model_cls: type[nn.Module], + safetensors_list: list[str], + ) -> None: + self.nunchaku_config = nunchaku_config + self.model_cls = model_cls + self.safetensors_list = safetensors_list + + @staticmethod + def _validate_nunchaku_checkpoint_matches_model( + nunchaku_config: NunchakuConfig, model_cls: type[nn.Module] + ) -> None: + metadata = get_metadata_from_safetensors_file( + nunchaku_config.transformer_weights_path + ) + original_dit_cls_name = json.loads(metadata.get("config"))["_class_name"] + specified_dit_cls_name = str(model_cls.__name__) + if original_dit_cls_name != specified_dit_cls_name: + raise Exception( + f"Class name of DiT specified in nunchaku transformer_weights_path: " + f"{original_dit_cls_name} does not match that of specified DiT name: " + f"{specified_dit_cls_name}" + ) + + def prepare(self) -> None: + self.nunchaku_config.model_cls = self.model_cls + _NunchakuQuantAdapter._validate_nunchaku_checkpoint_matches_model( + nunchaku_config=self.nunchaku_config, + model_cls=self.model_cls, + ) + + def get_post_load_hooks(self) -> list[PostLoadHook]: + return [partial(_patch_nunchaku_scales, safetensors_list=self.safetensors_list)] + + +class _Flux2Nvfp4FallbackAdapter(_TransformerQuantAdapter): + """Adapter for black-forest-labs/FLUX.2-dev-NVFP4""" + + def __init__( + self, + *, + cls_name: str, + server_args: ServerArgs, + quant_config: Optional[QuantizationConfig], + ) -> None: + self.cls_name = cls_name + self.server_args = server_args + self.quant_config = quant_config + + @staticmethod + def _maybe_adjust_flux2_nvfp4_fallback_defaults( + cls_name: str, + server_args: ServerArgs, + quant_config: Optional[QuantizationConfig], + ) -> None: + if cls_name != "Flux2Transformer2DModel" or quant_config is None: + return + + quant_name_getter = getattr(type(quant_config), "get_name", None) + quant_name = quant_name_getter() if callable(quant_name_getter) else None + if quant_name != "modelopt_fp4": + return + + use_best_perf_kit = getattr( + current_platform, + "should_use_modelopt_fp4_best_performance_kit", + None, + ) + if callable(use_best_perf_kit) and use_best_perf_kit(): + return + + weights_path = os.path.basename(server_args.transformer_weights_path or "") + if not weights_path.endswith("-mixed.safetensors") or server_args.tp_size <= 1: + return + + if server_args.dit_cpu_offload or server_args.text_encoder_cpu_offload: + server_args.dit_cpu_offload = False + server_args.text_encoder_cpu_offload = False + logger.warning( + "FLUX.2 mixed NVFP4 is using the generic ModelOpt FP4 fallback with " + "tp_size=%d; disabling dit/text-encoder CPU offload to avoid TP " + "all-gather launch failures. Override the offload flags explicitly if " + "you need the old behavior.", + server_args.tp_size, + ) + + def prepare(self) -> None: + _Flux2Nvfp4FallbackAdapter._maybe_adjust_flux2_nvfp4_fallback_defaults( + cls_name=self.cls_name, + server_args=self.server_args, + quant_config=self.quant_config, + ) + + +def resolve_transformer_safetensors_to_load( + server_args: ServerArgs, component_model_path: str +) -> list[str]: + """Resolve transformer weights from the base component path or an override.""" + quantized_path = server_args.transformer_weights_path + + if quantized_path: + quantized_path = maybe_download_model(quantized_path) + logger.info("using quantized transformer weights from: %s", quantized_path) + if os.path.isfile(quantized_path) and quantized_path.endswith(".safetensors"): + safetensors_list = [quantized_path] + else: + safetensors_list = _list_safetensors_files(quantized_path) + else: + safetensors_list = _list_safetensors_files(component_model_path) + + if not safetensors_list: + raise ValueError( + f"no safetensors files found in {quantized_path or component_model_path}" + ) + + return safetensors_list + + +def resolve_transformer_quant_load_spec( + *, + hf_config: dict, + server_args: ServerArgs, + safetensors_list: list[str], + component_model_path: str, + model_cls: type[nn.Module], + cls_name: str, +) -> TransformerQuantLoadSpec: + quant_config = _resolve_quant_config( + hf_config=hf_config, + server_args=server_args, + safetensors_list=safetensors_list, + component_model_path=component_model_path, + ) + nunchaku_config = server_args.nunchaku_config + + # resolve target param dtype + param_dtype = _resolve_target_param_dtype( + quant_config=quant_config, + nunchaku_config=nunchaku_config, + server_args=server_args, + ) + + adapters = _build_transformer_quant_adapters( + cls_name=cls_name, + server_args=server_args, + quant_config=quant_config, + nunchaku_config=nunchaku_config, + model_cls=model_cls, + safetensors_list=safetensors_list, + ) + for adapter in adapters: + adapter.prepare() + + # collect post-load hooks from built adapters + post_load_hooks: list[PostLoadHook] = [] + for adapter in adapters: + post_load_hooks.extend(adapter.get_post_load_hooks()) + + return TransformerQuantLoadSpec( + safetensors_list=safetensors_list, + quant_config=quant_config, + nunchaku_config=nunchaku_config, + param_dtype=param_dtype, + post_load_hooks=post_load_hooks, + ) + + +def _build_transformer_quant_adapters( + *, + cls_name: str, + server_args: ServerArgs, + quant_config: Optional[QuantizationConfig], + nunchaku_config: Optional[NunchakuConfig], + model_cls: type[nn.Module], + safetensors_list: list[str], +) -> list[_TransformerQuantAdapter]: + adapters: list[_TransformerQuantAdapter] = [ + _Flux2Nvfp4FallbackAdapter( + cls_name=cls_name, + server_args=server_args, + quant_config=quant_config, + ) + ] + if nunchaku_config is not None: + adapters.append( + _NunchakuQuantAdapter( + nunchaku_config=nunchaku_config, + model_cls=model_cls, + safetensors_list=safetensors_list, + ) + ) + return adapters + + +def _resolve_quant_config( + *, + hf_config: dict, + server_args: ServerArgs, + safetensors_list: list[str], + component_model_path: str, +) -> Optional[QuantizationConfig]: + """ + resolve quant config from checkpoints' metadata + priority: model config.json -> safetensors metadata -> format-specific fallback + """ + quant_config = get_quant_config(hf_config, component_model_path) + if quant_config is None and server_args.transformer_weights_path: + for safetensors_file in safetensors_list: + quant_config = get_quant_config_from_safetensors_metadata(safetensors_file) + if quant_config is not None: + return quant_config + + param_names_mapping_dict = ( + server_args.pipeline_config.dit_config.arch_config.param_names_mapping + ) + quant_config = build_nvfp4_config_from_safetensors_list( + safetensors_list, param_names_mapping_dict + ) + if quant_config is not None: + return quant_config + + return quant_config + + +def _resolve_target_param_dtype( + *, + quant_config: Optional[QuantizationConfig], + nunchaku_config: Optional[NunchakuConfig], + server_args: ServerArgs, +) -> Optional[torch.dtype]: + if quant_config is not None or nunchaku_config is not None: + return None + return PRECISION_TO_TYPE[server_args.pipeline_config.dit_precision] diff --git a/python/sglang/multimodal_gen/runtime/pipelines/flux_2_nvfp4.py b/python/sglang/multimodal_gen/runtime/pipelines/flux_2_nvfp4.py index bee30d255..9a7df8426 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines/flux_2_nvfp4.py +++ b/python/sglang/multimodal_gen/runtime/pipelines/flux_2_nvfp4.py @@ -1,7 +1,7 @@ # SPDX-License-Identifier: Apache-2.0 - import glob import os +from dataclasses import dataclass from functools import lru_cache from typing import Any, cast @@ -15,40 +15,85 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger logger = init_logger(__name__) + +@dataclass(frozen=True) +class Flux2Nvfp4ModelResolution: + base_model_name: str + base_model_path: str + transformer_weights_path: str + + _FLUX2_BASE_MODEL = "black-forest-labs/FLUX.2-dev" -def _find_mixed_safetensors(local_dir: str) -> str | None: - """Return the path to the *-mixed.safetensors file in a directory, or None.""" - mixed_files = sorted(glob.glob(os.path.join(local_dir, "*-mixed.safetensors"))) - return mixed_files[0] if mixed_files else None - - @lru_cache(maxsize=1) def _resolve_flux2_base_model_path() -> str: - # The NVFP4 repo only provides the quantized transformer weights. - # We still load model_index.json and the non-transformer components from the base repo. return maybe_download_model(_FLUX2_BASE_MODEL, force_diffusers_model=True) +def _find_mixed_safetensors(local_dir: str) -> str | None: + mixed_files = sorted(glob.glob(os.path.join(local_dir, "*-mixed.safetensors"))) + return mixed_files[0] if mixed_files else None + + +def _resolve_nvfp4_transformer_weights_path( + server_args: ServerArgs, model_path: str +) -> str: + if server_args.transformer_weights_path is not None: + return server_args.transformer_weights_path + + local_nvfp4_path = maybe_download_model(model_path) + mixed_file = _find_mixed_safetensors(local_nvfp4_path) + if mixed_file is not None: + logger.info("Using mixed-precision NVFP4 weights: %s", mixed_file) + return mixed_file + + logger.warning( + "No *-mixed.safetensors found in %s; falling back to full directory", + local_nvfp4_path, + ) + return local_nvfp4_path + + +def resolve_flux2_nvfp4_model( + server_args: ServerArgs, model_path: str +) -> Flux2Nvfp4ModelResolution: + transformer_weights_path = _resolve_nvfp4_transformer_weights_path( + server_args, model_path + ) + return Flux2Nvfp4ModelResolution( + base_model_name=_FLUX2_BASE_MODEL, + base_model_path=_resolve_flux2_base_model_path(), + transformer_weights_path=transformer_weights_path, + ) + + class Flux2NvfpPipeline(Flux2Pipeline): pipeline_name = "Flux2NvfpPipeline" - _base_model_path: str | None = None + _model_resolution: Flux2Nvfp4ModelResolution | None = None - def _get_base_model_path(self) -> str: - if self._base_model_path is None: - self._base_model_path = _resolve_flux2_base_model_path() - return self._base_model_path + def _get_model_resolution( + self, server_args: ServerArgs | None = None + ) -> Flux2Nvfp4ModelResolution: + if self._model_resolution is None: + if server_args is None: + raise ValueError( + "server_args is required to resolve FLUX.2 NVFP4 paths" + ) + self._model_resolution = resolve_flux2_nvfp4_model( + server_args, self.model_path + ) + return self._model_resolution def _load_config(self) -> dict[str, Any]: - base_model_path = self._get_base_model_path() + model_resolution = self._get_model_resolution(self.server_args) logger.info("Model path: %s", self.model_path) logger.info( "Using base model '%s' at %s for config and non-transformer components", - _FLUX2_BASE_MODEL, - base_model_path, + model_resolution.base_model_name, + model_resolution.base_model_path, ) - config = verify_model_config_and_directory(base_model_path) + config = verify_model_config_and_directory(model_resolution.base_model_path) return cast(dict[str, Any], config) def _resolve_component_path( @@ -63,7 +108,7 @@ class Flux2NvfpPipeline(Flux2Pipeline): # transformer weights: ...FLUX.2-dev-NVFP4/.../flux2-dev-nvfp4-mixed.safetensors # text_encoder path: ...FLUX.2-dev/.../text_encoder component_model_path = os.path.join( - self._get_base_model_path(), load_module_name + self._get_model_resolution(server_args).base_model_path, load_module_name ) logger.debug("Resolved component path: %s", component_model_path) return component_model_path @@ -73,21 +118,11 @@ class Flux2NvfpPipeline(Flux2Pipeline): server_args: ServerArgs, loaded_modules: dict | None = None, ) -> dict: - if server_args.transformer_weights_path is None: - local_nvfp4_path = maybe_download_model(self.model_path) - mixed_file = _find_mixed_safetensors(local_nvfp4_path) - if mixed_file: - logger.info("Using mixed-precision NVFP4 weights: %s", mixed_file) - server_args.transformer_weights_path = mixed_file - else: - logger.warning( - "No *-mixed.safetensors found in %s; falling back to full directory", - local_nvfp4_path, - ) - server_args.transformer_weights_path = local_nvfp4_path - + model_resolution = self._get_model_resolution(server_args) + server_args.transformer_weights_path = model_resolution.transformer_weights_path logger.info( - "NVFP4 transformer weights: %s", server_args.transformer_weights_path + "NVFP4 transformer weights: %s", + model_resolution.transformer_weights_path, ) return super().load_modules(server_args, loaded_modules) diff --git a/python/sglang/multimodal_gen/runtime/platforms/cuda.py b/python/sglang/multimodal_gen/runtime/platforms/cuda.py index f2a721ff2..b913c387d 100644 --- a/python/sglang/multimodal_gen/runtime/platforms/cuda.py +++ b/python/sglang/multimodal_gen/runtime/platforms/cuda.py @@ -189,7 +189,7 @@ class CudaPlatformBase(Platform): logger.warning( "best performance kit (comfy-kitchen) is not installed. " "Blackwell NVFP4 will fall back to the generic ModelOpt FP4 path. " - "Install it with `pip install comfy-kitchen`." + "Install it with `pip install comfy-kitchen[cublas]`." ) @classmethod diff --git a/python/sglang/multimodal_gen/runtime/server_args.py b/python/sglang/multimodal_gen/runtime/server_args.py index 0e09c9921..f5b60be57 100644 --- a/python/sglang/multimodal_gen/runtime/server_args.py +++ b/python/sglang/multimodal_gen/runtime/server_args.py @@ -22,7 +22,7 @@ import yaml from sglang.multimodal_gen import envs from sglang.multimodal_gen.configs.models.encoders import T5Config from sglang.multimodal_gen.configs.pipeline_configs.base import PipelineConfig -from sglang.multimodal_gen.configs.quantization import NunchakuSVDQuantArgs +from sglang.multimodal_gen.configs.quantization.nunchaku import NunchakuSVDQuantArgs from sglang.multimodal_gen.runtime.layers.quantization.configs.nunchaku_config import ( NunchakuConfig, ) @@ -275,27 +275,20 @@ class ServerArgs: self.input_save_path = None def _adjust_quant_config(self): - """validate and adjust""" + """ + resolve, validate and adjust quantization config + + handles only nunchaku for now + """ - # nunchaku ncfg = self.nunchaku_config if ncfg is None or isinstance(ncfg, NunchakuConfig): return - ncfg.validate() - # propagate the path to server_args - if ncfg.transformer_weights_path: - self.transformer_weights_path = ncfg.transformer_weights_path - - if not ncfg.enable_svdquant or not ncfg.transformer_weights_path: - self.nunchaku_config = None - else: - self.nunchaku_config = NunchakuConfig( - precision=self.nunchaku_config.quantization_precision, - rank=self.nunchaku_config.quantization_rank, - act_unsigned=self.nunchaku_config.quantization_act_unsigned, - transformer_weights_path=self.nunchaku_config.transformer_weights_path, - ) + resolution = ncfg.resolve_runtime_config() + if resolution.transformer_weights_path: + self.transformer_weights_path = resolution.transformer_weights_path + self.nunchaku_config = resolution.nunchaku_config def adjust_pipeline_config(self): # enable parallel folding when SP is enabled