[diffusion] quant: support nvfp4 for Flux.2 (#20137)

Co-authored-by: zcnrex <zcnrex@gmail.com>
Co-authored-by: BBuf <1182563586@qq.com>
Co-authored-by: Yikang Cai <dcai@catalyst-fleet1.cs.cmu.edu>
Co-authored-by: CHEN Xi <78632976+RubiaCx@users.noreply.github.com>
Co-authored-by: RubiaCx <1084281732@qq.com>
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
Co-authored-by: Mick <mickjagger19@icloud.com>
This commit is contained in:
ykcai-daniel
2026-03-25 08:28:25 +08:00
committed by GitHub
co-authored by zcnrex BBuf Yikang Cai CHEN Xi RubiaCx gemini-code-assist[bot] Mick
parent 37420dce0b
commit 281fe10b5e
20 changed files with 1341 additions and 102 deletions
@@ -18,7 +18,7 @@ class FluxArchConfig(DiTArchConfig):
num_attention_heads: int = 24
joint_attention_dim: int = 4096
pooled_projection_dim: int = 768
guidance_embeds: bool = False
guidance_embeds: bool = True
axes_dims_rope: Tuple[int, int, int] = (16, 56, 56)
stacked_params_mapping: list[tuple[str, str, str]] = field(default_factory=list)
@@ -35,13 +35,49 @@ class FluxArchConfig(DiTArchConfig):
# nunchaku checkpoint uses different weight names; map to sglang flux layout
param_names_mapping: dict = field(
default_factory=lambda: {
# HF diffusers format
# HF diffusers format: strip leading "transformer." prefix
r"^transformer\.(\w*)\.(.*)$": r"\1.\2",
# FLUX2-nvfp4 format: double blocks - image attention QKV (packed, fused)
r"^double_blocks\.(\d+)\.img_attn\.qkv\.(.*)$": r"transformer_blocks.\1.attn.to_qkv.\2",
r"^double_blocks\.(\d+)\.img_attn\.proj\.(.*)$": r"transformer_blocks.\1.attn.to_out.0.\2",
r"^double_blocks\.(\d+)\.img_attn\.norm\.query_norm\.(.*)$": r"transformer_blocks.\1.attn.norm_q.\2",
r"^double_blocks\.(\d+)\.img_attn\.norm\.key_norm\.(.*)$": r"transformer_blocks.\1.attn.norm_k.\2",
# FLUX2-nvfp4 format: double blocks - text/context attention QKV (packed, fused)
r"^double_blocks\.(\d+)\.txt_attn\.qkv\.(.*)$": r"transformer_blocks.\1.attn.to_added_qkv.\2",
r"^double_blocks\.(\d+)\.txt_attn\.proj\.(.*)$": r"transformer_blocks.\1.attn.to_add_out.\2",
r"^double_blocks\.(\d+)\.txt_attn\.norm\.query_norm\.(.*)$": r"transformer_blocks.\1.attn.norm_added_q.\2",
r"^double_blocks\.(\d+)\.txt_attn\.norm\.key_norm\.(.*)$": r"transformer_blocks.\1.attn.norm_added_k.\2",
# FLUX2-nvfp4 format: double blocks - image MLP
r"^double_blocks\.(\d+)\.img_mlp\.0\.(.*)$": r"transformer_blocks.\1.ff.linear_in.\2",
r"^double_blocks\.(\d+)\.img_mlp\.2\.(.*)$": r"transformer_blocks.\1.ff.linear_out.\2",
# FLUX2-nvfp4 format: double blocks - text/context MLP
r"^double_blocks\.(\d+)\.txt_mlp\.0\.(.*)$": r"transformer_blocks.\1.ff_context.linear_in.\2",
r"^double_blocks\.(\d+)\.txt_mlp\.2\.(.*)$": r"transformer_blocks.\1.ff_context.linear_out.\2",
# FLUX2-nvfp4 format: single blocks
r"^single_blocks\.(\d+)\.linear1\.(.*)$": r"single_transformer_blocks.\1.attn.to_qkv_mlp_proj.\2",
r"^single_blocks\.(\d+)\.linear2\.(.*)$": r"single_transformer_blocks.\1.attn.to_out.\2",
r"^single_blocks\.(\d+)\.norm\.query_norm\.(.*)$": r"single_transformer_blocks.\1.attn.norm_q.\2",
r"^single_blocks\.(\d+)\.norm\.key_norm\.(.*)$": r"single_transformer_blocks.\1.attn.norm_k.\2",
# FLUX2-nvfp4 format: non-block input/output projections
r"^img_in\.(.*)$": r"x_embedder.\1",
r"^txt_in\.(.*)$": r"context_embedder.\1",
r"^time_in\.in_layer\.(.*)$": r"time_guidance_embed.timestep_embedder.linear_1.\1",
r"^time_in\.out_layer\.(.*)$": r"time_guidance_embed.timestep_embedder.linear_2.\1",
r"^guidance_in\.in_layer\.(.*)$": r"time_guidance_embed.guidance_embedder.linear_1.\1",
r"^guidance_in\.out_layer\.(.*)$": r"time_guidance_embed.guidance_embedder.linear_2.\1",
r"^double_stream_modulation_img\.lin\.(.*)$": r"double_stream_modulation_img.linear.\1",
r"^double_stream_modulation_txt\.lin\.(.*)$": r"double_stream_modulation_txt.linear.\1",
r"^single_stream_modulation\.lin\.(.*)$": r"single_stream_modulation.linear.\1",
r"^final_layer\.adaLN_modulation\.1\.(.*)$": r"norm_out.linear.\1",
r"^final_layer\.linear\.(.*)$": r"proj_out.\1",
# FLUX2-nvfp4 format: RMSNorm uses "scale" parameter; rename to "weight" (model uses .weight)
r"^(.*)\.scale$": r"\1.weight",
# transformer_blocks nunchaku format (raw export - before internal conversion)
r"^transformer_blocks\.(\d+)\.mlp_fc1\.(.*)$": r"transformer_blocks.\1.ff.net.0.proj.\2",
r"^transformer_blocks\.(\d+)\.mlp_fc2\.(.*)$": r"transformer_blocks.\1.ff.net.2.\2",
r"^transformer_blocks\.(\d+)\.mlp_context_fc1\.(.*)$": r"transformer_blocks.\1.ff_context.net.0.proj.\2",
r"^transformer_blocks\.(\d+)\.mlp_context_fc2\.(.*)$": r"transformer_blocks.\1.ff_context.net.2.\2",
# nunchaku packed QKV → fused to_qkv / to_added_qkv (matches use_fused_qkv in model)
r"^transformer_blocks\.(\d+)\.qkv_proj\.(.*)$": r"transformer_blocks.\1.attn.to_qkv.\2",
r"^transformer_blocks\.(\d+)\.qkv_proj_context\.(.*)$": r"transformer_blocks.\1.attn.to_added_qkv.\2",
r"^transformer_blocks\.(\d+)\.out_proj\.(.*)$": r"transformer_blocks.\1.attn.to_out.0.\2",
@@ -50,11 +86,11 @@ class FluxArchConfig(DiTArchConfig):
r"^transformer_blocks\.(\d+)\.norm_k\.(.*)$": r"transformer_blocks.\1.attn.norm_k.\2",
r"^transformer_blocks\.(\d+)\.norm_added_q\.(.*)$": r"transformer_blocks.\1.attn.norm_added_q.\2",
r"^transformer_blocks\.(\d+)\.norm_added_k\.(.*)$": r"transformer_blocks.\1.attn.norm_added_k.\2",
# transformer_blocks nunchaku format (already converted with convert_flux_state_dict)
# nunchaku format (already converted): add_qkv_proj → fused to_added_qkv
r"^transformer_blocks\.(\d+)\.attn\.add_qkv_proj\.(.*)$": r"transformer_blocks.\1.attn.to_added_qkv.\2",
# single_transformer_blocks nunchaku format (raw export - before internal conversion)
r"^single_transformer_blocks\.(\d+)\.qkv_proj\.(.*)$": r"single_transformer_blocks.\1.attn.to_qkv.\2",
r"^single_transformer_blocks\.(\d+)\.out_proj\.(.*)$": r"single_transformer_blocks.\1.attn.to_out.0.\2",
r"^single_transformer_blocks\.(\d+)\.qkv_proj\.(.*)$": r"single_transformer_blocks.\1.attn.to_qkv_mlp_proj.\2",
r"^single_transformer_blocks\.(\d+)\.out_proj\.(.*)$": r"single_transformer_blocks.\1.attn.to_out.\2",
r"^single_transformer_blocks\.(\d+)\.norm_q\.(.*)$": r"single_transformer_blocks.\1.attn.norm_q.\2",
r"^single_transformer_blocks\.(\d+)\.norm_k\.(.*)$": r"single_transformer_blocks.\1.attn.norm_k.\2",
# nunchaku quantization parameter name conversions (apply to all blocks)
@@ -167,3 +167,35 @@ sglang generate \
SGLang-Diffusion will automatically read the `quantization_config` metadata embedded in the safetensors file header (if present). For the quant config to be auto-detected, the file's metadata must contain a JSON-encoded `quantization_config` key with at least a `quant_method` field (e.g. `"fp8"`).
Note: this feature is a WIP
#### Option 3: NVFP4 transformer checkpoint / repo
NVFP4 support is currently for `FLUX.2-dev-NVFP4` style checkpoints.
Recommended usage:
```bash
sglang generate \
--model-path black-forest-labs/FLUX.2-dev \
--transformer-weights-path black-forest-labs/FLUX.2-dev-NVFP4 \
--prompt "a curious pikachu"
```
This keeps the CLI semantics aligned with other quantization modes:
SGLang also supports passing the NVFP4 repo or local directory directly as `--model-path`.
In that case, SGLang keeps the user-provided NVFP4 path as the model identity, uses `black-forest-labs/FLUX.2-dev` as the base model for `model_index.json` and non-transformer components, and auto-resolves the quantized transformer weights from the NVFP4 repo or local directory.
Example with direct `--model-path`:
```bash
sglang generate \
--model-path /path/to/FLUX.2-dev-NVFP4 \
--prompt "a curious pikachu"
```
Notes:
- If `--transformer-weights-path` is provided explicitly, it still takes precedence.
- For automatic resolution from a local directory, SGLang looks for `*-mixed.safetensors` first, then falls back to the whole directory.
- On Blackwell, if `comfy-kitchen` is not installed, SGLang falls back to the generic ModelOpt FP4 path and prints a warning.
+3
View File
@@ -277,6 +277,9 @@ environment_variables: dict[str, Callable[[], Any]] = {
"SGLANG_USE_RUNAI_MODEL_STREAMER": _lazy_bool(
"SGLANG_USE_RUNAI_MODEL_STREAMER", "true"
),
"SGLANG_DIFFUSION_FLASHINFER_FP4_GEMM_BACKEND": _lazy_str(
"SGLANG_DIFFUSION_FLASHINFER_FP4_GEMM_BACKEND"
),
# ROCm: use AITer GroupNorm in VAE for improved performance
"SGLANG_USE_ROCM_VAE": _lazy_bool("SGLANG_USE_ROCM_VAE"),
}
+31 -1
View File
@@ -256,6 +256,14 @@ def get_model_short_name(model_id: str) -> str:
return model_id
def _normalize_hf_cache_path(path: str) -> str:
"""Normalize a local HuggingFace cache path before substring matching.
We match registered repo ids like ``org/repo`` against cache fragments like ``models--org--repo`` that appear in snapshot/blob paths.
"""
return os.path.normpath(path).lower().replace("\\", "/")
@lru_cache(maxsize=1)
def _get_config_info(
model_path: str, model_id: Optional[str] = None
@@ -297,12 +305,32 @@ def _get_config_info(
model_id = _MODEL_HF_PATH_TO_NAME[registered_model_hf_id]
return _CONFIG_REGISTRY.get(model_id)
# 2b. Match local HuggingFace cache snapshot/blob paths such as:
# .../models--org--repo/snapshots/<hash>
# This lets users pass a local HF cache snapshot directory directly even
# when its basename is only the snapshot hash.
# Example:
# /xxx/models--black-forest-labs--FLUX.2-dev-NVFP4/snapshots/142b87e70bc3006937b7093d89ff287b5f59f071
# -> models--black-forest-labs--flux.2-dev-nvfp4 (to match with cache_repo_fragment)
normalized_model_path = _normalize_hf_cache_path(model_path)
for registered_model_hf_id in all_model_hf_paths:
cache_repo_fragment = (
f"models--{registered_model_hf_id.lower().replace('/', '--')}"
)
if cache_repo_fragment in normalized_model_path:
logger.debug(
"Resolved HuggingFace cache path '%s' to registered model '%s'.",
model_path,
registered_model_hf_id,
)
model_id = _MODEL_HF_PATH_TO_NAME[registered_model_hf_id]
return _CONFIG_REGISTRY.get(model_id)
# 3. Use detectors
if os.path.exists(model_path):
config = verify_model_config_and_directory(model_path)
else:
config = maybe_download_model_index(model_path)
pipeline_name = config.get("_class_name", "").lower()
matched_model_names = []
@@ -699,6 +727,7 @@ def _register_configs():
pipeline_config_cls=Flux2PipelineConfig,
hf_model_paths=[
"black-forest-labs/FLUX.2-dev",
"black-forest-labs/FLUX.2-dev-NVFP4",
],
model_detectors=[
lambda hf_id: "flux.2" in hf_id.lower() and "klein" not in hf_id.lower()
@@ -847,6 +876,7 @@ _register_configs()
# Maps pattern -> pipeline_name for models that don't have model_index.json
_NON_DIFFUSERS_MULTIMODAL_PATTERNS: Dict[str, str] = {
"hunyuan3d": "Hunyuan3D2Pipeline",
"flux.2-dev-nvfp4": "Flux2NvfpPipeline",
}
@@ -54,6 +54,8 @@ WEIGHT_LOADER_V2_SUPPORTED = [
"GPTQLinearMethod",
"FBGEMMFp8LinearMethod",
"ModelOptFp8LinearMethod",
"ModelOptFp4LinearMethod",
"ComfyUIFp4LinearMethod",
"IPEXAWQLinearMethod",
"IPEXGPTQLinearMethod",
"HQQMarlinMethod",
@@ -6,14 +6,18 @@ from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config impor
QuantizationConfig,
)
from sglang.multimodal_gen.runtime.layers.quantization.fp8 import Fp8Config
from sglang.multimodal_gen.runtime.layers.quantization.modelopt_quant import (
ModelOptFp4Config,
)
from sglang.multimodal_gen.runtime.layers.quantization.modelslim import ModelSlimConfig
QuantizationMethods = Literal["fp8", "modelslim"]
QuantizationMethods = Literal["fp8", "modelopt_fp4", "modelslim"]
QUANTIZATION_METHODS: list[str] = list(get_args(QuantizationMethods))
# The customized quantization methods which will be added to this dict.
_CUSTOMIZED_METHOD_TO_QUANT_CONFIG = {
"modelopt_fp4": ModelOptFp4Config,
"modelslim": ModelSlimConfig,
"fp8": Fp8Config,
}
@@ -0,0 +1,562 @@
# Adapted from https://github.com/sgl-project/sglang/blob/main/python/sglang/srt/layers/quantization/modelopt_quant.py
from __future__ import annotations
import logging
from functools import lru_cache
from typing import Any, Dict, List, Optional
import torch
from sglang.multimodal_gen.runtime.layers.linear import (
LinearMethodBase,
UnquantizedLinearMethod,
)
from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config import (
QuantizationConfig,
QuantizeMethodBase,
)
from sglang.multimodal_gen.runtime.models.parameter import (
ModelWeightParameter,
PerTensorScaleParameter,
)
from sglang.multimodal_gen.runtime.models.utils import set_weight_attrs
from sglang.multimodal_gen.runtime.platforms import current_platform
from sglang.srt.layers.quantization.modelopt_quant import (
pad_nvfp4_activation_for_cutlass,
pad_nvfp4_weight,
slice_nvfp4_output,
)
from sglang.srt.layers.quantization.utils import is_layer_skipped
from sglang.srt.layers.utils.common import copy_or_rebind_param
from sglang.srt.utils.common import round_up
logger = logging.getLogger(__name__)
@lru_cache(maxsize=1)
def _get_fp4_quantize_op():
return current_platform.get_modelopt_fp4_quantize_op()
@lru_cache(maxsize=1)
def _get_fp4_gemm_op():
return current_platform.get_modelopt_fp4_gemm_op()
@lru_cache(maxsize=1)
def _get_comfy_kitchen_cuda_backend():
try:
import comfy_kitchen.backends.cuda as ck_cuda
return ck_cuda
except Exception:
return None
class ModelOptQuantConfig(QuantizationConfig):
def __init__(
self,
exclude_modules: Optional[List[str]],
packed_modules_mapping: Optional[Dict[str, List[str]]],
):
super().__init__()
self.packed_modules_mapping = packed_modules_mapping or {}
self.exclude_modules = exclude_modules or []
def _get_quant_method(
self,
layer: torch.nn.Module,
prefix: str,
*,
Linear: type[LinearMethodBase],
) -> Optional[QuantizeMethodBase]:
from sglang.multimodal_gen.runtime.layers.linear import LinearBase
if isinstance(layer, LinearBase):
if self.is_layer_excluded(prefix) or (
self.packed_modules_mapping
and is_layer_skipped(prefix, [], self.packed_modules_mapping)
):
return UnquantizedLinearMethod()
return Linear(self)
return None
@classmethod
def get_config_filenames(cls) -> List[str]:
return ["hf_quant_config.json"]
def get_scaled_act_names(self) -> List[str]:
return []
@classmethod
def override_quantization_method(cls, hf_quant_config, user_quant) -> Optional[str]:
if hf_quant_config is None:
return None
quant_algo = hf_quant_config.get("quant_algo", "").upper()
if user_quant == "modelopt":
if not ("NVFP4" in quant_algo or "FP4" in quant_algo):
logger.warning(
f"Unsupported quant_algo '{quant_algo}' for 'modelopt'; defaulting to modelopt_fp4."
)
return "modelopt_fp4"
return None
class ModelOptFp4Config(ModelOptQuantConfig):
"""Config class for NVFP4."""
def __init__(
self,
is_checkpoint_nvfp4_serialized: bool = False,
group_size: int = None,
exclude_modules: List[str] = None,
packed_modules_mapping: Optional[Dict[str, List[str]]] = None,
) -> None:
super().__init__(exclude_modules, packed_modules_mapping)
self.is_checkpoint_nvfp4_serialized = is_checkpoint_nvfp4_serialized
if is_checkpoint_nvfp4_serialized:
logger.warning(
"Detected nvfp4 checkpoint. Please note that the "
"format is experimental and subject to change."
)
self.group_size = group_size
@classmethod
def get_name(cls) -> str:
return "modelopt_fp4"
@classmethod
def get_supported_act_dtypes(cls) -> List[torch.dtype]:
return [torch.bfloat16, torch.half, torch.float8_e4m3fn]
@classmethod
def get_min_capability(cls) -> int:
return 100
@staticmethod
def common_group_size(cfg: dict) -> int:
"""Return the unique group_size across the config; raise if missing/mismatched."""
sizes = set()
def _add_group_size_from_dict(config: dict):
group_size = config.get("group_size")
if isinstance(group_size, int):
sizes.add(group_size)
# Top-level and 'quantization' block
_add_group_size_from_dict(cfg)
quantization = cfg.get("quantization")
if isinstance(quantization, dict):
_add_group_size_from_dict(quantization)
# config_groups: accept group-level or nested dicts (e.g., weights/input_activations)
for config_groups in (cfg.get("config_groups") or {}).values():
if isinstance(config_groups, dict):
_add_group_size_from_dict(config_groups)
for config_group in config_groups.values():
if isinstance(config_group, dict):
_add_group_size_from_dict(config_group)
if not sizes:
raise ValueError("No group_size found in config.")
if len(sizes) > 1:
raise ValueError(f"Inconsistent group_size values: {sorted(sizes)}")
return next(iter(sizes))
@classmethod
def from_config(cls, config: Dict[str, Any]) -> ModelOptFp4Config:
group_size = None
exclude_modules = []
# Flat format (config.json quantization_config)
quant_method = config.get("quant_algo")
if quant_method is not None:
group_size = config.get("group_size")
if group_size is None:
config_groups = config.get("config_groups", {})
if config_groups:
first_group = next(iter(config_groups.values()), {})
group_size = first_group.get("weights", {}).get("group_size")
exclude_modules = config.get("ignore", [])
else:
# Nested format (hf_quant_config.json)
try:
quant_config = cls.get_from_keys(config, ["quantization"])
quant_method = quant_config["quant_algo"]
group_size = ModelOptFp4Config.common_group_size(config)
exclude_modules = quant_config.get("exclude_modules", [])
except (ValueError, KeyError):
raise ValueError("Cannot find 'quant_algo' in quantization config.")
if quant_method not in ["NVFP4"]:
raise ValueError(
f"Only NVFP4 quantization is supported for diffusion, got '{quant_method}'."
)
if group_size is None or exclude_modules is None:
raise ValueError(
"NVFP4 quantization requires group_size and exclude_modules "
"in the quantization config"
)
return cls(
is_checkpoint_nvfp4_serialized=True,
group_size=group_size,
exclude_modules=exclude_modules,
packed_modules_mapping=config.get("packed_modules_mapping"),
)
def is_layer_excluded(self, prefix: str):
import regex as re
fused_patterns = ["q_a_proj", "q_b_proj", "kv_a_proj_with_mqa", "kv_b_proj"]
prefix_split = prefix.split(".")
for pattern in self.exclude_modules:
regex_str = pattern.replace(".", r"\.").replace("*", r".*")
pattern_split = pattern.split(".")
if re.fullmatch(regex_str, prefix):
return True
elif (
pattern_split[-1] in fused_patterns
and pattern_split[-1] in prefix_split[-1]
):
assert len(prefix_split) == 5 and len(pattern_split) == 5
return True
return False
def get_quant_method(self, layer: torch.nn.Module, prefix: str):
should_use_best_perf_kit = getattr(
current_platform, "should_use_modelopt_fp4_best_performance_kit", None
)
warn_missing_best_perf_kit = getattr(
current_platform, "warn_if_modelopt_fp4_best_performance_kit_missing", None
)
if callable(should_use_best_perf_kit) and should_use_best_perf_kit():
linear_cls = ComfyUIFp4LinearMethod
else:
if callable(warn_missing_best_perf_kit):
warn_missing_best_perf_kit()
linear_cls = ModelOptFp4LinearMethod
return self._get_quant_method(layer, prefix, Linear=linear_cls)
class ModelOptFp4LinearMethod(LinearMethodBase):
"""NVFP4 linear method using CUTLASS FP4 GEMM."""
def __init__(self, quant_config: ModelOptFp4Config):
self.quant_config = quant_config
def create_weights(
self,
layer: torch.nn.Module,
input_size_per_partition: int,
output_partition_sizes: List[int],
input_size: int,
output_size: int,
params_dtype: torch.dtype,
**extra_weight_attrs,
):
del input_size, output_size
if not self.quant_config.is_checkpoint_nvfp4_serialized:
raise ValueError(
"NVFP4 quantization was selected, "
" dynamic quantization is not supported."
)
if input_size_per_partition % 16 != 0:
raise ValueError(
f"Unsupported model when input features size is {input_size_per_partition}, not multiple of 16, for NVFP4 quantization."
)
output_size_per_partition = sum(output_partition_sizes)
weight_loader = extra_weight_attrs.get("weight_loader")
layer.logical_widths = output_partition_sizes
layer.input_size_per_partition = input_size_per_partition
layer.output_size_per_partition = output_size_per_partition
weight_dtype = (
torch.float8_e4m3fn
if self.quant_config.is_checkpoint_nvfp4_serialized
else params_dtype
)
weight = ModelWeightParameter(
data=torch.empty(
output_size_per_partition,
input_size_per_partition // 2,
dtype=torch.uint8,
),
input_dim=1,
output_dim=0,
weight_loader=weight_loader,
)
layer.register_parameter("weight", weight)
input_scale = PerTensorScaleParameter(
data=torch.empty(len(output_partition_sizes), dtype=torch.float32),
weight_loader=weight_loader,
)
set_weight_attrs(input_scale, {"missing_param_init": "ones"})
layer.register_parameter("input_scale", input_scale)
weight_scale_2 = PerTensorScaleParameter(
data=torch.empty(len(output_partition_sizes), dtype=torch.float32),
weight_loader=weight_loader,
)
layer.register_parameter("weight_scale_2", weight_scale_2)
weight_scale = ModelWeightParameter(
data=torch.empty(
output_size_per_partition,
input_size_per_partition // self.quant_config.group_size,
dtype=weight_dtype,
),
input_dim=1,
output_dim=0,
weight_loader=weight_loader,
)
layer.register_parameter("weight_scale", weight_scale)
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
input_scale_2 = layer.input_scale.max().to(torch.float32)
weight_scale_2 = layer.weight_scale_2.max().to(torch.float32)
copy_or_rebind_param(
layer, "alpha", (input_scale_2 * weight_scale_2).to(torch.float32)
)
copy_or_rebind_param(
layer, "input_scale_inv", (1 / input_scale_2).to(torch.float32)
)
layer.output_size_per_partition = layer.weight.shape[0]
# Swap nibbles: (byte >> 4) | (byte << 4).
w = layer.weight.data
w_swapped = ((w >> 4) | (w << 4)).contiguous()
weight, weights_padding_cols = pad_nvfp4_weight(w_swapped)
layer.weights_padding_cols = weights_padding_cols
copy_or_rebind_param(layer, "weight", weight)
scales = layer.weight_scale
scale_ndim = scales.ndim
if scale_ndim == 2:
scales = scales.unsqueeze(0)
assert scales.ndim == 3
B, M, K = scales.shape
M_padded = round_up(M, 128)
K_padded = round_up(K, 4)
padded_scales = torch.zeros((B, M_padded, K_padded), dtype=scales.dtype)
padded_scales[:B, :M, :K] = scales
padded_scales = padded_scales.contiguous().cuda()
padded_scales = (
padded_scales.reshape(M_padded, K_padded)
if scale_ndim == 2
else padded_scales.reshape(B, M_padded, K_padded)
)
copy_or_rebind_param(layer, "weight_scale_interleaved", padded_scales)
def apply(
self,
layer: torch.nn.Module,
x: torch.Tensor,
bias: Optional[torch.Tensor] = None,
) -> torch.Tensor:
output_dtype = x.dtype
input_shape = x.shape
x = x.view(-1, input_shape[-1])
output_size = layer.output_size_per_partition
output_shape = list(input_shape[:-1]) + [output_size]
fp4_quantize = _get_fp4_quantize_op()
if fp4_quantize is None:
raise RuntimeError(
"No FP4 quantization kernel available. Install flashinfer or sgl_kernel."
)
x_fp4, x_scale_interleaved = fp4_quantize(x, layer.input_scale_inv)
weights_padding_cols = getattr(layer, "weights_padding_cols", 0)
x_fp4 = pad_nvfp4_activation_for_cutlass(x_fp4, weights_padding_cols)
w = layer.weight
w_scale_interleaved = layer.weight_scale_interleaved
if x_scale_interleaved.dtype == torch.uint8:
x_scale_interleaved = x_scale_interleaved.view(torch.float8_e4m3fn)
if w_scale_interleaved.dtype == torch.uint8:
w_scale_interleaved = w_scale_interleaved.view(torch.float8_e4m3fn)
fp4_gemm, flashinfer_backend = _get_fp4_gemm_op()
if flashinfer_backend is not None:
out = fp4_gemm(
x_fp4,
w.T,
x_scale_interleaved,
w_scale_interleaved.T,
layer.alpha,
output_dtype,
backend=flashinfer_backend,
)
elif fp4_gemm is not None:
out = fp4_gemm(
x_fp4,
w,
x_scale_interleaved,
w_scale_interleaved,
layer.alpha,
output_dtype,
)
else:
raise RuntimeError(
"No FP4 GEMM kernel available. Install flashinfer or sgl_kernel."
)
out = slice_nvfp4_output(out, output_size)
if bias is not None:
out = out + bias
return out.view(*output_shape)
class ComfyUIFp4LinearMethod(LinearMethodBase):
"""NVFP4 linear method using comfy-kitchen cuBLAS kernels (Blackwell)."""
def __init__(self, quant_config: ModelOptFp4Config):
self.quant_config = quant_config
def create_weights(
self,
layer: torch.nn.Module,
input_size_per_partition: int,
output_partition_sizes: List[int],
input_size: int,
output_size: int,
params_dtype: torch.dtype,
**extra_weight_attrs,
):
del input_size, output_size
if not self.quant_config.is_checkpoint_nvfp4_serialized:
raise ValueError(
"NVFP4 quantization was selected, "
"dynamic quantization is not supported."
)
if input_size_per_partition % 16 != 0:
raise ValueError(
f"Unsupported model when input features size is {input_size_per_partition}, "
"not multiple of 16, for NVFP4 quantization."
)
output_size_per_partition = sum(output_partition_sizes)
weight_loader = extra_weight_attrs.get("weight_loader")
layer.logical_widths = output_partition_sizes
layer.input_size_per_partition = input_size_per_partition
layer.output_size_per_partition = output_size_per_partition
weight = ModelWeightParameter(
data=torch.empty(
output_size_per_partition,
input_size_per_partition // 2,
dtype=torch.uint8,
),
input_dim=1,
output_dim=0,
weight_loader=weight_loader,
)
layer.register_parameter("weight", weight)
input_scale = PerTensorScaleParameter(
data=torch.empty(len(output_partition_sizes), dtype=torch.float32),
weight_loader=weight_loader,
)
set_weight_attrs(input_scale, {"missing_param_init": "ones"})
layer.register_parameter("input_scale", input_scale)
weight_scale_2 = PerTensorScaleParameter(
data=torch.empty(len(output_partition_sizes), dtype=torch.float32),
weight_loader=weight_loader,
)
layer.register_parameter("weight_scale_2", weight_scale_2)
weight_scale = ModelWeightParameter(
data=torch.empty(
output_size_per_partition,
input_size_per_partition // self.quant_config.group_size,
dtype=torch.float8_e4m3fn,
),
input_dim=1,
output_dim=0,
weight_loader=weight_loader,
)
layer.register_parameter("weight_scale", weight_scale)
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
from comfy_kitchen.float_utils import from_blocked, to_blocked
input_scale = layer.input_scale.max().to(torch.float32)
weight_scale_2 = layer.weight_scale_2.max().to(torch.float32)
copy_or_rebind_param(layer, "input_scale_ck", input_scale.cuda())
copy_or_rebind_param(layer, "weight_scale_2_ck", weight_scale_2.cuda())
layer.output_size_per_partition = layer.weight.shape[0]
copy_or_rebind_param(layer, "weight", layer.weight.data.contiguous().cuda())
# Checkpoint block scales are already in cuBLAS tiled layout.
# Pad to (roundup(N, 128), roundup(K//16, 4)) if needed.
scales = layer.weight_scale.data
N, Ks = scales.shape
N_padded = round_up(N, 128)
Ks_padded = round_up(Ks, 4)
if N == N_padded and Ks == Ks_padded:
weight_scale_ck = scales.cuda()
else:
scales_rm = from_blocked(scales, num_rows=N, num_cols=Ks)
padded_rm = torch.zeros((N_padded, Ks_padded), dtype=scales.dtype)
padded_rm[:N, :Ks] = scales_rm
weight_scale_ck = to_blocked(padded_rm, flatten=False).cuda()
copy_or_rebind_param(layer, "weight_scale_ck", weight_scale_ck)
def apply(
self,
layer: torch.nn.Module,
x: torch.Tensor,
bias: Optional[torch.Tensor] = None,
) -> torch.Tensor:
ck_cuda = _get_comfy_kitchen_cuda_backend()
if ck_cuda is None:
raise RuntimeError(
"comfy_kitchen is not available. "
"Install it to use ComfyUIFp4LinearMethod."
)
output_dtype = x.dtype
input_shape = x.shape
x_2d = x.view(-1, input_shape[-1]) # [M, K]
M = x_2d.shape[0]
output_size = layer.output_size_per_partition
output_shape = list(input_shape[:-1]) + [output_size]
if not x_2d.is_contiguous():
x_2d = x_2d.contiguous()
x_fp4, x_block_scale = ck_cuda.quantize_nvfp4(
x_2d, layer.input_scale_ck, pad_16x=True
)
out = ck_cuda.scaled_mm_nvfp4(
x_fp4,
layer.weight,
tensor_scale_a=layer.input_scale_ck,
tensor_scale_b=layer.weight_scale_2_ck,
block_scale_a=x_block_scale,
block_scale_b=layer.weight_scale_ck,
bias=bias,
out_dtype=output_dtype,
)
out = out[:M, :output_size].contiguous()
return out.view(*output_shape)
@@ -26,11 +26,13 @@ from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
)
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()
@@ -81,8 +83,8 @@ class TransformerLoader(ComponentLoader):
server_args: ServerArgs,
safetensors_list: list[str],
component_model_path: str,
) -> Optional[dict]:
# priority: model config.json → safetensors metadata → nunchaku config
) -> 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
@@ -91,7 +93,18 @@ class TransformerLoader(ComponentLoader):
safetensors_file
)
if quant_config:
break
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(
@@ -153,6 +153,7 @@ def maybe_load_fsdp_model(
quant_method.process_weights_after_loading(module)
if _is_npu:
torch.npu.empty_cache()
model.post_load_weights()
for n, p in chain(model.named_parameters(), model.named_buffers()):
if p.is_meta:
@@ -227,7 +228,7 @@ def shard_model(
fully_shard(model, **fsdp_kwargs)
# TODO(PY): device mesh for cfg parallel
# TODO(mick): need refactor, to move out checkpoint-specific adjustments
def load_model_from_full_model_state_dict(
model: FSDPModule | torch.nn.Module,
full_sd_iterator: Generator[tuple[str, torch.Tensor], None, None],
@@ -295,6 +296,33 @@ def load_model_from_full_model_state_dict(
else:
target_dtype = meta_sharded_param.dtype
_QUANTIZED_DTYPES = (
torch.uint8,
torch.float8_e4m3fn,
torch.float8_e5m2,
torch.int8,
)
if full_tensor.dtype != target_dtype:
if (
full_tensor.dtype in _QUANTIZED_DTYPES
or target_dtype in _QUANTIZED_DTYPES
):
logger.warning(
"Dtype mismatch for quantized parameter %s: "
"checkpoint has %s, model expects %s",
target_param_name,
full_tensor.dtype,
target_dtype,
)
else:
logger.warning(
"Dtype mismatch for %s: checkpoint has %s, model expects %s. "
"Casting checkpoint tensor to the target dtype during load.",
target_param_name,
full_tensor.dtype,
target_dtype,
)
if not hasattr(meta_sharded_param, "device_mesh"):
full_tensor = full_tensor.to(device=device, dtype=target_dtype)
actual_param = param_dict.get(target_param_name)
@@ -370,36 +398,52 @@ def load_model_from_full_model_state_dict(
if unused_keys:
logger.warning("Found unloaded parameters in meta state dict: %s", unused_keys)
# for nunchaku; norm_q/norm_k for SANA QK normalization layers
ALLOWED_NEW_PARAM_PATTERNS = [
# Legacy allowlist for parameter families synthesized after loading.
# New formats should declare missing_param_init on the parameter instead.
LEGACY_ALLOWED_NEW_PARAM_PATTERNS = [
"gate_compress",
"wcscales",
"wtscale",
"input_scale",
"bias",
"norm_q",
"norm_k",
]
for new_param_name in unused_keys:
if not any(pattern in new_param_name for pattern in ALLOWED_NEW_PARAM_PATTERNS):
meta_sharded_param = meta_sd.get(new_param_name)
meta_sharded_param_dtype = meta_sharded_param.dtype
actual_param = param_dict.get(new_param_name)
missing_param_init = (
getattr(actual_param, "missing_param_init", None)
if actual_param is not None
else None
)
if missing_param_init is None and not any(
pattern in new_param_name for pattern in LEGACY_ALLOWED_NEW_PARAM_PATTERNS
):
logger.error(
"Unsupported new parameter: %s. Allowed patterns: %s",
"Unsupported new parameter: %s. Allowed legacy patterns: %s",
new_param_name,
ALLOWED_NEW_PARAM_PATTERNS,
LEGACY_ALLOWED_NEW_PARAM_PATTERNS,
)
raise ValueError(
f"New parameter '{new_param_name}' is not supported. "
f"Currently only parameters containing {ALLOWED_NEW_PARAM_PATTERNS} are allowed."
"Checkpoint-specific synthesized parameters should either match "
f"{LEGACY_ALLOWED_NEW_PARAM_PATTERNS} or declare missing_param_init."
)
meta_sharded_param = meta_sd.get(new_param_name)
meta_sharded_param_dtype = meta_sharded_param.dtype
if any(
p in new_param_name for p in ("wcscales", "wtscale", "norm_q", "norm_k")
if missing_param_init == "ones" or any(
p in new_param_name
for p in ("wcscales", "wtscale", "input_scale", "norm_q", "norm_k")
):
init_like = torch.ones_like
else:
elif missing_param_init == "zeros" or missing_param_init is None:
init_like = torch.zeros_like
else:
raise ValueError(
f"Unsupported missing_param_init={missing_param_init!r} for {new_param_name}"
)
if not hasattr(meta_sharded_param, "device_mesh"):
sharded_tensor = init_like(
@@ -18,6 +18,13 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__)
_QUANTIZED_DTYPES = {
torch.uint8,
torch.float8_e4m3fn,
torch.float8_e5m2,
torch.int8,
}
@contextlib.contextmanager
def set_default_torch_dtype(dtype: torch.dtype):
@@ -135,6 +142,25 @@ def hf_to_custom_state_dict(
del to_merge_params[target_param_name]
else:
continue
existing_tensor = custom_param_sd.get(target_param_name)
if existing_tensor is not None and existing_tensor.dtype != full_tensor.dtype:
existing_is_quantized = existing_tensor.dtype in _QUANTIZED_DTYPES
current_is_quantized = full_tensor.dtype in _QUANTIZED_DTYPES
if existing_is_quantized and not current_is_quantized:
logger.debug(
"Keeping quantized duplicate for %s: existing=%s new=%s",
target_param_name,
existing_tensor.dtype,
full_tensor.dtype,
)
continue
if current_is_quantized and not existing_is_quantized:
logger.debug(
"Replacing non-quantized duplicate for %s: existing=%s new=%s",
target_param_name,
existing_tensor.dtype,
full_tensor.dtype,
)
custom_param_sd[target_param_name] = full_tensor
return custom_param_sd, reverse_param_names_mapping
@@ -73,6 +73,10 @@ class BaseDiT(nn.Module, ABC):
f"Subclasses of BaseDiT must define '{attr}' instance variable"
)
def post_load_weights(self) -> None:
"""Run model-specific post-load weight fixups after all parameters are materialized."""
return None
@property
def supported_attention_backends(self) -> set[AttentionBackendEnum]:
return self._supported_attention_backends
@@ -192,7 +192,7 @@ def _get_qkv_projections(
encoder_query = encoder_key = encoder_value = None
if encoder_hidden_states is not None and attn.added_kv_proj_dim is not None:
if getattr(attn, "use_fused_added_qkv", False):
if attn.use_fused_added_qkv:
added_qkv, _ = attn.to_added_qkv(encoder_hidden_states)
encoder_query, encoder_key, encoder_value = [
x.contiguous() for x in added_qkv.chunk(3, dim=-1)
@@ -21,12 +21,20 @@ from diffusers.models.embeddings import TimestepEmbedding, Timesteps
from diffusers.models.normalization import AdaLayerNormContinuous
from sglang.multimodal_gen.configs.models.dits.flux import FluxConfig
from sglang.multimodal_gen.runtime.distributed import divide, get_tp_world_size
from sglang.multimodal_gen.runtime.layers.attention import USPAttention
from sglang.multimodal_gen.runtime.layers.layernorm import RMSNorm, apply_qk_norm
from sglang.multimodal_gen.runtime.layers.linear import ColumnParallelLinear
from sglang.multimodal_gen.runtime.layers.linear import (
ColumnParallelLinear,
MergedColumnParallelLinear,
RowParallelLinear,
)
from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config import (
QuantizationConfig,
)
from sglang.multimodal_gen.runtime.layers.quantization.modelopt_quant import (
ModelOptFp4Config,
)
from sglang.multimodal_gen.runtime.layers.rotary_embedding import (
NDRotaryEmbedding,
apply_flashinfer_rope_qk_inplace,
@@ -42,15 +50,25 @@ logger = init_logger(__name__) # pylint: disable=invalid-name
def _get_qkv_projections(
attn: "Flux2Attention", hidden_states, encoder_hidden_states=None
):
query, _ = attn.to_q(hidden_states)
key, _ = attn.to_k(hidden_states)
value, _ = attn.to_v(hidden_states)
if attn.use_fused_qkv:
qkv, _ = attn.to_qkv(hidden_states)
query, key, value = [t.contiguous() for t in qkv.chunk(3, dim=-1)]
else:
query, _ = attn.to_q(hidden_states)
key, _ = attn.to_k(hidden_states)
value, _ = attn.to_v(hidden_states)
encoder_query = encoder_key = encoder_value = None
if encoder_hidden_states is not None and attn.added_kv_proj_dim is not None:
encoder_query, _ = attn.add_q_proj(encoder_hidden_states)
encoder_key, _ = attn.add_k_proj(encoder_hidden_states)
encoder_value, _ = attn.add_v_proj(encoder_hidden_states)
if attn.use_fused_added_qkv:
added_qkv, _ = attn.to_added_qkv(encoder_hidden_states)
encoder_query, encoder_key, encoder_value = [
t.contiguous() for t in added_qkv.chunk(3, dim=-1)
]
else:
encoder_query, _ = attn.add_q_proj(encoder_hidden_states)
encoder_key, _ = attn.add_k_proj(encoder_hidden_states)
encoder_value, _ = attn.add_v_proj(encoder_hidden_states)
return query, key, value, encoder_query, encoder_key, encoder_value
@@ -80,6 +98,7 @@ class Flux2FeedForward(nn.Module):
inner_dim: Optional[int] = None,
bias: bool = False,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
):
super().__init__()
if inner_dim is None:
@@ -87,12 +106,22 @@ class Flux2FeedForward(nn.Module):
dim_out = dim_out or dim
# Flux2SwiGLU will reduce the dimension by half
self.linear_in = ColumnParallelLinear(
dim, inner_dim * 2, bias=bias, gather_output=True, quant_config=quant_config
self.linear_in = MergedColumnParallelLinear(
dim,
[inner_dim, inner_dim],
bias=bias,
gather_output=False,
quant_config=quant_config,
prefix=f"{prefix}.linear_in" if prefix else "linear_in",
)
self.act_fn = Flux2SwiGLU()
self.linear_out = ColumnParallelLinear(
inner_dim, dim_out, bias=bias, gather_output=True, quant_config=quant_config
self.linear_out = RowParallelLinear(
inner_dim,
dim_out,
bias=bias,
input_is_parallel=True,
quant_config=quant_config,
prefix=f"{prefix}.linear_out" if prefix else "linear_out",
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
@@ -117,6 +146,7 @@ class Flux2Attention(torch.nn.Module, AttentionModuleMixin):
out_dim: int = None,
elementwise_affine: bool = True,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
):
super().__init__()
@@ -125,6 +155,9 @@ class Flux2Attention(torch.nn.Module, AttentionModuleMixin):
self.query_dim = query_dim
self.out_dim = out_dim if out_dim is not None else query_dim
self.heads = out_dim // dim_head if out_dim is not None else num_heads
self.tp_size = get_tp_world_size()
self.local_heads = divide(self.heads, self.tp_size)
self.local_inner_dim = divide(self.inner_dim, self.tp_size)
self.use_bias = bias
self.dropout = dropout
@@ -132,27 +165,45 @@ class Flux2Attention(torch.nn.Module, AttentionModuleMixin):
self.added_kv_proj_dim = added_kv_proj_dim
self.added_proj_bias = added_proj_bias
self.to_q = ColumnParallelLinear(
query_dim,
self.inner_dim,
bias=bias,
gather_output=True,
quant_config=quant_config,
)
self.to_k = ColumnParallelLinear(
query_dim,
self.inner_dim,
bias=bias,
gather_output=True,
quant_config=quant_config,
)
self.to_v = ColumnParallelLinear(
query_dim,
self.inner_dim,
bias=bias,
gather_output=True,
quant_config=quant_config,
)
# Fuse Q/K/V into a single linear when using NVFP4: the checkpoint stores them
# packed as one tensor, so a fused layer avoids splitting during weight loading.
self.use_fused_qkv = isinstance(quant_config, ModelOptFp4Config)
self.use_fused_added_qkv = self.use_fused_qkv
if self.use_fused_qkv:
self.to_qkv = MergedColumnParallelLinear(
query_dim,
[self.inner_dim] * 3,
bias=bias,
gather_output=False,
quant_config=quant_config,
prefix=f"{prefix}.to_qkv" if prefix else "to_qkv",
)
else:
self.to_q = ColumnParallelLinear(
query_dim,
self.inner_dim,
bias=bias,
gather_output=False,
quant_config=quant_config,
prefix=f"{prefix}.to_q" if prefix else "to_q",
)
self.to_k = ColumnParallelLinear(
query_dim,
self.inner_dim,
bias=bias,
gather_output=False,
quant_config=quant_config,
prefix=f"{prefix}.to_k" if prefix else "to_k",
)
self.to_v = ColumnParallelLinear(
query_dim,
self.inner_dim,
bias=bias,
gather_output=False,
quant_config=quant_config,
prefix=f"{prefix}.to_v" if prefix else "to_v",
)
# QK Norm
self.norm_q = RMSNorm(dim_head, eps=eps)
@@ -160,12 +211,13 @@ class Flux2Attention(torch.nn.Module, AttentionModuleMixin):
self.to_out = torch.nn.ModuleList([])
self.to_out.append(
ColumnParallelLinear(
RowParallelLinear(
self.inner_dim,
self.out_dim,
bias=out_bias,
gather_output=True,
input_is_parallel=True,
quant_config=quant_config,
prefix=f"{prefix}.to_out.0" if prefix else "to_out.0",
)
)
self.to_out.append(torch.nn.Dropout(dropout))
@@ -173,37 +225,52 @@ class Flux2Attention(torch.nn.Module, AttentionModuleMixin):
if added_kv_proj_dim is not None:
self.norm_added_q = RMSNorm(dim_head, eps=eps)
self.norm_added_k = RMSNorm(dim_head, eps=eps)
self.add_q_proj = ColumnParallelLinear(
added_kv_proj_dim,
self.inner_dim,
bias=added_proj_bias,
gather_output=True,
quant_config=quant_config,
)
self.add_k_proj = ColumnParallelLinear(
added_kv_proj_dim,
self.inner_dim,
bias=added_proj_bias,
gather_output=True,
quant_config=quant_config,
)
self.add_v_proj = ColumnParallelLinear(
added_kv_proj_dim,
self.inner_dim,
bias=added_proj_bias,
gather_output=True,
quant_config=quant_config,
)
self.to_add_out = ColumnParallelLinear(
if self.use_fused_added_qkv:
# txt_attn.qkv is always BF16 in the NVFP4 checkpoint — no quant needed
self.to_added_qkv = MergedColumnParallelLinear(
added_kv_proj_dim,
[self.inner_dim] * 3,
bias=added_proj_bias,
gather_output=False,
quant_config=None,
prefix=f"{prefix}.to_added_qkv" if prefix else "to_added_qkv",
)
else:
self.add_q_proj = ColumnParallelLinear(
added_kv_proj_dim,
self.inner_dim,
bias=added_proj_bias,
gather_output=False,
quant_config=quant_config,
prefix=f"{prefix}.add_q_proj" if prefix else "add_q_proj",
)
self.add_k_proj = ColumnParallelLinear(
added_kv_proj_dim,
self.inner_dim,
bias=added_proj_bias,
gather_output=False,
quant_config=quant_config,
prefix=f"{prefix}.add_k_proj" if prefix else "add_k_proj",
)
self.add_v_proj = ColumnParallelLinear(
added_kv_proj_dim,
self.inner_dim,
bias=added_proj_bias,
gather_output=False,
quant_config=quant_config,
prefix=f"{prefix}.add_v_proj" if prefix else "add_v_proj",
)
self.to_add_out = RowParallelLinear(
self.inner_dim,
query_dim,
bias=out_bias,
gather_output=True,
input_is_parallel=True,
quant_config=quant_config,
prefix=f"{prefix}.to_add_out" if prefix else "to_add_out",
)
self.attn = USPAttention(
num_heads=num_heads,
num_heads=self.local_heads,
head_size=self.head_dim,
dropout_rate=0,
softmax_scale=None,
@@ -220,9 +287,9 @@ class Flux2Attention(torch.nn.Module, AttentionModuleMixin):
_get_qkv_projections(self, hidden_states, encoder_hidden_states)
)
query = query.unflatten(-1, (self.heads, -1))
key = key.unflatten(-1, (self.heads, -1))
value = value.unflatten(-1, (self.heads, -1))
query = query.unflatten(-1, (self.local_heads, -1))
key = key.unflatten(-1, (self.local_heads, -1))
value = value.unflatten(-1, (self.local_heads, -1))
query, key = apply_qk_norm(
q=query,
@@ -234,9 +301,9 @@ class Flux2Attention(torch.nn.Module, AttentionModuleMixin):
)
if self.added_kv_proj_dim is not None:
encoder_query = encoder_query.unflatten(-1, (self.heads, -1))
encoder_key = encoder_key.unflatten(-1, (self.heads, -1))
encoder_value = encoder_value.unflatten(-1, (self.heads, -1))
encoder_query = encoder_query.unflatten(-1, (self.local_heads, -1))
encoder_key = encoder_key.unflatten(-1, (self.local_heads, -1))
encoder_value = encoder_value.unflatten(-1, (self.local_heads, -1))
encoder_query, encoder_key = apply_qk_norm(
q=encoder_query,
@@ -317,6 +384,7 @@ class Flux2ParallelSelfAttention(torch.nn.Module, AttentionModuleMixin):
mlp_ratio: float = 4.0,
mlp_mult_factor: int = 2,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
):
super().__init__()
@@ -325,21 +393,27 @@ class Flux2ParallelSelfAttention(torch.nn.Module, AttentionModuleMixin):
self.query_dim = query_dim
self.out_dim = out_dim if out_dim is not None else query_dim
self.heads = out_dim // dim_head if out_dim is not None else num_heads
self.tp_size = get_tp_world_size()
self.local_heads = divide(self.heads, self.tp_size)
self.local_inner_dim = divide(self.inner_dim, self.tp_size)
self.use_bias = bias
self.dropout = dropout
self.mlp_ratio = mlp_ratio
self.mlp_hidden_dim = int(query_dim * self.mlp_ratio)
self.local_mlp_hidden_dim = divide(self.mlp_hidden_dim, self.tp_size)
self.mlp_mult_factor = mlp_mult_factor
# Fused QKV projections + MLP input projection
self.to_qkv_mlp_proj = ColumnParallelLinear(
self.to_qkv_mlp_proj = MergedColumnParallelLinear(
self.query_dim,
self.inner_dim * 3 + self.mlp_hidden_dim * self.mlp_mult_factor,
[self.inner_dim, self.inner_dim, self.inner_dim]
+ [self.mlp_hidden_dim] * self.mlp_mult_factor,
bias=bias,
gather_output=True,
gather_output=False,
quant_config=quant_config,
prefix=f"{prefix}.to_qkv_mlp_proj" if prefix else "to_qkv_mlp_proj",
)
self.mlp_act_fn = Flux2SwiGLU()
@@ -348,16 +422,17 @@ class Flux2ParallelSelfAttention(torch.nn.Module, AttentionModuleMixin):
self.norm_k = RMSNorm(dim_head, eps=eps)
# Fused attention output projection + MLP output projection
self.to_out = ColumnParallelLinear(
self.to_out = RowParallelLinear(
self.inner_dim + self.mlp_hidden_dim,
self.out_dim,
bias=out_bias,
gather_output=True,
input_is_parallel=True,
quant_config=quant_config,
prefix=f"{prefix}.to_out" if prefix else "to_out",
)
self.attn = USPAttention(
num_heads=num_heads,
num_heads=self.local_heads,
head_size=self.head_dim,
dropout_rate=0,
softmax_scale=None,
@@ -376,16 +451,19 @@ class Flux2ParallelSelfAttention(torch.nn.Module, AttentionModuleMixin):
hidden_states, _ = self.to_qkv_mlp_proj(hidden_states)
qkv, mlp_hidden_states = torch.split(
hidden_states,
[3 * self.inner_dim, self.mlp_hidden_dim * self.mlp_mult_factor],
[
3 * self.local_inner_dim,
self.local_mlp_hidden_dim * self.mlp_mult_factor,
],
dim=-1,
)
# Handle the attention logic
query, key, value = qkv.chunk(3, dim=-1)
query = query.unflatten(-1, (self.heads, -1))
key = key.unflatten(-1, (self.heads, -1))
value = value.unflatten(-1, (self.heads, -1))
query = query.unflatten(-1, (self.local_heads, -1))
key = key.unflatten(-1, (self.local_heads, -1))
value = value.unflatten(-1, (self.local_heads, -1))
query = self.norm_q(query)
key = self.norm_k(key)
@@ -428,6 +506,7 @@ class Flux2SingleTransformerBlock(nn.Module):
eps: float = 1e-6,
bias: bool = False,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
):
super().__init__()
@@ -447,6 +526,7 @@ class Flux2SingleTransformerBlock(nn.Module):
mlp_ratio=mlp_ratio,
mlp_mult_factor=2,
quant_config=quant_config,
prefix=f"{prefix}.attn" if prefix else "attn",
)
def forward(
@@ -502,6 +582,7 @@ class Flux2TransformerBlock(nn.Module):
eps: float = 1e-6,
bias: bool = False,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
):
super().__init__()
self.mlp_hidden_dim = int(dim * mlp_ratio)
@@ -520,16 +601,27 @@ class Flux2TransformerBlock(nn.Module):
out_bias=bias,
eps=eps,
quant_config=quant_config,
prefix=f"{prefix}.attn" if prefix else "attn",
)
self.norm2 = nn.LayerNorm(dim, elementwise_affine=False, eps=eps)
self.ff = Flux2FeedForward(
dim=dim, dim_out=dim, mult=mlp_ratio, bias=bias, quant_config=quant_config
dim=dim,
dim_out=dim,
mult=mlp_ratio,
bias=bias,
quant_config=quant_config,
prefix=f"{prefix}.ff" if prefix else "ff",
)
self.norm2_context = nn.LayerNorm(dim, elementwise_affine=False, eps=eps)
self.ff_context = Flux2FeedForward(
dim=dim, dim_out=dim, mult=mlp_ratio, bias=bias, quant_config=quant_config
dim=dim,
dim_out=dim,
mult=mlp_ratio,
bias=bias,
quant_config=quant_config,
prefix=f"{prefix}.ff_context" if prefix else "ff_context",
)
def forward(
@@ -706,6 +798,30 @@ class Flux2Transformer2DModel(CachableDiT, OffloadableDiTMixin):
"""
param_names_mapping = FluxConfig().arch_config.param_names_mapping
scale_shift_swap_params = ("norm_out.linear.weight", "norm_out.linear.bias")
def post_load_weights(self) -> None:
if not isinstance(getattr(self, "quant_config", None), ModelOptFp4Config):
return
# BFL/ComfyUI checkpoints store AdaLN modulation params as [scale, shift],
# while diffusers expects [shift, scale].
for param_name in self.scale_shift_swap_params:
parts = param_name.split(".")
module = self
for part in parts[:-1]:
module = getattr(module, part)
param = getattr(module, parts[-1], None)
if param is None:
continue
half = param.shape[0] // 2
with torch.no_grad():
first_half = param[:half].clone()
param[:half] = param[half:]
param[half:] = first_half
logger.info(
"Swapped scale/shift order for %s (BFL → diffusers)", param_name
)
def __init__(
self,
@@ -731,6 +847,8 @@ class Flux2Transformer2DModel(CachableDiT, OffloadableDiTMixin):
self.out_channels = out_channels or in_channels
self.inner_dim = num_attention_heads * attention_head_dim
self.guidance_embeds = guidance_embeds
quant_config = quant_config if quant_config is not None else config.quant_config
self.quant_config = quant_config
# 1. Sinusoidal positional embedding for RoPE on image and text tokens
self.rotary_emb = Flux2PosEmbed(theta=rope_theta, axes_dim=axes_dims_rope)
@@ -775,8 +893,9 @@ class Flux2Transformer2DModel(CachableDiT, OffloadableDiTMixin):
eps=eps,
bias=False,
quant_config=quant_config,
prefix=f"transformer_blocks.{i}",
)
for _ in range(num_layers)
for i in range(num_layers)
]
)
@@ -791,8 +910,9 @@ class Flux2Transformer2DModel(CachableDiT, OffloadableDiTMixin):
eps=eps,
bias=False,
quant_config=quant_config,
prefix=f"single_transformer_blocks.{i}",
)
for _ in range(num_single_layers)
for i in range(num_single_layers)
]
)
@@ -809,6 +929,8 @@ class Flux2Transformer2DModel(CachableDiT, OffloadableDiTMixin):
patch_size * patch_size * self.out_channels,
bias=False,
gather_output=True,
quant_config=quant_config,
prefix="proj_out",
)
self.layer_names = ["transformer_blocks", "single_transformer_blocks"]
@@ -0,0 +1,95 @@
# SPDX-License-Identifier: Apache-2.0
import glob
import os
from functools import lru_cache
from typing import Any, cast
from sglang.multimodal_gen.runtime.pipelines.flux_2 import Flux2Pipeline
from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
maybe_download_model,
verify_model_config_and_directory,
)
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__)
_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)
class Flux2NvfpPipeline(Flux2Pipeline):
pipeline_name = "Flux2NvfpPipeline"
_base_model_path: str | 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 _load_config(self) -> dict[str, Any]:
base_model_path = self._get_base_model_path()
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,
)
config = verify_model_config_and_directory(base_model_path)
return cast(dict[str, Any], config)
def _resolve_component_path(
self, server_args: ServerArgs, module_name: str, load_module_name: str
) -> str:
override_path = server_args.component_paths.get(module_name)
if override_path is not None:
return maybe_download_model(override_path)
# get non-transformer components from the base FLUX.2 repo explicitly.
# e.g.:
# 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
)
logger.debug("Resolved component path: %s", component_model_path)
return component_model_path
def load_modules(
self,
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
logger.info(
"NVFP4 transformer weights: %s", server_args.transformer_weights_path
)
return super().load_modules(server_args, loaded_modules)
EntryClass = Flux2NvfpPipeline
@@ -103,6 +103,95 @@ class CudaPlatformBase(Platform):
return False
return True
@classmethod
@lru_cache(maxsize=1)
def get_modelopt_fp4_quantize_op(cls) -> Callable | None:
try:
from flashinfer import fp4_quantize
return fp4_quantize
except ImportError:
pass
try:
from sgl_kernel import scaled_fp4_quant as fp4_quantize
return fp4_quantize
except ImportError:
return None
@classmethod
@lru_cache(maxsize=1)
def get_modelopt_fp4_gemm_op(cls) -> tuple[Callable | None, str | None]:
if cls.is_blackwell():
try:
from flashinfer import mm_fp4 as flashinfer_mm_fp4
return flashinfer_mm_fp4, "cudnn"
except ImportError:
pass
try:
from sgl_kernel import cutlass_scaled_fp4_mm as cutlass_fp4_gemm
return cutlass_fp4_gemm, None
except ImportError:
pass
try:
from flashinfer import mm_fp4 as flashinfer_mm_fp4
return flashinfer_mm_fp4, "auto"
except ImportError:
return None, None
@classmethod
@lru_cache(maxsize=1)
def has_modelopt_fp4_best_performance_kit(cls) -> bool:
try:
import comfy_kitchen.backends.cuda # noqa: F401
return True
except Exception:
return False
@classmethod
@lru_cache(maxsize=1)
def can_use_modelopt_fp4_best_performance_kit(cls) -> bool:
if not cls.is_blackwell() or not cls.has_modelopt_fp4_best_performance_kit():
return False
try:
import comfy_kitchen.backends.cuda as ck_cuda
device = cls.get_local_torch_device()
x = torch.zeros((16, 16), dtype=torch.bfloat16, device=device)
scale = torch.ones((), dtype=torch.float32, device=device)
ck_cuda.quantize_nvfp4(x, scale, pad_16x=True)
return True
except Exception as e:
logger.warning(
"best performance kit (comfy-kitchen) is installed but unusable on "
"this system (%s). Blackwell NVFP4 will fall back to the generic "
"ModelOpt FP4 path.",
e,
)
return False
@classmethod
def should_use_modelopt_fp4_best_performance_kit(cls) -> bool:
return cls.can_use_modelopt_fp4_best_performance_kit()
@classmethod
@lru_cache(maxsize=1)
def warn_if_modelopt_fp4_best_performance_kit_missing(cls) -> None:
if cls.is_blackwell() and not cls.has_modelopt_fp4_best_performance_kit():
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`."
)
@classmethod
def is_full_nvlink(cls, device_ids: list[int]) -> bool:
raise NotImplementedError
@@ -6,6 +6,7 @@ from __future__ import annotations
import enum
import random
from collections.abc import Callable
from functools import lru_cache
from typing import TYPE_CHECKING, Any, NamedTuple
@@ -198,6 +199,30 @@ class Platform:
def is_amp_supported(cls) -> bool:
return True
@classmethod
def get_modelopt_fp4_quantize_op(cls) -> Callable | None:
return None
@classmethod
def get_modelopt_fp4_gemm_op(cls) -> tuple[Callable | None, str | None]:
return None, None
@classmethod
def has_modelopt_fp4_best_performance_kit(cls) -> bool:
return False
@classmethod
def can_use_modelopt_fp4_best_performance_kit(cls) -> bool:
return False
@classmethod
def should_use_modelopt_fp4_best_performance_kit(cls) -> bool:
return False
@classmethod
def warn_if_modelopt_fp4_best_performance_kit_missing(cls) -> None:
pass
@classmethod
def get_local_torch_device(cls) -> torch.device:
raise NotImplementedError
@@ -167,3 +167,134 @@ def get_metadata_from_safetensors_file(file_path: str):
return metadata
except Exception as e:
logger.warning(e)
def _build_nvfp4_config_from_safetensors_files(
file_paths: list[str],
param_names_mapping_dict: Optional[dict] = None,
) -> Optional[QuantizationConfig]:
"""Build a single NVFP4 config by aggregating metadata across multiple files.
Some checkpoints split BF16 fallback layers and NVFP4 layers across multiple
safetensors. Building the config from only the first matching file can
incorrectly exclude layers that are quantized in a later shard.
"""
import torch
group_size = None
quantized_bfl_modules: set[str] = set()
non_quantized_bfl_modules: set[str] = set()
files_with_nvfp4_metadata: list[str] = []
for file_path in file_paths:
metadata = get_metadata_from_safetensors_file(file_path)
if not metadata:
continue
quant_config_str = metadata.get("_quantization_metadata")
if not quant_config_str:
continue
quant_config_dict = json.loads(quant_config_str)
if (
"format_version" not in quant_config_dict
or "layers" not in quant_config_dict
):
continue
layers = quant_config_dict.get("layers", {})
file_quantized_modules = {
layer_name
for layer_name, layer_cfg in layers.items()
if isinstance(layer_cfg, dict) and layer_cfg.get("format") == "nvfp4"
}
if not file_quantized_modules:
continue
files_with_nvfp4_metadata.append(file_path)
quantized_bfl_modules.update(file_quantized_modules)
with safe_open(file_path, framework="pt", device="cpu") as f:
all_keys = set(f.keys())
if group_size is None:
for layer_name in file_quantized_modules:
weight_key = f"{layer_name}.weight"
scale_key = f"{layer_name}.weight_scale"
if weight_key in all_keys and scale_key in all_keys:
w = f.get_tensor(weight_key)
s = f.get_tensor(scale_key)
input_size = w.shape[1] * 2
group_size = input_size // s.shape[1]
break
for k in sorted(all_keys):
if not k.endswith(".weight"):
continue
t = f.get_tensor(k)
if t.dtype != torch.uint8:
non_quantized_bfl_modules.add(k[: -len(".weight")])
if not files_with_nvfp4_metadata:
return None
if group_size is None:
logger.warning(
"Could not infer group_size from NVFP4 safetensors: %s",
", ".join(files_with_nvfp4_metadata),
)
return None
exclude_bfl_modules = sorted(non_quantized_bfl_modules - quantized_bfl_modules)
exclude_modules = []
if param_names_mapping_dict:
from sglang.multimodal_gen.runtime.loader.utils import get_param_names_mapping
mapping_fn = get_param_names_mapping(param_names_mapping_dict)
for module_bfl in exclude_bfl_modules:
mapped, _, _ = mapping_fn(f"{module_bfl}.weight")
exclude_modules.append(
mapped[: -len(".weight")] if mapped.endswith(".weight") else mapped
)
else:
exclude_modules = exclude_bfl_modules
try:
quant_cls = get_quantization_config("modelopt_fp4")
result = quant_cls.from_config(
{"quant_algo": "NVFP4", "group_size": group_size, "ignore": exclude_modules}
)
logger.info(
"Built NVFP4 quant config from %d safetensors: group_size=%d, %d excluded modules",
len(files_with_nvfp4_metadata),
group_size,
len(exclude_modules),
)
return result
except Exception as e:
logger.warning(
"Failed to build NVFP4 config from %s: %s",
", ".join(files_with_nvfp4_metadata),
e,
)
return None
def build_nvfp4_config_from_safetensors(
file_path: str,
param_names_mapping_dict: Optional[dict] = None,
) -> Optional[QuantizationConfig]:
"""Backward-compatible wrapper for a single safetensors file."""
return _build_nvfp4_config_from_safetensors_files(
[file_path], param_names_mapping_dict
)
def build_nvfp4_config_from_safetensors_list(
file_paths: list[str],
param_names_mapping_dict: Optional[dict] = None,
) -> Optional[QuantizationConfig]:
return _build_nvfp4_config_from_safetensors_files(
file_paths, param_names_mapping_dict
)
@@ -798,6 +798,18 @@ if not current_platform.is_hip():
)
)
# TODO: enable on 4090/5090/b200
ONE_GPU_CASES_C = [
DiffusionTestCase(
"flux_2_nvfp4_t2i",
DiffusionServerArgs(
model_path="black-forest-labs/FLUX.2-dev-NVFP4",
modality="image",
),
T2I_sampling_params,
)
]
TWO_GPU_CASES_A = [
DiffusionTestCase(
"wan2_2_i2v_a14b_2gpu",
@@ -65,6 +65,15 @@ class TestModelIdResolution(unittest.TestCase):
info = _get_config_info(expanded, model_id="Qwen-Image")
self.assertIsNotNone(info)
def test_hf_cache_snapshot_path_resolves_registered_nvfp4_model(self):
path = (
"/root/.cache/huggingface/hub/"
"models--black-forest-labs--FLUX.2-dev-NVFP4/"
"snapshots/142b87e70bc3006937b7093d89ff287b5f59f071"
)
info = _get_config_info(path)
self.assertIsNotNone(info)
def test_model_id_unknown_falls_back_without_crash(self):
# unrecognized model_id: should warn and fall back to path-based detection
# with an unresolvable path, expect RuntimeError from the detector step