diff --git a/docs_new/cards/logos/ideogram.png b/docs_new/cards/logos/ideogram.png
new file mode 100644
index 000000000..f39ef1145
Binary files /dev/null and b/docs_new/cards/logos/ideogram.png differ
diff --git a/docs_new/cookbook/diffusion/Ideogram/Ideogram4.mdx b/docs_new/cookbook/diffusion/Ideogram/Ideogram4.mdx
new file mode 100644
index 000000000..1432c4a77
--- /dev/null
+++ b/docs_new/cookbook/diffusion/Ideogram/Ideogram4.mdx
@@ -0,0 +1,82 @@
+---
+title: Ideogram 4
+metatags:
+ description: "Deploy Ideogram 4 with SGLang Diffusion for high-aesthetic text-to-image generation."
+---
+
+## 1. Model introduction
+
+[Ideogram 4](https://huggingface.co/ideogram-ai/ideogram-4-nf4) is Ideogram's text-to-image diffusion model. SGLang Diffusion supports the official NF4 checkpoint, the official FP8 checkpoint, and the Comfy-Org NVFP4 transformer checkpoint.
+
+Compared with previous open-source image models, Ideogram 4 provides a significant aesthetic lift, with stronger composition, more polished visual style, and better typography-aware generation.
+
+| Variant | Hugging Face model ID | Notes |
+| --- | --- | --- |
+| NF4 | `ideogram-ai/ideogram-4-nf4` | Official bitsandbytes NF4 checkpoint. Use this path first for low-memory deployment. |
+| FP8 | `ideogram-ai/ideogram-4-fp8` | Official FP8 checkpoint. |
+| NVFP4 | `Comfy-Org/Ideogram-4` | Comfy-Org NVFP4 transformer weights. SGLang loads non-transformer components from `ideogram-ai/ideogram-4-fp8`. |
+
+## 2. Prerequisites
+
+- NVIDIA CUDA GPU.
+- SGLang installed with diffusion dependencies.
+- `bitsandbytes>=0.46.1` for the NF4 checkpoint.
+- `HF_TOKEN` with access to the Ideogram 4 gated repositories.
+
+## 3. Serve the model
+
+NF4:
+
+```bash Command
+HF_TOKEN=$HF_TOKEN sglang serve \
+ --model-path ideogram-ai/ideogram-4-nf4 \
+ --num-gpus 1 \
+ --performance-mode auto \
+ --port 30010
+```
+
+FP8:
+
+```bash Command
+HF_TOKEN=$HF_TOKEN sglang serve \
+ --model-path ideogram-ai/ideogram-4-fp8 \
+ --num-gpus 1 \
+ --performance-mode auto \
+ --port 30010
+```
+
+Comfy-Org NVFP4:
+
+```bash Command
+HF_TOKEN=$HF_TOKEN sglang serve \
+ --model-path Comfy-Org/Ideogram-4 \
+ --num-gpus 1 \
+ --performance-mode auto \
+ --port 30010
+```
+
+Use B200 or another Blackwell GPU for NVFP4.
+
+## 4. Generate an image
+
+```python Example
+import base64
+from openai import OpenAI
+
+client = OpenAI(api_key="EMPTY", base_url="http://localhost:30010/v1")
+
+response = client.images.generate(
+ model="ideogram-ai/ideogram-4-nf4",
+ prompt="A cinematic poster of a quiet bookstore at dusk with elegant hand-lettered signage",
+ size="1024x1024",
+ n=1,
+ response_format="b64_json",
+ extra_body={"preset": "V4_QUALITY_48", "seed": 0},
+)
+
+image_bytes = base64.b64decode(response.data[0].b64_json)
+with open("ideogram4.png", "wb") as f:
+ f.write(image_bytes)
+```
+
+Ideogram 4 presets are `V4_DEFAULT_20`, `V4_QUALITY_48`, and `V4_TURBO_12`. The preset controls both `num_inference_steps` and guidance, so do not set those fields directly.
diff --git a/docs_new/cookbook/diffusion/intro.mdx b/docs_new/cookbook/diffusion/intro.mdx
index a4235d52d..06475776d 100644
--- a/docs_new/cookbook/diffusion/intro.mdx
+++ b/docs_new/cookbook/diffusion/intro.mdx
@@ -23,6 +23,12 @@ Offline models generate each image or video request as a bounded denoising job.
href="/cookbook/diffusion/FLUX/FLUX"
img="/cards/logos/flux.png"
/>
+
torch.distributed.ProcessGroup | None:
- if config.parallel_folding:
- if config.parallel_folding_mode == "sp":
- return get_sp_group()
- elif config.parallel_folding_mode == "ulysses":
- return get_sp_group().ulysses_group
- elif config.parallel_folding_mode == "ring":
- return get_sp_group().ring_group
- return get_tp_group()
diff --git a/python/sglang/multimodal_gen/runtime/layers/quantization/__init__.py b/python/sglang/multimodal_gen/runtime/layers/quantization/__init__.py
index 1bce2a374..b37e02a8a 100644
--- a/python/sglang/multimodal_gen/runtime/layers/quantization/__init__.py
+++ b/python/sglang/multimodal_gen/runtime/layers/quantization/__init__.py
@@ -2,6 +2,9 @@
from typing import Literal, get_args
+from sglang.multimodal_gen.runtime.layers.quantization.bitsandbytes import (
+ BitsAndBytesConfig,
+)
from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config import (
QuantizationConfig,
)
@@ -25,6 +28,7 @@ QuantizationMethods = Literal[
"modelopt",
"modelopt_fp8",
"modelopt_fp4",
+ "bitsandbytes",
"modelslim",
"mxfp8",
"mxfp4",
@@ -38,6 +42,7 @@ _CUSTOMIZED_METHOD_TO_QUANT_CONFIG = {
"modelopt": ModelOptFp8DiffusionConfig,
"modelopt_fp8": ModelOptFp8Config,
"modelopt_fp4": ModelOptFp4Config,
+ "bitsandbytes": BitsAndBytesConfig,
"modelslim": ModelSlimConfig,
"fp8": Fp8Config,
"mxfp4": Mxfp4Config,
diff --git a/python/sglang/multimodal_gen/runtime/layers/quantization/bitsandbytes.py b/python/sglang/multimodal_gen/runtime/layers/quantization/bitsandbytes.py
new file mode 100644
index 000000000..f5726bbbd
--- /dev/null
+++ b/python/sglang/multimodal_gen/runtime/layers/quantization/bitsandbytes.py
@@ -0,0 +1,383 @@
+# SPDX-License-Identifier: Apache-2.0
+
+from __future__ import annotations
+
+from typing import Any, Optional
+
+import torch
+import torch.nn as nn
+from packaging import version
+
+from sglang.multimodal_gen.runtime.layers.linear import (
+ LinearBase,
+ LinearMethodBase,
+ UnquantizedLinearMethod,
+)
+from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config import (
+ QuantizationConfig,
+ QuantizeMethodBase,
+)
+from sglang.multimodal_gen.runtime.models.utils import set_weight_attrs
+
+
+def _require_bitsandbytes() -> None:
+ try:
+ import bitsandbytes
+
+ if version.parse(bitsandbytes.__version__) < version.parse("0.46.1"):
+ raise ImportError(
+ "bitsandbytes version is wrong. Please install bitsandbytes>=0.46.1."
+ )
+ except ImportError as err:
+ raise ImportError(
+ "Please install bitsandbytes>=0.46.1 via "
+ "`pip install bitsandbytes>=0.46.1` to use bitsandbytes quantizer."
+ ) from err
+
+
+def _calculate_quant_ratio(dtype: torch.dtype) -> int:
+ if dtype.is_floating_point:
+ return torch.finfo(dtype).bits // torch.iinfo(torch.uint8).bits
+ return torch.iinfo(dtype).bits // torch.iinfo(torch.uint8).bits
+
+
+def _is_layer_skipped(prefix: str, skipped_modules: list[str]) -> bool:
+ components = prefix.split(".")
+ if any(module_name in components for module_name in skipped_modules):
+ return True
+
+ prefixes = {".".join(components[: i + 1]) for i in range(len(components))}
+ return bool(set(skipped_modules) & prefixes)
+
+
+class BitsAndBytesConfig(QuantizationConfig):
+ """Config class for pre-quantized bitsandbytes 4-bit checkpoints."""
+
+ def __init__(
+ self,
+ load_in_8bit: bool = False,
+ load_in_4bit: bool = True,
+ bnb_4bit_compute_dtype: str = "float32",
+ bnb_4bit_quant_storage: str = "uint8",
+ bnb_4bit_quant_type: str = "fp4",
+ bnb_4bit_use_double_quant: bool = False,
+ llm_int8_enable_fp32_cpu_offload: bool = False,
+ llm_int8_has_fp16_weight: bool = False,
+ llm_int8_skip_modules: list[str] | None = None,
+ llm_int8_threshold: float = 6.0,
+ ) -> None:
+ super().__init__()
+ self.load_in_8bit = load_in_8bit
+ self.load_in_4bit = load_in_4bit
+ self.bnb_4bit_compute_dtype = bnb_4bit_compute_dtype
+ self.bnb_4bit_quant_storage = bnb_4bit_quant_storage
+ self.bnb_4bit_quant_type = bnb_4bit_quant_type
+ self.bnb_4bit_use_double_quant = bnb_4bit_use_double_quant
+ self.llm_int8_enable_fp32_cpu_offload = llm_int8_enable_fp32_cpu_offload
+ self.llm_int8_has_fp16_weight = llm_int8_has_fp16_weight
+ self.llm_int8_skip_modules = llm_int8_skip_modules or []
+ self.llm_int8_threshold = llm_int8_threshold
+
+ if self.load_in_8bit or not self.load_in_4bit:
+ raise ValueError("SGLang diffusion only supports bitsandbytes 4-bit.")
+ if self.bnb_4bit_quant_storage != "uint8":
+ raise ValueError(
+ f"Unsupported bnb_4bit_quant_storage: {self.bnb_4bit_quant_storage}"
+ )
+
+ @classmethod
+ def get_name(cls) -> str:
+ return "bitsandbytes"
+
+ def get_scaled_act_names(self) -> list[str]:
+ return []
+
+ @classmethod
+ def get_supported_act_dtypes(cls) -> list[torch.dtype]:
+ return [torch.float32, torch.float16, torch.bfloat16]
+
+ @classmethod
+ def get_min_capability(cls) -> int:
+ return 70
+
+ @staticmethod
+ def get_config_filenames() -> list[str]:
+ return []
+
+ @classmethod
+ def from_config(cls, config: dict[str, Any]) -> "BitsAndBytesConfig":
+ def get_safe_value(keys, default_value=None):
+ try:
+ value = QuantizationConfig.get_from_keys(config, keys)
+ return value if value is not None else default_value
+ except ValueError:
+ return default_value
+
+ return cls(
+ load_in_8bit=get_safe_value(["load_in_8bit"], False),
+ load_in_4bit=get_safe_value(["load_in_4bit"], True),
+ bnb_4bit_compute_dtype=get_safe_value(
+ ["bnb_4bit_compute_dtype"], "float32"
+ ),
+ bnb_4bit_quant_storage=get_safe_value(["bnb_4bit_quant_storage"], "uint8"),
+ bnb_4bit_quant_type=get_safe_value(["bnb_4bit_quant_type"], "fp4"),
+ bnb_4bit_use_double_quant=get_safe_value(
+ ["bnb_4bit_use_double_quant"], False
+ ),
+ llm_int8_enable_fp32_cpu_offload=get_safe_value(
+ ["llm_int8_enable_fp32_cpu_offload"], False
+ ),
+ llm_int8_has_fp16_weight=get_safe_value(
+ ["llm_int8_has_fp16_weight"], False
+ ),
+ llm_int8_skip_modules=get_safe_value(["llm_int8_skip_modules"], []),
+ llm_int8_threshold=get_safe_value(["llm_int8_threshold"], 6.0),
+ )
+
+ def get_quant_method(
+ self, layer: torch.nn.Module, prefix: str
+ ) -> Optional[QuantizeMethodBase]:
+ if isinstance(layer, LinearBase):
+ if _is_layer_skipped(prefix, self.llm_int8_skip_modules):
+ return UnquantizedLinearMethod()
+ return BitsAndBytesLinearMethod(self)
+ return None
+
+
+class BitsAndBytesLinearMethod(LinearMethodBase):
+ """Linear method for pre-quantized bitsandbytes 4-bit weights."""
+
+ def __init__(self, quant_config: BitsAndBytesConfig):
+ _require_bitsandbytes()
+ 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,
+ ) -> None:
+ del input_size, output_size
+ quant_ratio = _calculate_quant_ratio(params_dtype)
+ total_size = input_size_per_partition * sum(output_partition_sizes)
+ if total_size % quant_ratio != 0:
+ raise ValueError(
+ "The input size is not aligned with the quantized weight shape."
+ )
+
+ qweight = nn.Parameter(
+ torch.empty(total_size // quant_ratio, 1, dtype=torch.uint8),
+ requires_grad=False,
+ )
+ set_weight_attrs(
+ qweight,
+ {
+ "input_dim": 0,
+ "output_dim": 0,
+ "pack_factor": quant_ratio,
+ "use_bitsandbytes_4bit": True,
+ },
+ )
+ layer.register_parameter("weight", qweight)
+ set_weight_attrs(qweight, extra_weight_attrs)
+
+ def apply(
+ self,
+ layer: torch.nn.Module,
+ x: torch.Tensor,
+ bias: torch.Tensor | None = None,
+ ) -> torch.Tensor:
+ original_type = x.dtype
+ original_shape = x.shape
+ if x.ndim > 2:
+ x = x.reshape(-1, x.size(-1))
+
+ out_dim = sum(
+ quant_state.shape[0]
+ for quant_state in layer.weight.bnb_quant_state.values()
+ )
+ out = torch.empty(x.shape[0], out_dim, dtype=torch.bfloat16, device=x.device)
+ apply_bnb_4bit(x.to(torch.bfloat16), layer.weight, out)
+ out = out.to(original_type)
+
+ if len(original_shape) > 2:
+ out = out.view(*original_shape[:-1], out.size(-1))
+
+ if bias is not None:
+ out = out + bias
+ return out
+
+
+def apply_bnb_4bit(
+ x: torch.Tensor,
+ weight: torch.Tensor,
+ out: torch.Tensor,
+) -> None:
+ from bitsandbytes import matmul_4bit
+
+ offsets = weight.bnb_shard_offsets
+ quant_states = weight.bnb_quant_state
+ current_index = 0
+ for i in range(len(quant_states)):
+ output_size = quant_states[i].shape[0]
+ out[:, current_index : current_index + output_size] = matmul_4bit(
+ x,
+ weight[offsets[i] : offsets[i + 1]].t(),
+ quant_states[i],
+ )
+ current_index += output_size
+
+
+class BitsAndBytes4BitLinear(nn.Module):
+ """Storage-only bitsandbytes 4-bit linear for nn.Linear-based encoders."""
+
+ def __init__(
+ self,
+ in_features: int,
+ out_features: int,
+ bias: bool = True,
+ compute_dtype: torch.dtype | None = None,
+ ) -> None:
+ super().__init__()
+ _require_bitsandbytes()
+ self.in_features = in_features
+ self.out_features = out_features
+ self.compute_dtype = compute_dtype
+ quant_ratio = _calculate_quant_ratio(compute_dtype or torch.get_default_dtype())
+ total_size = in_features * out_features
+ if total_size % quant_ratio != 0:
+ raise ValueError(
+ "The input size is not aligned with the quantized weight shape."
+ )
+
+ self.weight = nn.Parameter(
+ torch.empty(total_size // quant_ratio, 1, dtype=torch.uint8),
+ requires_grad=False,
+ )
+ set_weight_attrs(
+ self.weight,
+ {
+ "pack_factor": quant_ratio,
+ "use_bitsandbytes_4bit": True,
+ },
+ )
+ if bias:
+ self.bias = nn.Parameter(
+ torch.empty(
+ out_features, dtype=compute_dtype or torch.get_default_dtype()
+ ),
+ requires_grad=False,
+ )
+ else:
+ self.register_parameter("bias", None)
+
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
+ original_type = x.dtype
+ original_shape = x.shape
+ if x.ndim > 2:
+ x = x.reshape(-1, x.size(-1))
+
+ out = torch.empty(
+ x.shape[0], self.out_features, dtype=torch.bfloat16, device=x.device
+ )
+ apply_bnb_4bit(x.to(torch.bfloat16), self.weight, out)
+ out = out.to(original_type)
+
+ if len(original_shape) > 2:
+ out = out.view(*original_shape[:-1], out.size(-1))
+
+ if self.bias is not None:
+ out = out + self.bias
+ return out
+
+
+def swap_linears_to_bitsandbytes_4bit(module: nn.Module) -> None:
+ for name, child in list(module.named_children()):
+ if isinstance(child, nn.Linear):
+ replacement = BitsAndBytes4BitLinear(
+ child.in_features,
+ child.out_features,
+ bias=child.bias is not None,
+ compute_dtype=child.weight.dtype,
+ )
+ setattr(module, name, replacement)
+ else:
+ swap_linears_to_bitsandbytes_4bit(child)
+
+
+_BNB_4BIT_STATE_SUFFIXES = {
+ "absmax",
+ "quant_map",
+ "nested_absmax",
+ "nested_quant_map",
+ "bitsandbytes",
+}
+
+
+def is_bitsandbytes_4bit_state_name(weight_name: str) -> bool:
+ suffix = weight_name.split(".")[-1]
+ return any(state_suffix in suffix for state_suffix in _BNB_4BIT_STATE_SUFFIXES)
+
+
+def split_bitsandbytes_4bit_state(
+ weights: Any,
+) -> tuple[list[tuple[str, torch.Tensor]], dict[str, torch.Tensor]]:
+ normal_weights: list[tuple[str, torch.Tensor]] = []
+ quant_state_dict: dict[str, torch.Tensor] = {}
+ for name, tensor in weights:
+ if is_bitsandbytes_4bit_state_name(name):
+ if "quant_state.bitsandbytes" in name:
+ tensor = tensor.cpu().data
+ quant_state_dict[name] = tensor
+ continue
+ normal_weights.append((name, tensor))
+ return normal_weights, quant_state_dict
+
+
+def build_bitsandbytes_4bit_quant_states(
+ normal_weight_names: list[str],
+ quant_state_dict: dict[str, torch.Tensor],
+ device: torch.device,
+ param_names_mapping=None,
+) -> dict[str, Any]:
+ from bitsandbytes.functional import QuantState
+
+ quant_states: dict[str, Any] = {}
+ device_str = str(device)
+ for source_name in normal_weight_names:
+ if (
+ f"{source_name}.quant_state.bitsandbytes__nf4" not in quant_state_dict
+ and f"{source_name}.quant_state.bitsandbytes__fp4" not in quant_state_dict
+ ):
+ continue
+ target_name = source_name
+ if param_names_mapping is not None:
+ target_name, _, _ = param_names_mapping(source_name)
+ state_tensors = {
+ name: tensor
+ for name, tensor in quant_state_dict.items()
+ if name.startswith(f"{source_name}.")
+ }
+ quant_states[target_name] = QuantState.from_dict(
+ state_tensors, device=device_str
+ )
+ return quant_states
+
+
+def attach_bitsandbytes_4bit_quant_states(
+ params_dict: dict[str, torch.nn.Parameter],
+ quant_states: dict[str, Any],
+) -> None:
+ for param_name, quant_state in quant_states.items():
+ param = params_dict.get(param_name)
+ if param is None:
+ raise ValueError(f"Parameter {param_name} not found in the model.")
+
+ state_by_shard = {0: quant_state}
+ set_weight_attrs(param, {"bnb_quant_state": state_by_shard})
+ offsets = torch.tensor([0, param.numel()]).cpu()
+ set_weight_attrs(param, {"bnb_shard_offsets": offsets})
diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/text_encoder_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/text_encoder_loader.py
index f5a59a9d0..85c10eec4 100644
--- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/text_encoder_loader.py
+++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/text_encoder_loader.py
@@ -306,6 +306,18 @@ class TextEncoderLoader(ComponentLoader):
fsdp_cpu_offload = False
should_offload = False
+ if (
+ getattr(
+ model_config.arch_config, "requires_gpu_resident_text_encoder", False
+ )
+ and should_offload
+ ):
+ logger.warning(
+ "Keeping bitsandbytes 4-bit text encoder GPU-resident; CUDA "
+ "weights and quant states are required for this checkpoint."
+ )
+ should_offload = False
+
if should_offload and not current_platform.is_mps():
model_device = torch.device("cpu")
else:
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 81cdf1187..3102bd4f6 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
@@ -34,6 +34,23 @@ def _server_args_for_transformer_component(
if component_name not in ("transformer_2", "unconditional_transformer"):
return server_args
+ # Some pipelines have secondary DiT components with their own quantized
+ # weight file. Keep the mapping model-owned and the loader generic.
+ component_weights_paths = getattr(
+ server_args, "component_transformer_weights_paths", {}
+ )
+ component_weights_path = component_weights_paths.get(component_name)
+ if component_weights_path is not None:
+ component_server_args = copy.copy(server_args)
+ component_server_args.transformer_weights_path = component_weights_path
+ component_server_args.nunchaku_config = None
+ logger.info(
+ "Using transformer_weights_path override for %s: %s",
+ component_name,
+ component_weights_path,
+ )
+ return component_server_args
+
if (
server_args.transformer_weights_path is None
and server_args.nunchaku_config is None
diff --git a/python/sglang/multimodal_gen/runtime/loader/fsdp_load.py b/python/sglang/multimodal_gen/runtime/loader/fsdp_load.py
index efff17ba7..e660fcc39 100644
--- a/python/sglang/multimodal_gen/runtime/loader/fsdp_load.py
+++ b/python/sglang/multimodal_gen/runtime/loader/fsdp_load.py
@@ -25,6 +25,11 @@ from torch.nn.modules.module import _IncompatibleKeys
from sglang.multimodal_gen.configs.models.fsdp import is_module_list_entry_in
from sglang.multimodal_gen.runtime.layers.linear import UnquantizedLinearMethod
+from sglang.multimodal_gen.runtime.layers.quantization.bitsandbytes import (
+ attach_bitsandbytes_4bit_quant_states,
+ build_bitsandbytes_4bit_quant_states,
+ split_bitsandbytes_4bit_state,
+)
from sglang.multimodal_gen.runtime.loader.utils import (
get_param_names_mapping,
hf_to_custom_state_dict,
@@ -51,6 +56,13 @@ _QUANTIZED_DTYPES = (
_DTYPE_MISMATCH_EXAMPLE_LIMIT = 3
+def _is_bitsandbytes_quant_config(quant_config: Any | None) -> bool:
+ if quant_config is None:
+ return False
+ quant_name_getter = getattr(type(quant_config), "get_name", None)
+ return bool(callable(quant_name_getter) and quant_name_getter() == "bitsandbytes")
+
+
def _format_dtype_mismatch_summary(
mismatch_counts: Counter[tuple[torch.dtype, torch.dtype]],
mismatch_examples: dict[tuple[torch.dtype, torch.dtype], list[str]],
@@ -244,11 +256,21 @@ def maybe_load_fsdp_model(
pin_cpu_memory=pin_cpu_memory,
)
+ param_names_mapping_fn = get_param_names_mapping(model.param_names_mapping)
weight_iterator = safetensors_weights_iterator(weight_dir_list)
preprocess_loaded_state_dict = getattr(model, "preprocess_loaded_state_dict", None)
if preprocess_loaded_state_dict is not None:
weight_iterator = preprocess_loaded_state_dict(weight_iterator)
- param_names_mapping_fn = get_param_names_mapping(model.param_names_mapping)
+ bnb_quant_states = None
+ if _is_bitsandbytes_quant_config(init_params.get("quant_config")):
+ normal_weights, raw_quant_state = split_bitsandbytes_4bit_state(weight_iterator)
+ bnb_quant_states = build_bitsandbytes_4bit_quant_states(
+ [name for name, _ in normal_weights],
+ raw_quant_state,
+ device,
+ param_names_mapping_fn,
+ )
+ weight_iterator = iter(normal_weights)
load_model_from_full_model_state_dict(
model,
weight_iterator,
@@ -258,6 +280,10 @@ def maybe_load_fsdp_model(
cpu_offload=cpu_offload,
param_names_mapping=param_names_mapping_fn,
)
+ if bnb_quant_states:
+ attach_bitsandbytes_4bit_quant_states(
+ dict(model.named_parameters()), bnb_quant_states
+ )
for _, module in model.named_modules():
quant_method = getattr(module, "quant_method", None)
diff --git a/python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py b/python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py
index fa3634727..21aced3f5 100644
--- a/python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py
+++ b/python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py
@@ -270,6 +270,46 @@ class _ModelOptFp8OffloadAdapter(_TransformerQuantAdapter):
)
+class _BitsAndBytes4BitAdapter(_TransformerQuantAdapter):
+ """Adapter for pre-quantized bitsandbytes 4-bit transformer checkpoints."""
+
+ def __init__(
+ self,
+ *,
+ server_args: ServerArgs,
+ quant_config: Optional[QuantizationConfig],
+ ) -> None:
+ self.server_args = server_args
+ self.quant_config = quant_config
+
+ @staticmethod
+ def _maybe_disable_incompatible_offload_modes(
+ server_args: ServerArgs,
+ quant_config: Optional[QuantizationConfig],
+ ) -> None:
+ if _get_quant_config_name(quant_config) != "bitsandbytes":
+ return
+
+ changed = []
+ if server_args.dit_cpu_offload:
+ server_args.dit_cpu_offload = False
+ changed.append("dit_cpu_offload=False")
+ if server_args.use_fsdp_inference:
+ server_args.use_fsdp_inference = False
+ changed.append("use_fsdp_inference=False")
+ if changed:
+ logger.warning(
+ "Keeping bitsandbytes 4-bit transformer GPU-resident: %s",
+ ", ".join(changed),
+ )
+
+ def prepare(self) -> None:
+ _BitsAndBytes4BitAdapter._maybe_disable_incompatible_offload_modes(
+ 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]:
@@ -441,6 +481,10 @@ def _build_transformer_quant_adapters(
server_args=server_args,
quant_config=quant_config,
),
+ _BitsAndBytes4BitAdapter(
+ server_args=server_args,
+ quant_config=quant_config,
+ ),
]
if nunchaku_config is not None:
adapters.append(
diff --git a/python/sglang/multimodal_gen/runtime/models/dits/ideogram.py b/python/sglang/multimodal_gen/runtime/models/dits/ideogram.py
index fd78b0a91..272196996 100644
--- a/python/sglang/multimodal_gen/runtime/models/dits/ideogram.py
+++ b/python/sglang/multimodal_gen/runtime/models/dits/ideogram.py
@@ -12,6 +12,10 @@ from sglang.multimodal_gen.runtime.layers.attention import (
USPAttention,
build_varlen_mask_meta,
)
+from sglang.multimodal_gen.runtime.layers.linear import ReplicatedLinear
+from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config import (
+ QuantizationConfig,
+)
from sglang.multimodal_gen.runtime.layers.quantization.weight_only_fp8 import (
WeightOnlyFP8Linear,
)
@@ -35,8 +39,27 @@ class Ideogram4RMSNorm(nn.Module):
return F.rms_norm(x, self.weight.shape, self.weight, self.eps)
-def _linear(in_features: int, out_features: int, bias: bool = True):
- return WeightOnlyFP8Linear(in_features, out_features, bias=bias)
+class Ideogram4QuantizedLinear(ReplicatedLinear):
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
+ return super().forward(x)[0]
+
+
+def _linear(
+ in_features: int,
+ out_features: int,
+ bias: bool = True,
+ quant_config: QuantizationConfig | None = None,
+ prefix: str = "",
+):
+ if quant_config is None:
+ return WeightOnlyFP8Linear(in_features, out_features, bias=bias)
+ return Ideogram4QuantizedLinear(
+ in_features,
+ out_features,
+ bias=bias,
+ quant_config=quant_config,
+ prefix=prefix,
+ )
class Ideogram4Attention(nn.Module):
@@ -46,12 +69,20 @@ class Ideogram4Attention(nn.Module):
num_heads: int,
eps: float,
supported_attention_backends,
+ quant_config: QuantizationConfig | None = None,
+ prefix: str = "",
) -> None:
super().__init__()
self.hidden_size = hidden_size
self.num_heads = num_heads
self.head_dim = hidden_size // num_heads
- self.qkv = _linear(hidden_size, hidden_size * 3, bias=False)
+ self.qkv = _linear(
+ hidden_size,
+ hidden_size * 3,
+ bias=False,
+ quant_config=quant_config,
+ prefix=f"{prefix}.qkv",
+ )
self.norm_q = Ideogram4RMSNorm(self.head_dim, eps=eps)
self.norm_k = Ideogram4RMSNorm(self.head_dim, eps=eps)
self.attn = USPAttention(
@@ -62,7 +93,13 @@ class Ideogram4Attention(nn.Module):
causal=False,
supported_attention_backends=supported_attention_backends,
)
- self.o = _linear(hidden_size, hidden_size, bias=False)
+ self.o = _linear(
+ hidden_size,
+ hidden_size,
+ bias=False,
+ quant_config=quant_config,
+ prefix=f"{prefix}.o",
+ )
def forward(self, x, cos, sin, attn_mask, attn_mask_meta):
batch_size, seq_len, _ = x.shape
@@ -77,11 +114,35 @@ class Ideogram4Attention(nn.Module):
class Ideogram4MLP(nn.Module):
- def __init__(self, dim: int, hidden_dim: int) -> None:
+ def __init__(
+ self,
+ dim: int,
+ hidden_dim: int,
+ quant_config: QuantizationConfig | None = None,
+ prefix: str = "",
+ ) -> None:
super().__init__()
- self.w1 = _linear(dim, hidden_dim, bias=False)
- self.w2 = _linear(hidden_dim, dim, bias=False)
- self.w3 = _linear(dim, hidden_dim, bias=False)
+ self.w1 = _linear(
+ dim,
+ hidden_dim,
+ bias=False,
+ quant_config=quant_config,
+ prefix=f"{prefix}.w1",
+ )
+ self.w2 = _linear(
+ hidden_dim,
+ dim,
+ bias=False,
+ quant_config=quant_config,
+ prefix=f"{prefix}.w2",
+ )
+ self.w3 = _linear(
+ dim,
+ hidden_dim,
+ bias=False,
+ quant_config=quant_config,
+ prefix=f"{prefix}.w3",
+ )
def forward(self, x):
return self.w2(F.silu(self.w1(x)) * self.w3(x))
@@ -96,6 +157,8 @@ class Ideogram4TransformerBlock(nn.Module):
norm_eps,
adaln_dim,
supported_attention_backends,
+ quant_config: QuantizationConfig | None = None,
+ prefix: str = "",
):
super().__init__()
self.attention = Ideogram4Attention(
@@ -103,13 +166,26 @@ class Ideogram4TransformerBlock(nn.Module):
num_heads,
eps=1e-5,
supported_attention_backends=supported_attention_backends,
+ quant_config=quant_config,
+ prefix=f"{prefix}.attention",
+ )
+ self.feed_forward = Ideogram4MLP(
+ hidden_size,
+ intermediate_size,
+ quant_config=quant_config,
+ prefix=f"{prefix}.feed_forward",
)
- self.feed_forward = Ideogram4MLP(hidden_size, intermediate_size)
self.attention_norm1 = Ideogram4RMSNorm(hidden_size, eps=norm_eps)
self.ffn_norm1 = Ideogram4RMSNorm(hidden_size, eps=norm_eps)
self.attention_norm2 = Ideogram4RMSNorm(hidden_size, eps=norm_eps)
self.ffn_norm2 = Ideogram4RMSNorm(hidden_size, eps=norm_eps)
- self.adaln_modulation = _linear(adaln_dim, 4 * hidden_size, bias=True)
+ self.adaln_modulation = _linear(
+ adaln_dim,
+ 4 * hidden_size,
+ bias=True,
+ quant_config=quant_config,
+ prefix=f"{prefix}.adaln_modulation",
+ )
def forward(self, x, cos, sin, adaln_input, attn_mask, attn_mask_meta):
scale_msa, gate_msa, scale_mlp, gate_mlp = self.adaln_modulation(
@@ -144,12 +220,30 @@ def _sinusoidal_embedding(t: torch.Tensor, dim: int, scale: float = 1e4):
class Ideogram4EmbedScalar(nn.Module):
- def __init__(self, dim: int, input_range: tuple[float, float]) -> None:
+ def __init__(
+ self,
+ dim: int,
+ input_range: tuple[float, float],
+ quant_config: QuantizationConfig | None = None,
+ prefix: str = "",
+ ) -> None:
super().__init__()
self.dim = dim
self.range_min, self.range_max = input_range
- self.mlp_in = _linear(dim, dim, bias=True)
- self.mlp_out = _linear(dim, dim, bias=True)
+ self.mlp_in = _linear(
+ dim,
+ dim,
+ bias=True,
+ quant_config=quant_config,
+ prefix=f"{prefix}.mlp_in",
+ )
+ self.mlp_out = _linear(
+ dim,
+ dim,
+ bias=True,
+ quant_config=quant_config,
+ prefix=f"{prefix}.mlp_out",
+ )
def forward(self, x):
compute_dtype = x.dtype
@@ -160,11 +254,30 @@ class Ideogram4EmbedScalar(nn.Module):
class Ideogram4FinalLayer(nn.Module):
- def __init__(self, hidden_size: int, out_channels: int, adaln_dim: int) -> None:
+ def __init__(
+ self,
+ hidden_size: int,
+ out_channels: int,
+ adaln_dim: int,
+ quant_config: QuantizationConfig | None = None,
+ prefix: str = "",
+ ) -> None:
super().__init__()
self.norm_final = nn.LayerNorm(hidden_size, eps=1e-6, elementwise_affine=False)
- self.linear = _linear(hidden_size, out_channels, bias=True)
- self.adaln_modulation = _linear(adaln_dim, hidden_size, bias=True)
+ self.linear = _linear(
+ hidden_size,
+ out_channels,
+ bias=True,
+ quant_config=quant_config,
+ prefix=f"{prefix}.linear",
+ )
+ self.adaln_modulation = _linear(
+ adaln_dim,
+ hidden_size,
+ bias=True,
+ quant_config=quant_config,
+ prefix=f"{prefix}.adaln_modulation",
+ )
def forward(self, x, c):
scale = 1.0 + self.adaln_modulation(F.silu(c))
@@ -180,12 +293,12 @@ class Ideogram4Transformer2DModel(BaseDiT):
)
param_names_mapping = {}
reverse_param_names_mapping = {}
- handles_checkpoint_quantization = True
def __init__(
self,
config: Ideogram4DiTConfig,
hf_config: dict[str, Any],
+ quant_config: QuantizationConfig | None = None,
**kwargs,
) -> None:
super().__init__(config, hf_config, **kwargs)
@@ -195,11 +308,34 @@ class Ideogram4Transformer2DModel(BaseDiT):
self.hidden_size = hidden_size
self.num_attention_heads = cfg.num_attention_heads
self.num_channels_latents = cfg.in_channels
- self.input_proj = _linear(cfg.in_channels, hidden_size, bias=True)
+ self.input_proj = _linear(
+ cfg.in_channels,
+ hidden_size,
+ bias=True,
+ quant_config=quant_config,
+ prefix="input_proj",
+ )
self.llm_cond_norm = Ideogram4RMSNorm(cfg.llm_features_dim, eps=1e-6)
- self.llm_cond_proj = _linear(cfg.llm_features_dim, hidden_size, bias=True)
- self.t_embedding = Ideogram4EmbedScalar(hidden_size, input_range=(0.0, 1.0))
- self.adaln_proj = _linear(hidden_size, cfg.adaln_dim, bias=True)
+ self.llm_cond_proj = _linear(
+ cfg.llm_features_dim,
+ hidden_size,
+ bias=True,
+ quant_config=quant_config,
+ prefix="llm_cond_proj",
+ )
+ self.t_embedding = Ideogram4EmbedScalar(
+ hidden_size,
+ input_range=(0.0, 1.0),
+ quant_config=quant_config,
+ prefix="t_embedding",
+ )
+ self.adaln_proj = _linear(
+ hidden_size,
+ cfg.adaln_dim,
+ bias=True,
+ quant_config=quant_config,
+ prefix="adaln_proj",
+ )
self.embed_image_indicator = nn.Embedding(2, hidden_size)
self.rotary_emb = Qwen3VLTextRotaryEmbedding(
head_dim=cfg.attention_head_dim,
@@ -215,14 +351,18 @@ class Ideogram4Transformer2DModel(BaseDiT):
norm_eps=cfg.norm_eps,
adaln_dim=cfg.adaln_dim,
supported_attention_backends=self._supported_attention_backends,
+ quant_config=quant_config,
+ prefix=f"layers.{i}",
)
- for _ in range(cfg.num_layers)
+ for i in range(cfg.num_layers)
]
)
self.final_layer = Ideogram4FinalLayer(
hidden_size=hidden_size,
out_channels=cfg.in_channels,
adaln_dim=cfg.adaln_dim,
+ quant_config=quant_config,
+ prefix="final_layer",
)
def post_load_weights(self) -> None:
diff --git a/python/sglang/multimodal_gen/runtime/models/encoders/ideogram.py b/python/sglang/multimodal_gen/runtime/models/encoders/ideogram.py
index d7df4c184..86a38475c 100644
--- a/python/sglang/multimodal_gen/runtime/models/encoders/ideogram.py
+++ b/python/sglang/multimodal_gen/runtime/models/encoders/ideogram.py
@@ -10,6 +10,12 @@ from sglang.multimodal_gen.configs.models.encoders import BaseEncoderOutput
from sglang.multimodal_gen.configs.models.encoders.ideogram import (
Ideogram4TextEncoderConfig,
)
+from sglang.multimodal_gen.runtime.layers.quantization.bitsandbytes import (
+ attach_bitsandbytes_4bit_quant_states,
+ build_bitsandbytes_4bit_quant_states,
+ is_bitsandbytes_4bit_state_name,
+ swap_linears_to_bitsandbytes_4bit,
+)
from sglang.multimodal_gen.runtime.layers.quantization.weight_only_fp8 import (
swap_linears_to_weight_only_fp8,
)
@@ -31,7 +37,12 @@ class IdeogramQwen3VLTextEncoder(TextEncoder):
if isinstance(text_config, dict):
text_config = Qwen3VLTextConfig(**text_config)
self.language_model = Qwen3VLTextModel(text_config)
- if getattr(arch_config, "ideogram_fp8_weight_only", False):
+ self._uses_bitsandbytes_4bit = getattr(
+ arch_config, "ideogram_bnb_4bit_weight_only", False
+ )
+ if self._uses_bitsandbytes_4bit:
+ swap_linears_to_bitsandbytes_4bit(self.language_model)
+ elif getattr(arch_config, "ideogram_fp8_weight_only", False):
swap_linears_to_weight_only_fp8(self.language_model)
@torch.no_grad()
@@ -100,6 +111,9 @@ class IdeogramQwen3VLTextEncoder(TextEncoder):
return features
def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
+ if self._uses_bitsandbytes_4bit:
+ return self._load_bitsandbytes_4bit_weights(weights)
+
loaded_params: set[str] = set()
params_dict = dict(self.named_parameters(remove_duplicate=False))
for name, loaded_weight in weights:
@@ -117,5 +131,51 @@ class IdeogramQwen3VLTextEncoder(TextEncoder):
loaded_params.add(name)
return loaded_params
+ def _load_bitsandbytes_4bit_weights(
+ self, weights: Iterable[Tuple[str, torch.Tensor]]
+ ):
+ params_dict = dict(self.named_parameters(remove_duplicate=False))
+ raw_quant_state: dict[str, torch.Tensor] = {}
+ normal_weight_names: list[str] = []
+ loaded_params: set[str] = set()
+ for name, loaded_weight in weights:
+ if is_bitsandbytes_4bit_state_name(name):
+ if "quant_state.bitsandbytes" in name:
+ loaded_weight = loaded_weight.cpu().data
+ raw_quant_state[name] = loaded_weight
+ continue
+ if name.startswith("visual."):
+ continue
+ if "rotary_emb.inv_freq" in name:
+ continue
+ param = params_dict.get(name)
+ if param is None:
+ raise KeyError(
+ f"Unexpected weight name while loading Ideogram text encoder: {name}"
+ )
+ weight_loader = getattr(param, "weight_loader", default_weight_loader)
+ weight_loader(param, loaded_weight.to(param.dtype))
+ normal_weight_names.append(name)
+ loaded_params.add(name)
+
+ quant_states = build_bitsandbytes_4bit_quant_states(
+ normal_weight_names,
+ raw_quant_state,
+ next(self.parameters()).device,
+ )
+ attach_bitsandbytes_4bit_quant_states(params_dict, quant_states)
+ quantized_params_missing_state = [
+ name
+ for name, param in params_dict.items()
+ if getattr(param, "use_bitsandbytes_4bit", False)
+ and name not in quant_states
+ ]
+ if quantized_params_missing_state:
+ raise ValueError(
+ "Missing bitsandbytes quant_state for Ideogram text encoder weights: "
+ f"{quantized_params_missing_state[:8]}"
+ )
+ return loaded_params
+
EntryClass = IdeogramQwen3VLTextEncoder
diff --git a/python/sglang/multimodal_gen/runtime/models/encoders/t5.py b/python/sglang/multimodal_gen/runtime/models/encoders/t5.py
index 5058de389..3a1b931ab 100644
--- a/python/sglang/multimodal_gen/runtime/models/encoders/t5.py
+++ b/python/sglang/multimodal_gen/runtime/models/encoders/t5.py
@@ -30,7 +30,7 @@ import torch.nn.functional as F
from torch import nn
from sglang.multimodal_gen.configs.models.encoders import BaseEncoderOutput, T5Config
-from sglang.multimodal_gen.runtime.distributed import _get_folding_tp_group
+from sglang.multimodal_gen.runtime.distributed import get_sp_group, get_tp_group
from sglang.multimodal_gen.runtime.layers.activation import get_act_fn
from sglang.multimodal_gen.runtime.layers.layernorm import RMSNorm
from sglang.multimodal_gen.runtime.layers.linear import (
@@ -48,6 +48,19 @@ from sglang.multimodal_gen.runtime.models.encoders.base import TextEncoder
from sglang.multimodal_gen.runtime.platforms import current_platform
+def _get_folding_tp_group(
+ config: T5Config,
+) -> torch.distributed.ProcessGroup | None:
+ if config.parallel_folding:
+ if config.parallel_folding_mode == "sp":
+ return get_sp_group()
+ elif config.parallel_folding_mode == "ulysses":
+ return get_sp_group().ulysses_group
+ elif config.parallel_folding_mode == "ring":
+ return get_sp_group().ring_group
+ return get_tp_group()
+
+
class AttentionType:
"""
Attention type.
diff --git a/python/sglang/multimodal_gen/runtime/pipelines/ideogram.py b/python/sglang/multimodal_gen/runtime/pipelines/ideogram.py
index 796f1d5d9..92c7c369c 100644
--- a/python/sglang/multimodal_gen/runtime/pipelines/ideogram.py
+++ b/python/sglang/multimodal_gen/runtime/pipelines/ideogram.py
@@ -1,5 +1,10 @@
# SPDX-License-Identifier: Apache-2.0
+import os
+from dataclasses import dataclass
+from functools import lru_cache
+from typing import Any, cast
+
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
from sglang.multimodal_gen.runtime.pipelines_core import LoRAPipeline
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
@@ -14,6 +19,88 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.i
Ideogram4TextEncodingStage,
)
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__)
+
+_IDEOGRAM4_BASE_MODEL = "ideogram-ai/ideogram-4-fp8"
+_IDEOGRAM4_NVFP4_COND_FILE = "diffusion_models/ideogram4_nvfp4_mixed.safetensors"
+_IDEOGRAM4_NVFP4_UNCOND_FILE = (
+ "diffusion_models/ideogram4_unconditional_nvfp4_mixed.safetensors"
+)
+
+
+@dataclass(frozen=True)
+class Ideogram4Nvfp4ModelResolution:
+ base_model_name: str
+ base_model_path: str
+ transformer_weights_path: str
+ unconditional_transformer_weights_path: str | None
+
+
+@lru_cache(maxsize=1)
+def _resolve_ideogram4_base_model_path() -> str:
+ return maybe_download_model(_IDEOGRAM4_BASE_MODEL, force_diffusers_model=True)
+
+
+def _resolve_ideogram4_unconditional_transformer_weights_path(
+ transformer_weights_path: str,
+) -> str | None:
+ if os.path.basename(transformer_weights_path) != os.path.basename(
+ _IDEOGRAM4_NVFP4_COND_FILE
+ ):
+ return None
+ return os.path.join(
+ os.path.dirname(transformer_weights_path),
+ os.path.basename(_IDEOGRAM4_NVFP4_UNCOND_FILE),
+ )
+
+
+def _resolve_ideogram4_nvfp4_transformer_weights_paths(
+ server_args: ServerArgs, model_path: str
+) -> tuple[str, str | None]:
+ if server_args.transformer_weights_path is not None:
+ transformer_weights_path = server_args.transformer_weights_path
+ return (
+ transformer_weights_path,
+ _resolve_ideogram4_unconditional_transformer_weights_path(
+ transformer_weights_path
+ ),
+ )
+
+ local_nvfp4_path = maybe_download_model(
+ model_path,
+ allow_patterns=[
+ _IDEOGRAM4_NVFP4_COND_FILE,
+ _IDEOGRAM4_NVFP4_UNCOND_FILE,
+ ],
+ )
+ return (
+ os.path.join(local_nvfp4_path, _IDEOGRAM4_NVFP4_COND_FILE),
+ os.path.join(local_nvfp4_path, _IDEOGRAM4_NVFP4_UNCOND_FILE),
+ )
+
+
+def resolve_ideogram4_nvfp4_model(
+ server_args: ServerArgs, model_path: str
+) -> Ideogram4Nvfp4ModelResolution:
+ (
+ transformer_weights_path,
+ unconditional_transformer_weights_path,
+ ) = _resolve_ideogram4_nvfp4_transformer_weights_paths(
+ server_args,
+ model_path,
+ )
+ return Ideogram4Nvfp4ModelResolution(
+ base_model_name=_IDEOGRAM4_BASE_MODEL,
+ base_model_path=_resolve_ideogram4_base_model_path(),
+ transformer_weights_path=transformer_weights_path,
+ unconditional_transformer_weights_path=unconditional_transformer_weights_path,
+ )
class Ideogram4Pipeline(LoRAPipeline, ComposedPipelineBase):
@@ -55,4 +142,85 @@ class Ideogram4Pipeline(LoRAPipeline, ComposedPipelineBase):
)
-EntryClass = Ideogram4Pipeline
+class Ideogram4Nvfp4Pipeline(Ideogram4Pipeline):
+ pipeline_name = "Ideogram4Nvfp4Pipeline"
+ _model_resolution: Ideogram4Nvfp4ModelResolution | None = None
+
+ def _get_model_resolution(
+ self,
+ server_args: ServerArgs | None = None,
+ ) -> Ideogram4Nvfp4ModelResolution:
+ if self._model_resolution is None:
+ if server_args is None:
+ raise ValueError(
+ "server_args is required to resolve Ideogram4 NVFP4 paths"
+ )
+ self._model_resolution = resolve_ideogram4_nvfp4_model(
+ server_args,
+ self.model_path,
+ )
+ return self._model_resolution
+
+ def _load_config(self) -> dict[str, Any]:
+ 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",
+ model_resolution.base_model_name,
+ model_resolution.base_model_path,
+ )
+ config = verify_model_config_and_directory(model_resolution.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)
+
+ component_model_path = os.path.join(
+ 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
+
+ def load_modules(
+ self,
+ server_args: ServerArgs,
+ loaded_modules: dict | None = None,
+ ) -> dict:
+ model_resolution = self._get_model_resolution(server_args)
+ server_args.transformer_weights_path = model_resolution.transformer_weights_path
+ if model_resolution.unconditional_transformer_weights_path is not None:
+ # The loader treats transformer_weights_path as the base DiT override.
+ # Route the sibling unconditional DiT weights through the generic
+ # per-component override map instead of hard-coding Ideogram there.
+ component_transformer_weights_paths = dict(
+ getattr(server_args, "component_transformer_weights_paths", {})
+ )
+ component_transformer_weights_paths.setdefault(
+ "unconditional_transformer",
+ model_resolution.unconditional_transformer_weights_path,
+ )
+ server_args.component_transformer_weights_paths = (
+ component_transformer_weights_paths
+ )
+ logger.info(
+ "NVFP4 transformer weights: %s",
+ model_resolution.transformer_weights_path,
+ )
+ logger.info(
+ "NVFP4 unconditional transformer weights: %s",
+ server_args.component_transformer_weights_paths.get(
+ "unconditional_transformer"
+ ),
+ )
+ return super().load_modules(server_args, loaded_modules)
+
+
+EntryClass = [Ideogram4Pipeline, Ideogram4Nvfp4Pipeline]
diff --git a/python/sglang/multimodal_gen/runtime/server_args.py b/python/sglang/multimodal_gen/runtime/server_args.py
index fcda04afd..bfacbb9fd 100644
--- a/python/sglang/multimodal_gen/runtime/server_args.py
+++ b/python/sglang/multimodal_gen/runtime/server_args.py
@@ -184,6 +184,11 @@ class ServerArgs(DisaggServerArgsMixin):
# path to pre-quantized transformer weights (single .safetensors or directory).
transformer_weights_path: str | None = None
+ # Per-component transformer weight overrides (key = model_index.json component name).
+ # Pipelines use this when a checkpoint ships separate quantized weights for
+ # secondary DiT components; the generic loader consumes it without model-specific
+ # filename logic.
+ component_transformer_weights_paths: dict[str, str] = field(default_factory=dict)
# Quantization method for online quantization
quantization: str | None = None
diff --git a/python/sglang/multimodal_gen/runtime/utils/quantization_utils.py b/python/sglang/multimodal_gen/runtime/utils/quantization_utils.py
index 168227e7b..4e1c4bf22 100644
--- a/python/sglang/multimodal_gen/runtime/utils/quantization_utils.py
+++ b/python/sglang/multimodal_gen/runtime/utils/quantization_utils.py
@@ -2,6 +2,7 @@ import glob
import json
import os
import re
+import struct
from pathlib import Path
from typing import Any, Dict, List, Optional
@@ -38,8 +39,15 @@ def normalize_flat_modelopt_quant_config(
def _infer_nvfp4_group_size_from_tensors(weight, scale) -> Optional[int]:
"""Infer NVFP4 group_size from serialized weight/scale tensor shapes."""
- weight_shape = tuple(getattr(weight, "shape", ()))
- scale_shape = tuple(getattr(scale, "shape", ()))
+ return _infer_nvfp4_group_size_from_shapes(
+ getattr(weight, "shape", ()),
+ getattr(scale, "shape", ()),
+ )
+
+
+def _infer_nvfp4_group_size_from_shapes(weight_shape, scale_shape) -> Optional[int]:
+ weight_shape = tuple(weight_shape or ())
+ scale_shape = tuple(scale_shape or ())
if len(weight_shape) < 2:
return None
@@ -67,9 +75,34 @@ def _infer_nvfp4_group_size_from_tensors(weight, scale) -> Optional[int]:
return None
+def _read_safetensors_tensor_metadata(file_path: str) -> dict[str, dict[str, Any]]:
+ with open(file_path, "rb") as f:
+ header_len = struct.unpack(" bool:
+ weight_metadata = tensor_metadata.get(f"{module_name}.weight")
+ scale_metadata = tensor_metadata.get(f"{module_name}.weight_scale")
+ if weight_metadata is None or scale_metadata is None:
+ return False
+
+ weight_dtype = str(weight_metadata.get("dtype", "")).upper()
+ scale_dtype = str(scale_metadata.get("dtype", "")).upper()
+ scale_shape = scale_metadata.get("shape", [])
+ return weight_dtype == "U8" and "F8_E4M3" in scale_dtype and len(scale_shape) >= 2
+
+
def _resolve_quant_method_name(quant_cfg: dict) -> str:
quant_cfg = normalize_flat_modelopt_quant_config(quant_cfg) or quant_cfg
quant_method = quant_cfg.get("quant_method")
+ if quant_method == "bitsandbytes":
+ return "bitsandbytes"
if quant_method != "modelopt":
return quant_method
@@ -285,6 +318,7 @@ def _build_nvfp4_config_from_safetensors_files(
non_quantized_bfl_modules: set[str] = set()
files_with_nvfp4_signal: list[str] = []
checkpoint_uses_packed_qkv = False
+ checkpoint_uses_comfy_quant = False
packed_qkv_pattern = re.compile(
r"^(double_blocks\.\d+\.(img|txt)_attn\.qkv|single_blocks\.\d+\.linear1)\."
)
@@ -322,21 +356,26 @@ def _build_nvfp4_config_from_safetensors_files(
if isinstance(layer_cfg, dict) and layer_cfg.get("format") == "nvfp4"
)
+ tensor_metadata = _read_safetensors_tensor_metadata(file_path)
with safe_open(file_path, framework="pt", device="cpu") as f:
all_keys = set(f.keys())
if any(packed_qkv_pattern.match(k) for k in all_keys):
checkpoint_uses_packed_qkv = True
+ if any(k.endswith(".comfy_quant") for k in all_keys):
+ checkpoint_uses_comfy_quant = True
# Some ModelOpt NVFP4 exports only store a flat config.json plus
# per-file metadata without the diffusers `layers` section. Infer
- # quantized modules directly from tensor families in that case:
- # quantized modules ship `.weight` + `.weight_scale`, while BF16
- # fallbacks only ship `.weight`.
+ # quantized modules directly from tensor families in that case.
+ # Mixed checkpoints may also contain FP8 fallback layers with scalar
+ # `.weight_scale`, so require packed uint8 weights and block scales.
file_quantized_modules.update(
key[: -len(".weight_scale")]
for key in all_keys
if key.endswith(".weight_scale")
- and f"{key[: -len('.weight_scale')]}.weight" in all_keys
+ and _is_nvfp4_tensor_family(
+ key[: -len(".weight_scale")], tensor_metadata
+ )
)
if file_quantized_modules or metadata_signals_nvfp4:
@@ -347,10 +386,13 @@ def _build_nvfp4_config_from_safetensors_files(
for layer_name in sorted(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)
- group_size = _infer_nvfp4_group_size_from_tensors(w, s)
+ weight_metadata = tensor_metadata.get(weight_key)
+ scale_metadata = tensor_metadata.get(scale_key)
+ if weight_metadata is not None and scale_metadata is not None:
+ group_size = _infer_nvfp4_group_size_from_shapes(
+ weight_metadata.get("shape"),
+ scale_metadata.get("shape"),
+ )
if group_size is not None:
break
@@ -432,28 +474,32 @@ def _build_nvfp4_config_from_safetensors_files(
try:
quant_cls = get_quantization_config("modelopt_fp4")
+ checkpoint_uses_swizzled_scales = (
+ checkpoint_uses_packed_qkv or checkpoint_uses_comfy_quant
+ )
result = quant_cls.from_config(
{
"quant_algo": "NVFP4",
"group_size": group_size,
"ignore": exclude_modules,
"checkpoint_uses_packed_qkv": checkpoint_uses_packed_qkv,
- # The official FLUX.2 mixed NVFP4 export is detected by its
- # packed QKV tensors and stores block scales in the
- # FlashInfer/CUTLASS-swizzled layout. SGLang-converted
- # transformer repos keep the linear layout.
+ # packed-QKV and Comfy NVFP4 checkpoints store serialized
+ # weights/scales in the FlashInfer/CUTLASS checkpoint layout
"checkpoint_weight_scale_layout": (
- "swizzled" if checkpoint_uses_packed_qkv else "linear"
+ "swizzled" if checkpoint_uses_swizzled_scales else "linear"
),
+ "swap_weight_nibbles": checkpoint_uses_swizzled_scales,
}
)
logger.info(
- "Built NVFP4 quant config from %d safetensors: group_size=%d, %d excluded modules, packed_qkv=%s, scale_layout=%s",
+ "Built NVFP4 quant config from %d safetensors: group_size=%d, %d excluded modules, packed_qkv=%s, comfy_quant=%s, scale_layout=%s, swap_nibbles=%s",
len(files_with_nvfp4_signal),
group_size,
len(exclude_modules),
checkpoint_uses_packed_qkv,
+ checkpoint_uses_comfy_quant,
getattr(result, "checkpoint_weight_scale_layout", "linear"),
+ getattr(result, "swap_weight_nibbles", False),
)
return result
except Exception as e:
diff --git a/python/sglang/multimodal_gen/test/server/gpu_cases.py b/python/sglang/multimodal_gen/test/server/gpu_cases.py
index 652aa18ab..09dc41fbe 100644
--- a/python/sglang/multimodal_gen/test/server/gpu_cases.py
+++ b/python/sglang/multimodal_gen/test/server/gpu_cases.py
@@ -553,6 +553,15 @@ else:
env_vars=MODELOPT_NVFP4_B200_ENV_VARS,
run_consistency_check=True,
),
+ _make_modelopt_ci_case(
+ "ideogram4_nvfp4_t2i",
+ model_path="Comfy-Org/Ideogram-4",
+ modality="image",
+ sampling_params=IDEOGRAM4_CI_sampling_params,
+ extras=[],
+ env_vars=MODELOPT_NVFP4_B200_ENV_VARS,
+ run_consistency_check=True,
+ ),
_make_modelopt_ci_case(
"wan22_modelopt_nvfp4_t2v",
model_path=MODELOPT_WAN22_NVFP4_MODEL,
diff --git a/python/sglang/multimodal_gen/test/test_utils.py b/python/sglang/multimodal_gen/test/test_utils.py
index 5f8b2492b..cb53038b6 100644
--- a/python/sglang/multimodal_gen/test/test_utils.py
+++ b/python/sglang/multimodal_gen/test/test_utils.py
@@ -33,7 +33,7 @@ if TYPE_CHECKING:
logger = init_logger(__name__)
-SGL_TEST_FILES_CI_DATA_REVISION = "50aa0d4d5d4d260302d74b80d97747efd0f0ae45"
+SGL_TEST_FILES_CI_DATA_REVISION = "af6e712a2c49ab5fcd81dde58e2f54c78e77683b"
SGL_TEST_FILES_CONSISTENCY_GT_ROOT = (
"https://raw.githubusercontent.com/"
f"sgl-project/ci-data/{SGL_TEST_FILES_CI_DATA_REVISION}/"
diff --git a/python/sglang/multimodal_gen/test/unit/test_ideogram4.py b/python/sglang/multimodal_gen/test/unit/test_ideogram4.py
index 020c49057..12e103e10 100644
--- a/python/sglang/multimodal_gen/test/unit/test_ideogram4.py
+++ b/python/sglang/multimodal_gen/test/unit/test_ideogram4.py
@@ -24,6 +24,11 @@ from sglang.multimodal_gen.registry import _get_config_info, get_model_info
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType, get_module_role
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
from sglang.multimodal_gen.runtime.layers.attention import USPAttention
+from sglang.multimodal_gen.runtime.layers.linear import UnquantizedLinearMethod
+from sglang.multimodal_gen.runtime.layers.quantization.modelopt_quant import (
+ ModelOptFp4Config,
+ ModelOptFp4LinearMethod,
+)
from sglang.multimodal_gen.runtime.layers.quantization.weight_only_fp8 import (
FP8_WEIGHT_DTYPE,
WeightOnlyFP8Linear,
@@ -47,6 +52,9 @@ from sglang.multimodal_gen.runtime.models.dits.ideogram import (
from sglang.multimodal_gen.runtime.models.encoders.ideogram import (
IdeogramQwen3VLTextEncoder,
)
+from sglang.multimodal_gen.runtime.pipelines.ideogram import (
+ _resolve_ideogram4_unconditional_transformer_weights_path,
+)
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.pipelines_core.stages.denoising import DenoisingStage
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.ideogram import (
@@ -153,6 +161,33 @@ class TestIdeogram4(unittest.TestCase):
self.assertIs(info.pipeline_config_cls, Ideogram4PipelineConfig)
self.assertIs(info.sampling_param_cls, Ideogram4SamplingParams)
+ def test_registry_resolves_comfy_nvfp4_repo_to_native_pipeline(self):
+ get_model_info.cache_clear()
+ _get_config_info.cache_clear()
+
+ info = get_model_info("Comfy-Org/Ideogram-4", backend="sglang")
+
+ self.assertEqual(info.pipeline_cls.__name__, "Ideogram4Nvfp4Pipeline")
+ self.assertIs(info.pipeline_config_cls, Ideogram4PipelineConfig)
+ self.assertIs(info.sampling_param_cls, Ideogram4SamplingParams)
+
+ def test_registry_resolves_official_nf4_repo_to_native_pipeline(self):
+ get_model_info.cache_clear()
+ _get_config_info.cache_clear()
+
+ with patch(
+ "sglang.multimodal_gen.registry.maybe_download_model_index",
+ return_value={
+ "_class_name": "Ideogram4Pipeline",
+ "_diffusers_version": "0.0.0",
+ },
+ ):
+ info = get_model_info("ideogram-ai/ideogram-4-nf4", backend="sglang")
+
+ self.assertEqual(info.pipeline_cls.__name__, "Ideogram4Pipeline")
+ self.assertIs(info.pipeline_config_cls, Ideogram4PipelineConfig)
+ self.assertIs(info.sampling_param_cls, Ideogram4SamplingParams)
+
def test_rowwise_fp8_dequant_uses_output_channel_scale(self):
weight = torch.tensor(
[[1.0, 2.0, -3.0], [4.0, -5.0, 6.0]], dtype=FP8_WEIGHT_DTYPE
@@ -305,6 +340,7 @@ class TestIdeogram4(unittest.TestCase):
server_args = SimpleNamespace(
transformer_weights_path="/unused/override.safetensors",
nunchaku_config={"enabled": True},
+ component_transformer_weights_paths={},
)
component_args = _server_args_for_transformer_component(
server_args, "unconditional_transformer"
@@ -313,6 +349,45 @@ class TestIdeogram4(unittest.TestCase):
self.assertIsNone(component_args.transformer_weights_path)
self.assertIsNone(component_args.nunchaku_config)
+ def test_transformer_component_uses_per_component_weights_override(self):
+ server_args = SimpleNamespace(
+ transformer_weights_path=(
+ "/ckpt/diffusion_models/ideogram4_nvfp4_mixed.safetensors"
+ ),
+ nunchaku_config={"enabled": True},
+ component_transformer_weights_paths={
+ "unconditional_transformer": (
+ "/ckpt/diffusion_models/"
+ "ideogram4_unconditional_nvfp4_mixed.safetensors"
+ )
+ },
+ )
+
+ component_args = _server_args_for_transformer_component(
+ server_args,
+ "unconditional_transformer",
+ )
+
+ self.assertIsNot(component_args, server_args)
+ self.assertEqual(
+ component_args.transformer_weights_path,
+ "/ckpt/diffusion_models/ideogram4_unconditional_nvfp4_mixed.safetensors",
+ )
+ self.assertIsNone(component_args.nunchaku_config)
+
+ def test_ideogram_nvfp4_unconditional_transformer_path_uses_sibling_file(self):
+ self.assertEqual(
+ _resolve_ideogram4_unconditional_transformer_weights_path(
+ "/ckpt/diffusion_models/ideogram4_nvfp4_mixed.safetensors"
+ ),
+ "/ckpt/diffusion_models/ideogram4_unconditional_nvfp4_mixed.safetensors",
+ )
+ self.assertIsNone(
+ _resolve_ideogram4_unconditional_transformer_weights_path(
+ "/ckpt/custom_transformer.safetensors"
+ )
+ )
+
def test_ideogram_denoiser_does_not_request_dtype_cast(self):
import sglang.multimodal_gen.runtime.server_args as server_args_module
@@ -554,6 +629,69 @@ class TestIdeogram4(unittest.TestCase):
)
self.assertEqual(state["layers.0.attention.qkv.weight"].dtype, FP8_WEIGHT_DTYPE)
+ def test_ideogram_dit_nvfp4_quant_config_uses_native_fp4_linears(self):
+ import sglang.multimodal_gen.runtime.server_args as server_args_module
+
+ quant_config = ModelOptFp4Config(
+ is_checkpoint_nvfp4_serialized=True,
+ group_size=16,
+ exclude_modules=[
+ "input_proj",
+ "llm_cond_proj",
+ "t_embedding.*",
+ "adaln_proj",
+ "layers.*.adaln_modulation",
+ "final_layer.*",
+ ],
+ )
+ prev_args = server_args_module._global_server_args
+ try:
+ set_global_server_args(
+ SimpleNamespace(attention_backend="torch_sdpa", comfyui_mode=False)
+ )
+ with patch(
+ "sglang.multimodal_gen.runtime.layers.attention.layer.get_ring_parallel_world_size",
+ return_value=1,
+ ):
+ with torch.device("meta"):
+ model = Ideogram4Transformer2DModel(
+ Ideogram4DiTConfig(),
+ {},
+ quant_config=quant_config,
+ )
+ finally:
+ set_global_server_args(prev_args)
+
+ self.assertEqual(model.layers[0].attention.qkv.prefix, "layers.0.attention.qkv")
+ self.assertIsInstance(
+ model.layers[0].attention.qkv.quant_method,
+ ModelOptFp4LinearMethod,
+ )
+ self.assertIsInstance(model.input_proj.quant_method, UnquantizedLinearMethod)
+
+ state = model.state_dict()
+ self.assertEqual(
+ tuple(state["layers.0.attention.qkv.weight"].shape),
+ (13824, 2304),
+ )
+ self.assertEqual(state["layers.0.attention.qkv.weight"].dtype, torch.uint8)
+ self.assertEqual(
+ tuple(state["layers.0.attention.qkv.weight_scale"].shape),
+ (13824, 288),
+ )
+ self.assertEqual(
+ state["layers.0.attention.qkv.weight_scale"].dtype,
+ FP8_WEIGHT_DTYPE,
+ )
+ self.assertEqual(
+ tuple(state["layers.0.attention.qkv.weight_scale_2"].shape),
+ (1,),
+ )
+ self.assertEqual(
+ tuple(state["layers.0.attention.qkv.input_scale"].shape),
+ (1,),
+ )
+
def test_missing_weight_only_fp8_scale_is_fatal(self):
with torch.device("meta"):
model = WeightOnlyFP8Linear(3, 2, bias=False)
@@ -613,6 +751,27 @@ class TestIdeogram4(unittest.TestCase):
config.arch_config.architectures, ["IdeogramQwen3VLTextEncoder"]
)
self.assertTrue(config.arch_config.ideogram_fp8_weight_only)
+ self.assertFalse(config.arch_config.ideogram_bnb_4bit_weight_only)
+ self.assertFalse(config.arch_config.requires_gpu_resident_text_encoder)
+
+ def test_ideogram_text_encoder_post_config_hook_uses_bnb_for_nf4(self):
+ config = Ideogram4TextEncoderConfig()
+ config.update_model_arch(
+ {
+ "quantization_config": {
+ "quant_method": "bitsandbytes",
+ "load_in_4bit": True,
+ "bnb_4bit_quant_type": "nf4",
+ }
+ }
+ )
+
+ self.assertEqual(
+ config.arch_config.architectures, ["IdeogramQwen3VLTextEncoder"]
+ )
+ self.assertTrue(config.arch_config.ideogram_bnb_4bit_weight_only)
+ self.assertFalse(config.arch_config.ideogram_fp8_weight_only)
+ self.assertTrue(config.arch_config.requires_gpu_resident_text_encoder)
def test_ideogram_text_encoder_swaps_linears_to_weight_only_fp8(self):
config = Ideogram4TextEncoderConfig()
diff --git a/python/sglang/multimodal_gen/test/unit/test_transformer_quant.py b/python/sglang/multimodal_gen/test/unit/test_transformer_quant.py
index 6d0ac4b9f..6e00922e0 100644
--- a/python/sglang/multimodal_gen/test/unit/test_transformer_quant.py
+++ b/python/sglang/multimodal_gen/test/unit/test_transformer_quant.py
@@ -11,6 +11,7 @@ from types import SimpleNamespace
from unittest.mock import patch
import torch
+from safetensors.torch import save_file
partial_json_parser = types.ModuleType("partial_json_parser")
partial_json_parser_core = types.ModuleType("partial_json_parser.core")
@@ -58,6 +59,10 @@ from sglang.multimodal_gen.runtime.loader.transformer_load_utils import (
resolve_transformer_safetensors_to_load,
)
from sglang.multimodal_gen.runtime.models.dits.flux import FluxSingleTransformerBlock
+from sglang.multimodal_gen.runtime.utils.quantization_utils import (
+ build_nvfp4_config_from_safetensors_list,
+ get_quant_config,
+)
from sglang.multimodal_gen.tools.build_modelopt_nvfp4_transformer import (
_updated_quant_config,
)
@@ -260,6 +265,99 @@ class TestTransformerQuantHelpers(unittest.TestCase):
self.assertFalse(config.swap_weight_nibbles)
+ def test_bitsandbytes_quant_config_resolves_from_hf_config(self):
+ config = get_quant_config(
+ {
+ "quantization_config": {
+ "quant_method": "bitsandbytes",
+ "load_in_4bit": True,
+ "bnb_4bit_quant_type": "nf4",
+ "bnb_4bit_quant_storage": "uint8",
+ }
+ },
+ "/unused/component/path",
+ )
+
+ self.assertEqual(config.get_name(), "bitsandbytes")
+ self.assertTrue(config.load_in_4bit)
+ self.assertEqual(config.bnb_4bit_quant_type, "nf4")
+
+ def test_nvfp4_safetensors_inference_ignores_fp8_fallback_scales(self):
+ with tempfile.NamedTemporaryFile(suffix=".safetensors") as f:
+ save_file(
+ {
+ "fallback.weight": torch.empty(
+ (4, 4),
+ dtype=torch.float8_e4m3fn,
+ ),
+ "fallback.weight_scale": torch.tensor(1.0, dtype=torch.float32),
+ "layers.0.attention.qkv.weight": torch.zeros(
+ (32, 8),
+ dtype=torch.uint8,
+ ),
+ "layers.0.attention.qkv.weight_scale": torch.empty(
+ (32, 1),
+ dtype=torch.float8_e4m3fn,
+ ),
+ "layers.0.attention.qkv.weight_scale_2": torch.tensor(
+ 1.0,
+ dtype=torch.float32,
+ ),
+ },
+ f.name,
+ )
+
+ config = build_nvfp4_config_from_safetensors_list([f.name])
+
+ self.assertIsInstance(config, ModelOptFp4Config)
+ self.assertEqual(config.group_size, 16)
+ self.assertIn("fallback", config.exclude_modules)
+ self.assertNotIn("layers.0.attention.qkv", config.exclude_modules)
+ self.assertEqual(config.checkpoint_weight_scale_layout, "linear")
+ self.assertFalse(config.swap_weight_nibbles)
+
+ def test_nvfp4_safetensors_inference_uses_comfy_checkpoint_layout(self):
+ with tempfile.NamedTemporaryFile(suffix=".safetensors") as f:
+ save_file(
+ {
+ "fallback.weight": torch.empty(
+ (4, 4),
+ dtype=torch.float8_e4m3fn,
+ ),
+ "fallback.weight_scale": torch.tensor(1.0, dtype=torch.float32),
+ "fallback.comfy_quant": torch.tensor(
+ list(b'{"format":"float8_e4m3fn"}'),
+ dtype=torch.uint8,
+ ),
+ "layers.0.attention.qkv.weight": torch.zeros(
+ (32, 8),
+ dtype=torch.uint8,
+ ),
+ "layers.0.attention.qkv.weight_scale": torch.empty(
+ (32, 1),
+ dtype=torch.float8_e4m3fn,
+ ),
+ "layers.0.attention.qkv.weight_scale_2": torch.tensor(
+ 1.0,
+ dtype=torch.float32,
+ ),
+ "layers.0.attention.qkv.comfy_quant": torch.tensor(
+ list(b'{"format":"nvfp4"}'),
+ dtype=torch.uint8,
+ ),
+ },
+ f.name,
+ )
+
+ config = build_nvfp4_config_from_safetensors_list([f.name])
+
+ self.assertIsInstance(config, ModelOptFp4Config)
+ self.assertEqual(config.group_size, 16)
+ self.assertIn("fallback", config.exclude_modules)
+ self.assertNotIn("layers.0.attention.qkv", config.exclude_modules)
+ self.assertEqual(config.checkpoint_weight_scale_layout, "swizzled")
+ self.assertTrue(config.swap_weight_nibbles)
+
def test_builder_adds_diffusers_quant_type_for_nvfp4(self):
updated = _updated_quant_config(
{
diff --git a/python/sglang/utils.py b/python/sglang/utils.py
index 33bc175de..c7dc1ae78 100644
--- a/python/sglang/utils.py
+++ b/python/sglang/utils.py
@@ -34,6 +34,8 @@ logger = logging.getLogger(__name__)
KNOWN_NON_DIFFUSERS_DIFFUSION_MODEL_PATTERNS: dict[str, str] = {
"hunyuan3d": "Hunyuan3D2Pipeline",
"flux.2-dev-nvfp4": "Flux2NvfpPipeline",
+ "comfy-org/ideogram-4": "Ideogram4Nvfp4Pipeline",
+ "comfy-org--ideogram-4": "Ideogram4Nvfp4Pipeline",
}