[diffusion] feat: support loading minimax h3 gguf text encoders (#36055)

This commit is contained in:
Mick
2026-08-24 23:00:49 +08:00
committed by GitHub
parent 9b0007ed19
commit 30f9ed09d1
10 changed files with 367 additions and 44 deletions
@@ -1006,6 +1006,21 @@ 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.
Qwen3-VL GGUF files use the same component option:
```bash Overlay
--component-paths.text_encoder \
DeepBeepMeep/MiniMax-H3/Qwen3-VL-32B-Instruct/qwen3vl-32B-MiniMax-H3-Q4_K_M.gguf
```
The file selects its own Q4/Q6 layout; do not add `--quantization`. SGLang keeps
the aligned vocabulary table and language linears in packed GGML storage. The
smaller Comfy vision tower, whose flattened blocks can cross row and TP shard
boundaries, is restored one matrix at a time to BF16 during loading.
This CUDA-only memory option supports encoder TP and layerwise placement, but
not encoder FSDP. Like the Quanto variant, it is approximate and is not a
throughput claim.
</Tab>
<Tab title="Encoder scheduling">
+8 -8
View File
@@ -162,11 +162,11 @@ backend.
</tr>
<tr>
<td><code>gguf</code></td>
<td>One selected GGUF DiT file</td>
<td><code>--transformer-weights-path</code></td>
<td>MiniMax-H3 original or pruned FL2VA / Ref2VA DiTs</td>
<td>One selected GGUF DiT or native encoder file</td>
<td><code>--transformer-weights-path</code>, or an explicit encoder file through <code>--component-paths.&lt;component&gt;</code></td>
<td>MiniMax-H3 original or pruned FL2VA / Ref2VA DiTs and its Qwen3-VL text encoder</td>
<td>None</td>
<td>CUDA only; auto-detected; supports standard and K-quant GGML types; FSDP and the separate Qwen3-VL text-encoder GGUF files are not supported</td>
<td>CUDA only and auto-detected; supports standard and K-quant GGML types. TP requires aligned GGML blocks; component/layerwise offload is supported and FSDP is not.</td>
</tr>
<tr>
<td><code>comfy-fp8</code></td>
@@ -242,11 +242,11 @@ backend.
</tr>
<tr>
<td><code>gguf</code></td>
<td>A single community <code>.gguf</code> holding the transformer</td>
<td><code>--transformer-weights-path</code></td>
<td>MiniMax-H3 <code>fl2va</code> (original and pruned AdaLN curve)</td>
<td>A single community <code>.gguf</code> holding a transformer or native encoder</td>
<td><code>--transformer-weights-path</code> or <code>--component-paths.&lt;component&gt;</code></td>
<td>MiniMax-H3 <code>fl2va</code> DiT (original and pruned AdaLN curve) and Qwen3-VL text encoder</td>
<td>None</td>
<td>CUDA only, no FSDP. TP requires GGML-aligned shard boundaries. Shrinks the download and the host memory offload pins (17.5 vs 61.7 GiB for H3) rather than peak VRAM, which offload already bounds. Dequantized per use, so it is not faster. See <a href="#gguf">GGUF</a>.</td>
<td>CUDA only, no FSDP. TP requires GGML-aligned shard boundaries. Packed weights reduce checkpoint and resident-weight memory; they are dequantized per use, so this is not a speed claim. See <a href="#gguf">GGUF</a>.</td>
</tr>
</tbody>
</table>
@@ -38,6 +38,7 @@ class QuantizationConfig(SRTQuantizationConfig):
checkpoint_uses_native_qkv_layout: bool = False
checkpoint_uses_comfy_quantization: bool = False
supports_srt_linear_layers: bool = False
supports_quantized_embeddings: bool = False
def get_scaled_act_names(self) -> list[str]:
return []
@@ -3,6 +3,7 @@
from __future__ import annotations
from collections.abc import Callable
from typing import Any
import gguf
@@ -18,11 +19,15 @@ from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config impor
QuantizationConfig,
QuantizeMethodBase,
)
from sglang.multimodal_gen.runtime.layers.vocab_parallel_embedding import (
VocabParallelEmbedding,
)
from sglang.multimodal_gen.runtime.loader.gguf_weights import GGUFTensorMeta
from sglang.multimodal_gen.runtime.utils.weight_attrs import set_weight_attrs
from sglang.srt.layers.quantization.gguf import (
DEQUANT_TYPES,
UNQUANTIZED_TYPES,
apply_gguf_embedding,
dequantize_gguf_weight,
)
@@ -30,10 +35,29 @@ from sglang.srt.layers.quantization.gguf import (
class GGUFConfig(QuantizationConfig):
"""Select a GGUF method from each checkpoint tensor's metadata."""
supports_quantized_embeddings = True
def __init__(self, gguf_file: str, tensor_meta: dict[str, GGUFTensorMeta]):
super().__init__()
self.gguf_file = gguf_file
self.tensor_meta = tensor_meta
self._refresh_quantized_prefixes()
self.selected: set[str] = set()
def retain_tensor_meta(self, key_filter: Callable[[str], bool]) -> None:
self.tensor_meta = {
name: metadata
for name, metadata in self.tensor_meta.items()
if key_filter(name)
}
self._refresh_quantized_prefixes()
def _refresh_quantized_prefixes(self) -> None:
self.quantized_prefixes = {
metadata.param_name.removesuffix(".qweight")
for metadata in self.tensor_meta.values()
if metadata.is_packed
}
@classmethod
def get_name(cls) -> str:
@@ -58,7 +82,11 @@ class GGUFConfig(QuantizationConfig):
def get_quant_method(
self, layer: nn.Module, prefix: str
) -> QuantizeMethodBase | None:
if not isinstance(layer, LinearBase):
if isinstance(layer, LinearBase):
unquantized_method = UnquantizedLinearMethod
elif isinstance(layer, VocabParallelEmbedding):
unquantized_method = None
else:
return None
metadata = self.tensor_meta.get(f"{prefix}.weight")
@@ -68,14 +96,28 @@ class GGUFConfig(QuantizationConfig):
f"{self.gguf_file!r}"
)
weight_type = metadata.weight_type
if weight_type in UNQUANTIZED_TYPES:
return UnquantizedLinearMethod()
if not metadata.is_packed or weight_type in UNQUANTIZED_TYPES:
if unquantized_method is None:
return None
return unquantized_method()
if weight_type not in DEQUANT_TYPES:
raise ValueError(
f"GGUF tensor {prefix}.weight uses unsupported type {weight_type}"
)
self.selected.add(prefix)
if isinstance(layer, VocabParallelEmbedding):
return GGUFEmbeddingMethod(metadata, prefix)
return GGUFLinearMethod(metadata, prefix)
def supports_input_partition(
self, prefix: str, input_size_per_partition: int
) -> bool:
metadata = self.tensor_meta.get(f"{prefix}.weight")
if metadata is None or not metadata.is_packed:
return True
block_size, _ = gguf.GGML_QUANT_SIZES[metadata.weight_type]
return input_size_per_partition % block_size == 0
class GGUFLinearMethod(LinearMethodBase):
"""Register TP-local packed weights and reuse SRT dequantization."""
@@ -95,6 +137,7 @@ class GGUFLinearMethod(LinearMethodBase):
params_dtype: torch.dtype,
**extra_weight_attrs: Any,
) -> None:
self.params_dtype = params_dtype
if self.metadata.logical_shape != (output_size, input_size):
raise ValueError(
f"GGUF tensor {self.prefix}.weight has logical shape "
@@ -130,4 +173,17 @@ class GGUFLinearMethod(LinearMethodBase):
return nn.functional.linear(x, weight, bias)
__all__ = ["GGUFConfig", "GGUFLinearMethod"]
class GGUFEmbeddingMethod(GGUFLinearMethod):
"""Use SRT's packed GGUF lookup for a diffusion vocabulary table."""
def embedding(self, layer: nn.Module, tokens: torch.Tensor) -> torch.Tensor:
return apply_gguf_embedding(
tokens,
layer.qweight,
self.weight_type,
self.metadata.logical_shape[1],
dtype=self.params_dtype,
)
__all__ = ["GGUFConfig", "GGUFEmbeddingMethod", "GGUFLinearMethod"]
@@ -43,6 +43,7 @@ from sglang.multimodal_gen.runtime.layers.quantization.configs.quanto_int8_confi
QuantoInt8Config,
inspect_quanto_int8_checkpoint,
)
from sglang.multimodal_gen.runtime.layers.quantization.gguf import GGUFConfig
from sglang.multimodal_gen.runtime.layers.quantization.quanto_int8 import (
normalize_quanto_int8_weights,
)
@@ -52,6 +53,12 @@ from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader imp
NativeComponentLoaderRequired,
uses_native_transformers_bnb4,
)
from sglang.multimodal_gen.runtime.loader.gguf_weights import (
gguf_weights_iterator,
names_gguf_checkpoint,
read_gguf_tensor_meta,
remap_gguf_tensor_meta,
)
from sglang.multimodal_gen.runtime.loader.utils import (
get_param_names_mapping,
set_default_torch_dtype,
@@ -171,22 +178,46 @@ def _get_encoder_quant_config(
quant_config = get_quant_config(component_config, component_model_path)
name_mapper = None
parameter_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")
def parameter_name_mapper(name: str) -> str:
mapped_name, merge_index, _ = mapping_fn(name)
if merge_index is not None:
raise ValueError(
"Serialized quantized component weights cannot use a "
"stacked parameter-name mapping"
)
return mapped_name
def name_mapper(name: str) -> str:
# Layer-prefix metadata omits the suffix that many model
# mappings use to delimit a parameter name.
mapped_name = parameter_name_mapper(f"{name}.weight")
return mapped_name.removesuffix(".weight")
if names_gguf_checkpoint(component_weights_path):
if quant_config is not None:
raise ValueError(
"A GGUF encoder checkpoint cannot be combined with a second "
"quantization declaration"
)
tensor_meta = read_gguf_tensor_meta(component_weights_path)
dequantize_prefixes = (
vars(model_cls).get("gguf_dequantize_prefixes", ())
if model_cls is not None
else ()
)
tensor_meta = remap_gguf_tensor_meta(
tensor_meta,
parameter_name_mapper or (lambda name: name),
dequantize_prefixes=dequantize_prefixes,
)
return GGUFConfig(component_weights_path, tensor_meta)
if (
quant_config is None
and component_weights_path != component_model_path
@@ -381,6 +412,9 @@ def _require_quantized_encoder_layers(
elif isinstance(quant_config, QuantoInt8Config):
expected = quant_config.layer_prefixes
selected = quant_config.selected
elif isinstance(quant_config, GGUFConfig):
expected = quant_config.quantized_prefixes
selected = quant_config.selected
else:
expected = set()
selected = set()
@@ -448,6 +482,17 @@ class TextEncoderLoader(ComponentLoader):
weights_override = server_args.component_weights_paths.get(component_name)
if weights_override is None:
return component_model_path
if names_gguf_checkpoint(weights_override):
if not current_platform.is_cuda():
raise ValueError(
"GGUF encoder checkpoints require CUDA; the GGML kernels have "
f"no {current_platform.device_type} implementation"
)
if server_args.should_use_fsdp_for_component(component_name):
raise ValueError(
f"GGUF encoder checkpoint {component_name!r} is incompatible "
"with FSDP; select resident or layerwise placement"
)
model_weights_path = materialize_weight(resolve_weight(weights_override))
logger.info(
"Using weight-file override for %s: %s",
@@ -601,19 +646,14 @@ class TextEncoderLoader(ComponentLoader):
def _get_all_weights(
self,
model: nn.Module,
model: EncoderTensorParallelMixin,
model_path: str,
to_cpu: bool,
) -> Generator[tuple[str, torch.Tensor], None, None]:
key_filter = cast(
Callable[[str], bool] | None,
getattr(model, "should_materialize_checkpoint_weight", None),
)
key_filter = model.should_materialize_checkpoint_weight
def include_checkpoint_weight(name: str) -> bool:
return not name.endswith(".comfy_quant") and (
key_filter is None or key_filter(name)
)
return not name.endswith(".comfy_quant") and key_filter(name)
primary_weights = TextEncoderLoader.Source(
model_path,
@@ -840,6 +880,10 @@ class TextEncoderLoader(ComponentLoader):
)
model.bind_encoder_tp_group(encoder_tp_group)
if isinstance(quant_config, GGUFConfig):
quant_config.retain_tensor_meta(
model.should_materialize_checkpoint_weight
)
if quant_config is not None:
_require_quantized_encoder_layers(
model, component_name, quant_config=quant_config
@@ -858,16 +902,23 @@ class TextEncoderLoader(ComponentLoader):
model._keep_checkpoint_mapping = True
weights_to_load = {name for name, _ in model.named_parameters()}
checkpoint_weights = self._get_all_weights(
model,
model_path,
to_cpu=component_starts_on_cpu,
)
if isinstance(quant_config, GGUFConfig):
checkpoint_weights = gguf_weights_iterator(
model_path,
quant_config.tensor_meta,
key_filter=model.should_materialize_checkpoint_weight,
)
else:
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:
if quant_config is not None and not isinstance(quant_config, GGUFConfig):
postprocess_device: torch.device | None = local_torch_device
if isinstance(quant_config, QuantoInt8Config) or (
isinstance(quant_config, KitchenInt8Config)
@@ -7,7 +7,7 @@ import math
import os
import warnings
from collections.abc import Callable, Generator
from dataclasses import dataclass
from dataclasses import dataclass, replace
from typing import TYPE_CHECKING, Any
import torch
@@ -34,6 +34,7 @@ class GGUFTensorMeta:
stored_shape: tuple[int, ...]
stored_dtype: torch.dtype
param_name: str
dequantize_on_load: bool = False
@property
def weight_type(self) -> WeightType:
@@ -45,6 +46,10 @@ class GGUFTensorMeta:
def is_quantized(self) -> bool:
return self.ggml_type not in _UNQUANTIZED_TYPES
@property
def is_packed(self) -> bool:
return self.is_quantized and not self.dequantize_on_load
def _gguf_module() -> Any:
try:
@@ -81,8 +86,20 @@ def read_gguf_tensor_meta(gguf_file: str) -> dict[str, GGUFTensorMeta]:
metadata: dict[str, GGUFTensorMeta] = {}
for tensor in reader.tensors:
weight_type = WeightType(tensor.tensor_type)
logical_shape = tuple(int(dim) for dim in reversed(tensor.shape))
shape_field = reader.fields.get(f"comfy.gguf.orig_shape.{tensor.name}")
logical_shape = (
tuple(int(dim) for dim in shape_field.contents())
if shape_field is not None
else tuple(int(dim) for dim in reversed(tensor.shape))
)
if math.prod(logical_shape) != tensor.n_elements:
raise ValueError(
f"GGUF tensor {tensor.name} declares original shape "
f"{logical_shape}, which contains {math.prod(logical_shape)} "
f"elements instead of {tensor.n_elements}"
)
is_quantized = int(weight_type) not in _UNQUANTIZED_TYPES
dequantize_on_load = False
if is_quantized:
if len(logical_shape) != 2 or not tensor.name.endswith(".weight"):
raise ValueError(
@@ -93,14 +110,18 @@ def read_gguf_tensor_meta(gguf_file: str) -> dict[str, GGUFTensorMeta]:
block_size, type_size = gguf.GGML_QUANT_SIZES[weight_type]
inner_dim = logical_shape[-1]
if inner_dim % block_size:
raise ValueError(
f"GGUF tensor {tensor.name} has inner dimension {inner_dim}, "
f"which is not a multiple of block size {block_size}"
if shape_field is None:
raise ValueError(
f"GGUF tensor {tensor.name} has inner dimension {inner_dim}, "
f"which is not a multiple of block size {block_size}"
)
dequantize_on_load = True
stored_shape = logical_shape
else:
stored_shape = (
*logical_shape[:-1],
inner_dim // block_size * type_size,
)
stored_shape = (
*logical_shape[:-1],
inner_dim // block_size * type_size,
)
if (
int(weight_type) in _SUPER_BLOCK_DEQUANT_TYPES
and math.prod(logical_shape) % _GGML_SUPER_BLOCK
@@ -109,7 +130,7 @@ def read_gguf_tensor_meta(gguf_file: str) -> dict[str, GGUFTensorMeta]:
f"GGUF tensor {tensor.name} is not aligned to "
f"{_GGML_SUPER_BLOCK}-element super blocks"
)
stored_dtype = torch.uint8
stored_dtype = torch.bfloat16 if dequantize_on_load else torch.uint8
else:
stored_shape = logical_shape
stored_dtype = {
@@ -120,7 +141,7 @@ def read_gguf_tensor_meta(gguf_file: str) -> dict[str, GGUFTensorMeta]:
param_name = (
f"{tensor.name.removesuffix('.weight')}.qweight"
if is_quantized
if is_quantized and not dequantize_on_load
else tensor.name
)
metadata[tensor.name] = GGUFTensorMeta(
@@ -129,11 +150,51 @@ def read_gguf_tensor_meta(gguf_file: str) -> dict[str, GGUFTensorMeta]:
stored_shape=stored_shape,
stored_dtype=stored_dtype,
param_name=param_name,
dequantize_on_load=dequantize_on_load,
)
return metadata
def remap_gguf_tensor_meta(
tensor_meta: dict[str, GGUFTensorMeta],
name_mapper: Callable[[str], str],
dequantize_prefixes: tuple[str, ...] = (),
) -> dict[str, GGUFTensorMeta]:
"""Map checkpoint tensor names while retaining raw lookup aliases."""
remapped: dict[str, GGUFTensorMeta] = {}
for checkpoint_name, metadata in tensor_meta.items():
if metadata.is_quantized and checkpoint_name.startswith(dequantize_prefixes):
metadata = replace(
metadata,
stored_shape=metadata.logical_shape,
stored_dtype=torch.bfloat16,
param_name=checkpoint_name,
dequantize_on_load=True,
)
parameter_name = name_mapper(checkpoint_name)
mapped_param_name = (
f"{parameter_name.removesuffix('.weight')}.qweight"
if metadata.is_packed
else parameter_name
)
mapped_metadata = replace(metadata, param_name=mapped_param_name)
for alias in (checkpoint_name, parameter_name):
previous = remapped.get(alias)
if previous is not None and previous != mapped_metadata:
raise ValueError(
f"GGUF tensors collide after parameter mapping at {alias!r}"
)
remapped[alias] = mapped_metadata
return remapped
def _tensor_to_torch(tensor, metadata: GGUFTensorMeta) -> torch.Tensor:
if metadata.dequantize_on_load:
gguf = _gguf_module()
value = gguf.dequantize(tensor.data, metadata.weight_type)
return torch.from_numpy(value.reshape(metadata.logical_shape)).to(
metadata.stored_dtype
)
with warnings.catch_warnings():
warnings.filterwarnings(
"ignore",
@@ -144,7 +205,7 @@ def _tensor_to_torch(tensor, metadata: GGUFTensorMeta) -> torch.Tensor:
if metadata.ggml_type == _GGML_BF16:
return value.view(torch.bfloat16).reshape(metadata.stored_shape).clone()
value = value.reshape(metadata.stored_shape)
return value.clone() if not metadata.is_quantized else value
return value.clone() if not metadata.is_packed else value
def gguf_weights_iterator(
@@ -182,4 +243,5 @@ __all__ = [
"gguf_weights_iterator",
"names_gguf_checkpoint",
"read_gguf_tensor_meta",
"remap_gguf_tensor_meta",
]
@@ -164,6 +164,10 @@ class EncoderTensorParallelMixin:
# states or sharding contract cannot use the generic loader lifecycle.
manages_checkpoint_quantization = False
@staticmethod
def should_materialize_checkpoint_weight(name: str) -> bool:
return True
@classmethod
def configure_component_paths(
cls,
@@ -208,6 +208,9 @@ class MiniMaxH3Qwen3VLEncoder(TextEncoder):
supports_dp_encode = True
param_names_mapping = _PARAM_NAMES_MAPPING
# Comfy packs the vision tower across whole tensors rather than rows. Keep
# its language/vocabulary matrices packed and restore this smaller tower.
gguf_dequantize_prefixes = ("visual.", "model.visual.")
@classmethod
def configure_component_paths(
@@ -33,6 +33,7 @@ from sglang.multimodal_gen.runtime.loader.gguf_weights import (
gguf_weights_iterator,
names_gguf_checkpoint,
read_gguf_tensor_meta,
remap_gguf_tensor_meta,
)
from sglang.srt.layers.quantization.gguf import UNQUANTIZED_TYPES
from sglang.srt.utils.hf_transformers import check_gguf_file
@@ -51,10 +52,18 @@ def _kv_string(key: str, value: str, bo: str = "<") -> bytes:
return out
def _kv_u64_array(key: str, values: list[int], bo: str = "<") -> bytes:
out = struct.pack(f"{bo}Q", len(key)) + key.encode()
out += struct.pack(f"{bo}IIQ", 9, 10, len(values))
out += b"".join(struct.pack(f"{bo}Q", value) for value in values)
return out
def _write_gguf(
path: Path,
tensors: list[tuple[str, list[int], int, bytes]],
byte_order: str = "<",
metadata: tuple[bytes, ...] = (),
) -> None:
"""Write a minimal GGUF v3 file containing ``tensors``.
@@ -65,8 +74,9 @@ def _write_gguf(
"""
bo = byte_order
header = b"GGUF" + struct.pack(f"{bo}I", 3)
header += struct.pack(f"{bo}QQ", len(tensors), 1)
header += struct.pack(f"{bo}QQ", len(tensors), 1 + len(metadata))
header += _kv_string("general.architecture", "test", bo)
header += b"".join(metadata)
# Tensor info blocks, then padded data.
infos = b""
@@ -117,6 +127,63 @@ class TestGGUFTensorMeta(unittest.TestCase):
# The layer registers `qweight`, so that is what the iterator must yield.
self.assertEqual(meta.param_name, "w.qweight")
def test_comfy_original_shape_restores_matrix_rows(self):
path = self.tmp / "comfy.gguf"
logical_shape = [4, 512]
payload = bytes(4 * 512 // _Q4_K_BLOCK * _Q4_K_TYPE_SIZE)
_write_gguf(
path,
[("w.weight", [256, 8], _Q4_K, payload)],
metadata=(_kv_u64_array("comfy.gguf.orig_shape.w.weight", logical_shape),),
)
meta = read_gguf_tensor_meta(str(path))["w.weight"]
self.assertEqual(meta.logical_shape, (4, 512))
self.assertEqual(meta.stored_shape, (4, 288))
def test_comfy_non_aligned_rows_dequantize_during_load(self):
path = self.tmp / "comfy-unaligned.gguf"
logical_shape = [2, 384]
payload = bytes(2 * 384 // _Q4_K_BLOCK * _Q4_K_TYPE_SIZE)
_write_gguf(
path,
[("vision.weight", [256, 3], _Q4_K, payload)],
metadata=(
_kv_u64_array("comfy.gguf.orig_shape.vision.weight", logical_shape),
),
)
metadata = read_gguf_tensor_meta(str(path))
loaded = dict(gguf_weights_iterator(str(path), metadata))["vision.weight"]
self.assertTrue(metadata["vision.weight"].dequantize_on_load)
self.assertEqual(loaded.dtype, torch.bfloat16)
self.assertEqual(tuple(loaded.shape), (2, 384))
def test_parameter_mapping_retains_checkpoint_lookup(self):
meta = GGUFTensorMeta(
ggml_type=int(_Q4_K),
logical_shape=(4, 512),
stored_shape=(4, 288),
stored_dtype=torch.uint8,
param_name="visual.block.qkv.qweight",
)
remapped = remap_gguf_tensor_meta(
{"visual.block.qkv.weight": meta},
lambda name: "model.visual.block.qkv_proj.weight",
)
self.assertIs(
remapped["visual.block.qkv.weight"],
remapped["model.visual.block.qkv_proj.weight"],
)
self.assertEqual(
remapped["visual.block.qkv.weight"].param_name,
"model.visual.block.qkv_proj.qweight",
)
def test_unquantized_layout_matches_logical_shape(self):
out_features, in_features = 3, 8
payload = np.zeros((out_features, in_features), dtype=np.float32).tobytes()
@@ -20,6 +20,7 @@ from sglang.multimodal_gen.runtime.layers.quantization.configs.kitchen_w4a8_conf
KitchenW4A8Config,
)
from sglang.multimodal_gen.runtime.layers.quantization.fp8 import Fp8Config
from sglang.multimodal_gen.runtime.layers.quantization.gguf import GGUFConfig
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
ComponentCheckpointUnsupportedError,
NativeComponentLoaderRequired,
@@ -27,11 +28,16 @@ from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader imp
from sglang.multimodal_gen.runtime.loader.component_loaders.text_encoder_loader import (
TextEncoderLoader,
_configure_encoder_quantization,
_get_encoder_quant_config,
_process_quantized_encoder_weights,
_require_quantized_encoder_layers,
_resolve_and_configure_encoder_quantization,
)
from sglang.multimodal_gen.runtime.models.encoders.base import TextEncoder
from sglang.multimodal_gen.runtime.loader.gguf_weights import GGUFTensorMeta
from sglang.multimodal_gen.runtime.models.encoders.base import (
EncoderTensorParallelMixin,
TextEncoder,
)
from sglang.multimodal_gen.runtime.models.encoders.minimax_h3_qwen3vl import (
MiniMaxH3ConditioningProjection,
MiniMaxH3Qwen3VLEncoder,
@@ -126,6 +132,13 @@ class TestTextEncoderClassResolution(unittest.TestCase):
def test_unknown_architecture_falls_back_to_automodel(self):
self.assertIs(self._resolve(True, ["NotARealClass"]), transformers.AutoModel)
def test_tensor_parallel_encoder_keeps_checkpoint_weights_by_default(self):
self.assertTrue(
EncoderTensorParallelMixin.should_materialize_checkpoint_weight(
"model.layers.0.self_attn.q_proj.weight"
)
)
def test_bitsandbytes_native_load_requires_resident_encoder(self):
loaded_encoder = nn.Linear(1, 1)
transformers_model_class = SimpleNamespace(
@@ -511,6 +524,57 @@ class TestTextEncoderQuantization(unittest.TestCase):
model_config.quant_config.layer_markers,
)
def test_gguf_maps_h3_names_and_drops_unused_language_layers(self):
self.get_quant_config.return_value = None
def meta(name: str) -> GGUFTensorMeta:
return GGUFTensorMeta(
ggml_type=12,
logical_shape=(2, 256),
stored_shape=(2, 144),
stored_dtype=torch.uint8,
param_name=f"{name.removesuffix('.weight')}.qweight",
)
checkpoint_meta = {
name: meta(name)
for name in (
"model.layers.49.self_attn.q_proj.weight",
"model.layers.50.self_attn.q_proj.weight",
"visual.blocks.0.attn.qkv.weight",
)
}
with mock.patch(
"sglang.multimodal_gen.runtime.loader.component_loaders."
"text_encoder_loader.read_gguf_tensor_meta",
return_value=checkpoint_meta,
):
config = _get_encoder_quant_config(
{},
"/model/text_encoder",
"/weights/encoder.gguf",
MiniMaxH3Qwen3VLEncoder,
)
self.assertIsInstance(config, GGUFConfig)
encoder = MiniMaxH3Qwen3VLEncoder.__new__(MiniMaxH3Qwen3VLEncoder)
encoder.selected_lm_layer = 50
config.retain_tensor_meta(encoder.should_materialize_checkpoint_weight)
self.assertEqual(
config.quantized_prefixes,
{"model.language_model.layers.49.self_attn.q_proj"},
)
vision_meta = config.tensor_meta["model.visual.blocks.0.attn.qkv_proj.weight"]
self.assertTrue(vision_meta.dequantize_on_load)
self.assertEqual(
vision_meta.param_name,
"model.visual.blocks.0.attn.qkv_proj.weight",
)
self.assertNotIn(
"model.language_model.layers.50.self_attn.q_proj.weight",
config.tensor_meta,
)
def test_encoder_must_use_native_loader(self):
model_config = SimpleNamespace(quant_config=None)
with self.assertRaisesRegex(