[diffusion] feat: support loading self-describing quanto int8 encoders (#36052)

This commit is contained in:
Mick
2026-08-24 16:37:09 +08:00
committed by GitHub
parent 5081ad5d4e
commit adc09a1f63
13 changed files with 536 additions and 43 deletions
@@ -969,6 +969,19 @@ Install `comfy-kitchen>=0.2.27` and omit `--quantization`. SGLang automatically
loads its W4A8 language linears and tensorwise INT8 embedding; the unmarked
vision tower remains BF16.
The same component option accepts a self-describing Quanto qint8 file without
an additional quantization flag:
```bash Overlay
--component-paths.text_encoder \
DeepBeepMeep/MiniMax-H3/Qwen3-VL-32B-Instruct/Qwen3-VL-32B-Instruct-layer50_quanto_bf16_int8.safetensors
```
This variant keeps the declared language and vision linear weights in qint8
storage, then dequantizes only the active matrix for BF16/FP16 linear math. Use
it as a memory option, not as an INT8 throughput claim. The embedded Quanto map
is the selector; adding `--quantization` would describe a different operation.
</Tab>
<Tab title="Encoder scheduling">
@@ -200,6 +200,14 @@ backend.
<td><code>comfy-kitchen&gt;=0.2.27</code></td>
<td>Auto-detected; omit <code>--quantization</code>. Requires SM80+ and validates packed weights, group/channel scales, and optional codebooks before model construction. Mixed encoder files may keep their embedding tensorwise INT8. TP must preserve ConvRot group boundaries; offload is supported and FSDP is not.</td>
</tr>
<tr>
<td><code>quanto-int8</code></td>
<td>One native encoder safetensors file with an embedded Quanto quantization map</td>
<td>An explicit weight file through <code>--component-paths.&lt;component&gt;</code></td>
<td>Native encoders whose mapped linear layers consume every declared qint8 entry; MiniMax-H3's Qwen3-VL encoder is supported</td>
<td>None</td>
<td>Auto-detected weight-only qint8 storage. Each active matrix is dequantized to the compute dtype for the ordinary linear operation, so this reduces stored/resident weight memory rather than promising INT8 GEMM speed. TP and offload are supported; FSDP is not.</td>
</tr>
<tr>
<td><code>qvg-kv</code></td>
<td>Unquantized model with runtime causal KV-cache compression</td>
@@ -36,6 +36,7 @@ class QuantizationConfig(SRTQuantizationConfig):
# for quantization frameworks with a separate quantized model provided, e.g. Nunchaku
quantized_model_path: str | None = None
checkpoint_uses_native_qkv_layout: bool = False
supports_srt_linear_layers: bool = False
def get_scaled_act_names(self) -> list[str]:
return []
@@ -0,0 +1,194 @@
# SPDX-License-Identifier: Apache-2.0
"""Config and checkpoint admission for Optimum Quanto qint8 weights."""
from __future__ import annotations
import base64
import json
from collections.abc import Callable
from typing import Any
import torch
from safetensors import safe_open
from sglang.multimodal_gen.runtime.layers.linear import (
LinearBase as DiffusionLinearBase,
)
from sglang.multimodal_gen.runtime.layers.linear import (
UnquantizedLinearMethod as DiffusionUnquantizedLinearMethod,
)
from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config import (
QuantizationConfig,
QuantizeMethodBase,
)
from sglang.multimodal_gen.runtime.layers.quantization.quanto_int8 import (
QuantoInt8LinearMethod,
)
from sglang.srt.layers.linear import LinearBase as SrtLinearBase
from sglang.srt.layers.quantization.unquant import (
UnquantizedLinearMethod as SrtUnquantizedLinearMethod,
)
_FLOAT_DTYPES = {"BF16", "F16", "F32"}
class QuantoInt8Config(QuantizationConfig):
"""Dispatch linears declared qint8 in an Optimum Quanto quantization map."""
supports_srt_linear_layers = True
def __init__(self, layer_prefixes: set[str]) -> None:
super().__init__()
self.layer_prefixes = layer_prefixes
self.selected: set[str] = set()
@classmethod
def get_name(cls) -> str:
return "quanto_int8"
@classmethod
def get_supported_act_dtypes(cls) -> list[torch.dtype]:
return [torch.bfloat16, torch.float16]
@classmethod
def get_min_capability(cls) -> int:
return 0
@staticmethod
def get_config_filenames() -> list[str]:
return []
@classmethod
def from_config(cls, config: dict[str, Any]) -> QuantoInt8Config:
raise ValueError(
"QuantoInt8Config must be constructed from safetensors metadata"
)
def get_quant_method(
self, layer: torch.nn.Module, prefix: str
) -> QuantizeMethodBase | None:
if isinstance(layer, DiffusionLinearBase):
unquantized_method = DiffusionUnquantizedLinearMethod
elif isinstance(layer, SrtLinearBase):
unquantized_method = SrtUnquantizedLinearMethod
else:
return None
if prefix not in self.layer_prefixes:
return unquantized_method()
self.selected.add(prefix)
return QuantoInt8LinearMethod()
def inspect_quanto_int8_checkpoint(
file_path: str,
param_name_mapper: Callable[[str], str] | None = None,
) -> QuantoInt8Config | None:
"""Validate a self-describing Quanto qint8 safetensors checkpoint."""
with safe_open(file_path, framework="pt", device="cpu") as checkpoint:
metadata = checkpoint.metadata() or {}
if metadata.get("quantization_format") != "quanto":
return None
encoded_map = metadata.get("quantization_map_base64")
if encoded_map is None:
raise ValueError("Quanto checkpoint is missing quantization_map_base64")
try:
quantization_map = json.loads(
base64.b64decode(encoded_map, validate=True).decode("utf-8")
)
except (ValueError, UnicodeDecodeError, json.JSONDecodeError) as error:
raise ValueError("Invalid Quanto quantization_map_base64") from error
if not isinstance(quantization_map, dict) or not quantization_map:
raise ValueError("Quanto quantization map must be a non-empty object")
if not all(
isinstance(prefix, str) and isinstance(spec, dict)
for prefix, spec in quantization_map.items()
):
raise ValueError("Quanto quantization map entries must be named objects")
checkpoint_keys = set(checkpoint.keys())
data_suffix = ".weight._data"
data_prefixes = {
name.removesuffix(data_suffix)
for name in checkpoint_keys
if name.endswith(data_suffix)
}
map_prefixes = set(quantization_map)
if data_prefixes != map_prefixes:
missing_map = data_prefixes - map_prefixes
missing_data = map_prefixes - data_prefixes
raise ValueError(
"Quanto tensor/map prefixes do not match: "
f"missing metadata={sorted(missing_map)[:5]}, "
f"missing tensors={sorted(missing_data)[:5]}"
)
mapped_prefixes: set[str] = set()
for prefix, quantization in quantization_map.items():
if quantization.get("weights") != "qint8":
raise ValueError(
f"Unsupported Quanto weight type for {prefix!r}: "
f"{quantization.get('weights')!r}"
)
if quantization.get("activations") != "none":
raise ValueError(
f"Quanto activation quantization is not supported for {prefix!r}"
)
names = {
"data": f"{prefix}.weight._data",
"scale": f"{prefix}.weight._scale",
"input": f"{prefix}.input_scale",
"output": f"{prefix}.output_scale",
}
missing = set(names.values()) - checkpoint_keys
if missing:
raise ValueError(
f"Quanto layer {prefix!r} is missing tensors: {sorted(missing)}"
)
if f"{prefix}.weight" in checkpoint_keys:
raise ValueError(
f"Quanto layer {prefix!r} contains both packed and dense weights"
)
data_slice = checkpoint.get_slice(names["data"])
scale_slice = checkpoint.get_slice(names["scale"])
data_shape = tuple(data_slice.get_shape())
scale_shape = tuple(scale_slice.get_shape())
if data_slice.get_dtype() != "I8" or len(data_shape) != 2:
raise ValueError(
f"Quanto layer {prefix!r} needs a 2D I8 weight, got "
f"{data_slice.get_dtype()} {data_shape}"
)
if scale_slice.get_dtype() not in _FLOAT_DTYPES or scale_shape != (
data_shape[0],
1,
):
raise ValueError(
f"Quanto layer {prefix!r} has incompatible scale "
f"{scale_slice.get_dtype()} {scale_shape}"
)
for scale_name in (names["input"], names["output"]):
scale = checkpoint.get_slice(scale_name)
if (
scale.get_dtype() not in _FLOAT_DTYPES
or tuple(scale.get_shape()) != ()
):
raise ValueError(
f"Quanto auxiliary scale {scale_name!r} must be a float scalar"
)
mapped_prefix = (
param_name_mapper(prefix) if param_name_mapper is not None else prefix
)
if mapped_prefix in mapped_prefixes:
raise ValueError(
f"Quanto layers collide after parameter mapping at {mapped_prefix!r}"
)
mapped_prefixes.add(mapped_prefix)
return QuantoInt8Config(mapped_prefixes)
__all__ = ["QuantoInt8Config", "inspect_quanto_int8_checkpoint"]
@@ -0,0 +1,87 @@
# SPDX-License-Identifier: Apache-2.0
"""Runtime operations for serialized Optimum Quanto qint8 weights."""
from __future__ import annotations
from collections.abc import Iterable, Iterator
from typing import Any
import torch
import torch.nn.functional as F
from torch.nn.parameter import Parameter
from sglang.multimodal_gen.runtime.layers.linear import (
LinearMethodBase,
)
from sglang.multimodal_gen.runtime.utils.weight_attrs import set_weight_attrs
class QuantoInt8LinearMethod(LinearMethodBase):
"""Keep qint8 weights packed and dequantize only the active matrix."""
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: Any,
) -> None:
weight = Parameter(
torch.empty(
sum(output_partition_sizes),
input_size_per_partition,
dtype=torch.int8,
),
requires_grad=False,
)
set_weight_attrs(weight, {"input_dim": 1, "output_dim": 0})
set_weight_attrs(weight, extra_weight_attrs)
layer.register_parameter("weight", weight)
weight_scale = Parameter(
torch.empty(
sum(output_partition_sizes),
1,
dtype=params_dtype,
),
requires_grad=False,
)
set_weight_attrs(weight_scale, {"output_dim": 0})
set_weight_attrs(weight_scale, extra_weight_attrs)
layer.register_parameter("weight_scale", weight_scale)
def apply(
self,
layer: torch.nn.Module,
x: torch.Tensor,
bias: torch.Tensor | None = None,
) -> torch.Tensor:
weight = layer.weight.to(dtype=x.dtype)
weight.mul_(layer.weight_scale.to(dtype=x.dtype))
return F.linear(x, weight, bias)
def normalize_quanto_int8_weights(
weights: Iterable[tuple[str, torch.Tensor]],
) -> Iterator[tuple[str, torch.Tensor]]:
"""Translate flattened Quanto tensors to native linear parameter names."""
for name, tensor in weights:
if name.endswith((".input_scale", ".output_scale")):
if tensor.numel() != 1 or tensor.item() != 1:
raise ValueError(f"Quanto weight-only scale {name!r} must equal 1")
continue
if name.endswith(".weight._data"):
name = name.removesuffix("._data")
elif name.endswith(".weight._scale"):
name = name.removesuffix("._scale") + "_scale"
yield name, tensor
__all__ = [
"QuantoInt8LinearMethod",
"normalize_quanto_int8_weights",
]
@@ -37,6 +37,13 @@ from sglang.multimodal_gen.runtime.layers.quantization.configs.kitchen_int8_conf
from sglang.multimodal_gen.runtime.layers.quantization.configs.kitchen_w4a8_config import (
KitchenW4A8Config,
)
from sglang.multimodal_gen.runtime.layers.quantization.configs.quanto_int8_config import (
QuantoInt8Config,
inspect_quanto_int8_checkpoint,
)
from sglang.multimodal_gen.runtime.layers.quantization.quanto_int8 import (
normalize_quanto_int8_weights,
)
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
ComponentCheckpointUnsupportedError,
ComponentLoader,
@@ -85,6 +92,10 @@ from sglang.multimodal_gen.runtime.weights.source import (
)
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
from sglang.srt.environ import envs
from sglang.srt.layers.linear import LinearBase as SrtLinearBase
from sglang.srt.layers.quantization.unquant import (
UnquantizedLinearMethod as SrtUnquantizedLinearMethod,
)
logger = init_logger(__name__)
@@ -122,6 +133,23 @@ def _get_encoder_quant_config(
model_cls: type[nn.Module] | None = None,
):
quant_config = get_quant_config(component_config, component_model_path)
name_mapper = None
if model_cls is not None:
mapping = vars(model_cls).get("param_names_mapping", {})
if mapping:
mapping_fn = get_param_names_mapping(mapping)
def name_mapper(name: str) -> str:
# Layer-prefix metadata omits the suffix that many model
# mappings use to delimit a parameter name.
mapped_name, merge_index, _ = mapping_fn(f"{name}.weight")
if merge_index is not None:
raise ValueError(
"Serialized quantized component weights cannot use a "
"stacked parameter-name mapping"
)
return mapped_name.removesuffix(".weight")
if (
quant_config is None
and component_weights_path != component_model_path
@@ -131,26 +159,16 @@ def _get_encoder_quant_config(
component_weights_path
)
if quant_config is None and component_weights_path.endswith(".safetensors"):
name_mapper = None
if model_cls is not None:
mapping = vars(model_cls).get("param_names_mapping", {})
if mapping:
mapping_fn = get_param_names_mapping(mapping)
def name_mapper(name: str) -> str:
mapped_name, merge_index, _ = mapping_fn(f"{name}.weight")
if merge_index is not None:
raise ValueError(
"Comfy quantized component weights cannot use a "
"stacked parameter-name mapping"
)
return mapped_name.removesuffix(".weight")
markers = inspect_comfy_quant_markers(
[component_weights_path],
quant_config = inspect_quanto_int8_checkpoint(
component_weights_path,
param_name_mapper=name_mapper,
)
quant_config = resolve_comfy_checkpoint_quantization(markers)
if quant_config is None:
markers = inspect_comfy_quant_markers(
[component_weights_path],
param_name_mapper=name_mapper,
)
quant_config = resolve_comfy_checkpoint_quantization(markers)
return quant_config
@@ -298,10 +316,13 @@ def _process_quantized_encoder_weights(
) -> int:
processed_layers = 0
for module in model.modules():
if not isinstance(module, LinearBase):
if not isinstance(module, (LinearBase, SrtLinearBase)):
continue
quant_method = module.quant_method
if quant_method is None or isinstance(quant_method, UnquantizedLinearMethod):
if quant_method is None or isinstance(
quant_method,
(UnquantizedLinearMethod, SrtUnquantizedLinearMethod),
):
continue
origin_device = _module_tensor_device(module)
@@ -334,9 +355,12 @@ def _require_quantized_encoder_layers(
quant_config: QuantizationConfig | None = None,
) -> None:
has_quantized_layers = any(
isinstance(module, LinearBase)
isinstance(module, (LinearBase, SrtLinearBase))
and module.quant_method is not None
and not isinstance(module.quant_method, UnquantizedLinearMethod)
and not isinstance(
module.quant_method,
(UnquantizedLinearMethod, SrtUnquantizedLinearMethod),
)
for module in model.modules()
)
if not has_quantized_layers:
@@ -345,11 +369,20 @@ def _require_quantized_encoder_layers(
f"quantized linear layers for {component_name!r}"
)
if isinstance(quant_config, (ComfyFp8Config, KitchenInt8Config, KitchenW4A8Config)):
missing = set(quant_config.layer_markers) - set(quant_config.selected)
expected = set(quant_config.layer_markers)
selected = set(quant_config.selected)
elif isinstance(quant_config, QuantoInt8Config):
expected = quant_config.layer_prefixes
selected = quant_config.selected
else:
expected = set()
selected = set()
if expected:
missing = expected - selected
if missing:
raise ComponentCheckpointUnsupportedError(
f"The native {type(model).__name__} implementation did not consume "
f"Comfy quantization markers for {component_name!r}: "
f"serialized quantization markers for {component_name!r}: "
f"{sorted(missing)[:5]}"
)
@@ -818,17 +851,18 @@ class TextEncoderLoader(ComponentLoader):
model._keep_checkpoint_mapping = True
weights_to_load = {name for name, _ in model.named_parameters()}
loaded_weights = model.load_weights(
self._get_all_weights(
model,
model_path,
to_cpu=component_starts_on_cpu,
)
checkpoint_weights = self._get_all_weights(
model,
model_path,
to_cpu=component_starts_on_cpu,
)
if isinstance(quant_config, QuantoInt8Config):
checkpoint_weights = normalize_quanto_int8_weights(checkpoint_weights)
loaded_weights = model.load_weights(checkpoint_weights)
if quant_config is not None:
postprocess_device: torch.device | None = local_torch_device
if (
if isinstance(quant_config, QuantoInt8Config) or (
isinstance(quant_config, KitchenInt8Config)
and quant_config.is_checkpoint_int8_serialized
):
@@ -32,6 +32,7 @@ MINIMAX_H3_QWEN3VL_HIDDEN_DIM = 5120
_LAYER_WEIGHT_RE = re.compile(r"^model\.language_model\.layers\.(\d+)\.")
_PARAM_NAMES_MAPPING = {
r"^model\.(embed_tokens|layers|norm|rotary_emb)\.": r"model.language_model.\1.",
r"^language_model\.": r"model.language_model.",
r"^visual\.": r"model.visual.",
r"^(model\.visual\.blocks\.\d+\.attn\.)qkv\.": r"\1qkv_proj.",
}
@@ -694,7 +694,16 @@ class Qwen3VLModel(nn.Module):
prefix: str = "",
):
super().__init__()
self.visual = Qwen3VLVisionTransformer(config.vision_config)
vision_quant_config = (
quant_config
if quant_config is not None and quant_config.supports_srt_linear_layers
else None
)
self.visual = Qwen3VLVisionTransformer(
config.vision_config,
quant_config=vision_quant_config,
prefix=add_prefix("visual", prefix),
)
self.language_model = Qwen3VLTextModel(
config.text_config,
quant_config=quant_config,
@@ -45,22 +45,30 @@ class Qwen3VLVisionRotaryEmbedding(nn.Module):
class Qwen3VLVisionBlock(nn.Module):
def __init__(self, config: Any, layer_idx: int) -> None:
def __init__(
self,
config: Any,
layer_idx: int,
quant_config: Any = None,
prefix: str = "visual",
) -> None:
super().__init__()
parallel = get_parallel()
self.norm1 = nn.LayerNorm(config.hidden_size, eps=1e-6)
self.norm2 = nn.LayerNorm(config.hidden_size, eps=1e-6)
self.attn = QwenVLVisionAttention(
config,
prefix=f"visual.blocks.{layer_idx}.attn",
prefix=f"{prefix}.blocks.{layer_idx}.attn",
model_name="Qwen3-VL",
quant_config=quant_config,
)
self.mlp = Qwen3_VisionMLP(
config.hidden_size,
config.intermediate_size,
bias=True,
hidden_act=config.hidden_act,
prefix=f"visual.blocks.{layer_idx}.mlp",
prefix=f"{prefix}.blocks.{layer_idx}.mlp",
quant_config=quant_config,
tp_rank=parallel.tp_rank,
tp_size=parallel.tp_size,
)
@@ -176,7 +184,12 @@ def _vision_cu_seqlens(grid_thw: torch.Tensor) -> torch.Tensor:
class Qwen3VLVisionTransformer(nn.Module):
def __init__(self, config: Any) -> None:
def __init__(
self,
config: Any,
quant_config: Any = None,
prefix: str = "visual",
) -> None:
super().__init__()
parallel = get_parallel()
self.config = config
@@ -191,7 +204,8 @@ class Qwen3VLVisionTransformer(nn.Module):
head_dim = config.hidden_size // config.num_heads
self.rotary_pos_emb = Qwen3VLVisionRotaryEmbedding(head_dim // 2)
self.blocks = nn.ModuleList(
Qwen3VLVisionBlock(config, layer_idx) for layer_idx in range(config.depth)
Qwen3VLVisionBlock(config, layer_idx, quant_config, prefix)
for layer_idx in range(config.depth)
)
self.merger = Qwen3VLMoeVisionPatchMerger(
dim=config.out_hidden_size,
@@ -59,7 +59,14 @@ def _apply_rotary_embedding(
class QwenVLVisionAttention(nn.Module):
def __init__(self, config: Any, *, prefix: str, model_name: str) -> None:
def __init__(
self,
config: Any,
*,
prefix: str,
model_name: str,
quant_config: Any = None,
) -> None:
super().__init__()
parallel = get_parallel()
self.num_heads = config.num_heads // parallel.tp_size
@@ -70,6 +77,7 @@ class QwenVLVisionAttention(nn.Module):
head_size=self.head_dim,
total_num_heads=config.num_heads,
bias=True,
quant_config=quant_config,
prefix=f"{prefix}.qkv_proj",
tp_rank=parallel.tp_rank,
tp_size=parallel.tp_size,
@@ -78,6 +86,7 @@ class QwenVLVisionAttention(nn.Module):
input_size=config.hidden_size,
output_size=config.hidden_size,
bias=True,
quant_config=quant_config,
prefix=f"{prefix}.proj",
tp_rank=parallel.tp_rank,
tp_size=parallel.tp_size,
@@ -0,0 +1,85 @@
import base64
import json
import pytest
import torch
import torch.nn.functional as F
from safetensors.torch import load_file, save_file
from sglang.multimodal_gen.runtime.layers.linear import ReplicatedLinear
from sglang.multimodal_gen.runtime.layers.quantization.configs.quanto_int8_config import (
QuantoInt8Config,
inspect_quanto_int8_checkpoint,
)
from sglang.multimodal_gen.runtime.layers.quantization.quanto_int8 import (
normalize_quanto_int8_weights,
)
def _save_quanto_checkpoint(path, *, activations="none"):
prefix = "language_model.layers.0.mlp.up_proj"
quantization_map = {prefix: {"weights": "qint8", "activations": activations}}
save_file(
{
f"{prefix}.weight._data": torch.tensor(
[[1, -2], [3, 4], [-5, 6]], dtype=torch.int8
),
f"{prefix}.weight._scale": torch.tensor(
[[0.5], [0.25], [0.125]], dtype=torch.bfloat16
),
f"{prefix}.input_scale": torch.tensor(1, dtype=torch.bfloat16),
f"{prefix}.output_scale": torch.tensor(1, dtype=torch.bfloat16),
},
path,
metadata={
"quantization_format": "quanto",
"quantization_map_base64": base64.b64encode(
json.dumps(quantization_map).encode()
).decode(),
},
)
def test_quanto_checkpoint_drives_native_linear_end_to_end(tmp_path):
checkpoint = tmp_path / "encoder.safetensors"
_save_quanto_checkpoint(checkpoint)
config = inspect_quanto_int8_checkpoint(
str(checkpoint), param_name_mapper=lambda name: f"model.{name}"
)
assert isinstance(config, QuantoInt8Config)
prefix = "model.language_model.layers.0.mlp.up_proj"
layer = ReplicatedLinear(
2,
3,
bias=False,
params_dtype=torch.bfloat16,
quant_config=config,
prefix=prefix,
)
tensors = dict(normalize_quanto_int8_weights(load_file(str(checkpoint)).items()))
raw_prefix = prefix.removeprefix("model.")
for suffix, parameter in (
("weight", layer.weight),
("weight_scale", layer.weight_scale),
):
parameter.weight_loader(parameter, tensors.pop(f"{raw_prefix}.{suffix}"))
x = torch.tensor([[2.0, -1.0]], dtype=torch.bfloat16)
expected_weight = layer.weight.to(torch.bfloat16) * layer.weight_scale
output, _ = layer(x)
torch.testing.assert_close(output, F.linear(x, expected_weight))
assert not tensors
assert config.selected == {prefix}
def test_quanto_checkpoint_rejects_activation_quantization(tmp_path):
checkpoint = tmp_path / "encoder.safetensors"
_save_quanto_checkpoint(checkpoint, activations="qint8")
with pytest.raises(ValueError, match="activation quantization"):
inspect_quanto_int8_checkpoint(str(checkpoint))
def test_quanto_weight_only_auxiliary_scales_must_be_identity():
with pytest.raises(ValueError, match="must equal 1"):
list(normalize_quanto_int8_weights([("layer.input_scale", torch.tensor(0.5))]))
@@ -4,6 +4,9 @@ import torch
from torch import nn
from sglang.multimodal_gen.configs.models.encoders.qwen3vl import Qwen3VLArchConfig
from sglang.multimodal_gen.runtime.layers.quantization.configs.quanto_int8_config import (
QuantoInt8Config,
)
from sglang.multimodal_gen.runtime.models.encoders.minimax_h3_qwen3vl import (
MiniMaxH3Qwen3VLEncoder,
)
@@ -86,6 +89,41 @@ def test_native_vision_keeps_checkpoint_parameter_names():
}
def test_native_vision_accepts_srt_linear_quantization():
config = SimpleNamespace(
hidden_size=16,
intermediate_size=24,
hidden_act="gelu_pytorch_tanh",
num_heads=2,
depth=1,
patch_size=2,
temporal_patch_size=1,
in_channels=3,
num_position_embeddings=16,
spatial_merge_size=2,
out_hidden_size=12,
deepstack_visual_indexes=[],
)
prefixes = {
"model.visual.blocks.0.attn.qkv_proj",
"model.visual.blocks.0.attn.proj",
"model.visual.blocks.0.mlp.linear_fc1",
"model.visual.blocks.0.mlp.linear_fc2",
}
quant_config = QuantoInt8Config(prefixes)
with get_parallel().override(tp_size=1, tp_rank=0):
model = Qwen3VLVisionTransformer(
config,
quant_config=quant_config,
prefix="model.visual",
)
assert quant_config.selected == prefixes
for name, parameter in model.blocks[0].named_parameters():
if name.endswith("weight") and not name.startswith("norm"):
assert parameter.dtype == torch.int8
def test_native_vision_keeps_position_math_in_fp32():
class PatchEmbed(nn.Module):
def __init__(self):
@@ -186,7 +186,7 @@ class TestMiniMaxH3CheckpointFilter(unittest.TestCase):
"model.layers.49.self_attn.q_proj.weight": True,
"model.layers.50.self_attn.q_proj.weight": False,
"visual.blocks.0.attn.qkv.weight": True,
"language_model.layers.63.mlp.down_proj.weight": True,
"language_model.layers.63.mlp.down_proj.weight": False,
"module.model.language_model.layers.63.mlp.down_proj.weight": True,
}
self.assertEqual(
@@ -383,11 +383,11 @@ class TestTextEncoderQuantization(unittest.TestCase):
with tempfile.NamedTemporaryFile(suffix=".safetensors") as checkpoint:
save_file(
{
"model.layers.0.self_attn.q_proj.weight": torch.ones(
"visual.blocks.0.attn.qkv.weight": torch.ones(
(2, 256), dtype=torch.int8
),
"model.layers.0.self_attn.q_proj.weight_scale": torch.ones((2, 1)),
"model.layers.0.self_attn.q_proj.comfy_quant": torch.tensor(
"visual.blocks.0.attn.qkv.weight_scale": torch.ones((2, 1)),
"visual.blocks.0.attn.qkv.comfy_quant": torch.tensor(
list(marker), dtype=torch.uint8
),
},
@@ -411,7 +411,7 @@ class TestTextEncoderQuantization(unittest.TestCase):
self.assertIsInstance(model_config.quant_config, KitchenInt8Config)
self.assertEqual(
set(model_config.quant_config.layer_markers),
{"model.language_model.layers.0.self_attn.q_proj"},
{"model.visual.blocks.0.attn.qkv_proj"},
)
def test_mixed_w4a8_weight_file_maps_embedding_and_linear_markers(self):