[diffusion] feat: support loading serialized comfy convrot int8 native encoders (#36023)
This commit is contained in:
@@ -169,10 +169,10 @@ backend.
|
|||||||
<tr>
|
<tr>
|
||||||
<td><code>comfy-int8-convrot</code></td>
|
<td><code>comfy-int8-convrot</code></td>
|
||||||
<td>One selected safetensors file with per-layer <code>int8_tensorwise</code> and ConvRot metadata</td>
|
<td>One selected safetensors file with per-layer <code>int8_tensorwise</code> and ConvRot metadata</td>
|
||||||
<td><code>--transformer-weights-path</code></td>
|
<td><code>--transformer-weights-path</code>, or an explicit weight file through <code>--component-paths.<component></code></td>
|
||||||
<td>MiniMax-H3 native DiT; pruned FL2VA is E2E-verified and Ref2VA has the same validated tensor contract</td>
|
<td>Native DiTs and encoders whose parameter mappings preserve each marked linear; MiniMax-H3 DiT and Qwen3-VL encoder checkpoints have validated tensor contracts</td>
|
||||||
<td><code>comfy-kitchen</code></td>
|
<td><code>comfy-kitchen</code></td>
|
||||||
<td>CUDA; auto-detected; uses the fused Kitchen INT8 kernel and validates weight/scale layout before model construction. TP requires every row-parallel input shard to preserve the checkpoint's ConvRot group boundary; the H3 256-group checkpoint supports TP1/2/4, not TP8. Offload is supported; FSDP is not.</td>
|
<td>CUDA; auto-detected; uses the fused Kitchen INT8 kernel and validates weight/scale layout before model construction. TP requires every row-parallel input shard to preserve the checkpoint's ConvRot group boundary. The H3 256-group DiT supports TP1/2/4, not TP8; its Qwen3-VL encoder keeps TP8 by replicating only incompatible row projections. Offload is supported; FSDP is not.</td>
|
||||||
</tr>
|
</tr>
|
||||||
<tr>
|
<tr>
|
||||||
<td><code>qvg-kv</code></td>
|
<td><code>qvg-kv</code></td>
|
||||||
@@ -345,9 +345,10 @@ sglang generate \
|
|||||||
|
|
||||||
### Kitchen INT8
|
### Kitchen INT8
|
||||||
|
|
||||||
Serialized Comfy ConvRot INT8 DiTs are selected through
|
Serialized Comfy ConvRot INT8 DiTs use `--transformer-weights-path`; compatible
|
||||||
`--transformer-weights-path` and auto-detected from their per-layer markers.
|
native encoders use an explicit file through `--component-paths.<component>` or
|
||||||
They load INT8 weights and row scales directly; omit `--quantization`.
|
its component alias. Both are auto-detected from per-layer markers and load
|
||||||
|
INT8 weights and row scales directly; omit `--quantization`.
|
||||||
|
|
||||||
For a BF16 checkpoint, `--quantization kitchen_int8` instead performs online
|
For a BF16 checkpoint, `--quantization kitchen_int8` instead performs online
|
||||||
quantization after loading:
|
quantization after loading:
|
||||||
|
|||||||
@@ -93,6 +93,7 @@ class ComfyFp8Config(QuantizationConfig):
|
|||||||
def __init__(self, layer_markers: dict[str, dict[str, Any]]) -> None:
|
def __init__(self, layer_markers: dict[str, dict[str, Any]]) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.layer_markers = layer_markers
|
self.layer_markers = layer_markers
|
||||||
|
self.selected: list[str] = []
|
||||||
self._fp8_config = Fp8Config(
|
self._fp8_config = Fp8Config(
|
||||||
is_checkpoint_fp8_serialized=True,
|
is_checkpoint_fp8_serialized=True,
|
||||||
activation_scheme="static",
|
activation_scheme="static",
|
||||||
@@ -136,6 +137,7 @@ class ComfyFp8Config(QuantizationConfig):
|
|||||||
marker = self.layer_markers.get(prefix)
|
marker = self.layer_markers.get(prefix)
|
||||||
if marker is None:
|
if marker is None:
|
||||||
return UnquantizedLinearMethod()
|
return UnquantizedLinearMethod()
|
||||||
|
self.selected.append(prefix)
|
||||||
if marker.get("full_precision_matrix_mult", False):
|
if marker.get("full_precision_matrix_mult", False):
|
||||||
return ComfyFullPrecisionFp8LinearMethod()
|
return ComfyFullPrecisionFp8LinearMethod()
|
||||||
return Fp8LinearMethod(self._fp8_config)
|
return Fp8LinearMethod(self._fp8_config)
|
||||||
|
|||||||
@@ -154,3 +154,9 @@ class QuantizationConfig(ABC):
|
|||||||
|
|
||||||
def get_cache_scale(self, name: str) -> str | None:
|
def get_cache_scale(self, name: str) -> str | None:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
def supports_input_partition(
|
||||||
|
self, prefix: str, input_size_per_partition: int
|
||||||
|
) -> bool:
|
||||||
|
"""Whether a row-parallel shard preserves this format's input layout."""
|
||||||
|
return True
|
||||||
|
|||||||
+11
@@ -155,3 +155,14 @@ class KitchenInt8Config(QuantizationConfig):
|
|||||||
|
|
||||||
def get_scaled_act_names(self) -> list[str]:
|
def get_scaled_act_names(self) -> list[str]:
|
||||||
return []
|
return []
|
||||||
|
|
||||||
|
def supports_input_partition(
|
||||||
|
self, prefix: str, input_size_per_partition: int
|
||||||
|
) -> bool:
|
||||||
|
group_size = self.group_size
|
||||||
|
if self.layer_markers is not None:
|
||||||
|
marker_group_size = self._serialized_group_sizes.get(prefix)
|
||||||
|
if marker_group_size is None:
|
||||||
|
return True
|
||||||
|
group_size = marker_group_size
|
||||||
|
return input_size_per_partition % group_size == 0
|
||||||
|
|||||||
+90
-21
@@ -27,6 +27,13 @@ from sglang.multimodal_gen.runtime.layers.linear import (
|
|||||||
LinearBase,
|
LinearBase,
|
||||||
UnquantizedLinearMethod,
|
UnquantizedLinearMethod,
|
||||||
)
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.layers.quantization.comfy_fp8 import ComfyFp8Config
|
||||||
|
from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config import (
|
||||||
|
QuantizationConfig,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.layers.quantization.configs.kitchen_int8_config import (
|
||||||
|
KitchenInt8Config,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
|
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
|
||||||
ComponentCheckpointUnsupportedError,
|
ComponentCheckpointUnsupportedError,
|
||||||
ComponentLoader,
|
ComponentLoader,
|
||||||
@@ -34,6 +41,7 @@ from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader imp
|
|||||||
uses_native_transformers_bnb4,
|
uses_native_transformers_bnb4,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.loader.utils import (
|
from sglang.multimodal_gen.runtime.loader.utils import (
|
||||||
|
get_param_names_mapping,
|
||||||
set_default_torch_dtype,
|
set_default_torch_dtype,
|
||||||
skip_init_modules,
|
skip_init_modules,
|
||||||
)
|
)
|
||||||
@@ -65,6 +73,8 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
|||||||
from sglang.multimodal_gen.runtime.utils.quantization_utils import (
|
from sglang.multimodal_gen.runtime.utils.quantization_utils import (
|
||||||
get_quant_config,
|
get_quant_config,
|
||||||
get_quant_config_from_safetensors_metadata,
|
get_quant_config_from_safetensors_metadata,
|
||||||
|
inspect_comfy_quant_markers,
|
||||||
|
resolve_comfy_checkpoint_quantization,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.weights.source import (
|
from sglang.multimodal_gen.runtime.weights.source import (
|
||||||
materialize_weight,
|
materialize_weight,
|
||||||
@@ -104,6 +114,7 @@ def _get_encoder_quant_config(
|
|||||||
component_config: dict,
|
component_config: dict,
|
||||||
component_model_path: str,
|
component_model_path: str,
|
||||||
component_weights_path: str,
|
component_weights_path: str,
|
||||||
|
model_cls: type[nn.Module] | None = None,
|
||||||
):
|
):
|
||||||
quant_config = get_quant_config(component_config, component_model_path)
|
quant_config = get_quant_config(component_config, component_model_path)
|
||||||
if (
|
if (
|
||||||
@@ -114,6 +125,27 @@ def _get_encoder_quant_config(
|
|||||||
quant_config = get_quant_config_from_safetensors_metadata(
|
quant_config = get_quant_config_from_safetensors_metadata(
|
||||||
component_weights_path
|
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(name)
|
||||||
|
if merge_index is not None:
|
||||||
|
raise ValueError(
|
||||||
|
"Comfy quantized component weights cannot use a "
|
||||||
|
"stacked parameter-name mapping"
|
||||||
|
)
|
||||||
|
return mapped_name
|
||||||
|
|
||||||
|
markers = inspect_comfy_quant_markers(
|
||||||
|
[component_weights_path],
|
||||||
|
param_name_mapper=name_mapper,
|
||||||
|
)
|
||||||
|
quant_config = resolve_comfy_checkpoint_quantization(markers)
|
||||||
return quant_config
|
return quant_config
|
||||||
|
|
||||||
|
|
||||||
@@ -140,8 +172,9 @@ def _configure_encoder_quantization(
|
|||||||
component_config,
|
component_config,
|
||||||
component_model_path,
|
component_model_path,
|
||||||
component_weights_path,
|
component_weights_path,
|
||||||
|
model_cls,
|
||||||
)
|
)
|
||||||
except (KeyError, TypeError, ValueError) as error:
|
except (KeyError, NotImplementedError, TypeError, ValueError) as error:
|
||||||
raise ComponentCheckpointUnsupportedError(
|
raise ComponentCheckpointUnsupportedError(
|
||||||
f"Cannot configure checkpoint quantization for {component_name!r}: {error}"
|
f"Cannot configure checkpoint quantization for {component_name!r}: {error}"
|
||||||
) from error
|
) from error
|
||||||
@@ -224,7 +257,7 @@ def _module_tensor_device(module: nn.Module) -> torch.device | None:
|
|||||||
|
|
||||||
def _process_quantized_encoder_weights(
|
def _process_quantized_encoder_weights(
|
||||||
model: nn.Module,
|
model: nn.Module,
|
||||||
process_device: torch.device,
|
process_device: torch.device | None,
|
||||||
component_name: str,
|
component_name: str,
|
||||||
) -> int:
|
) -> int:
|
||||||
processed_layers = 0
|
processed_layers = 0
|
||||||
@@ -236,7 +269,11 @@ def _process_quantized_encoder_weights(
|
|||||||
continue
|
continue
|
||||||
|
|
||||||
origin_device = _module_tensor_device(module)
|
origin_device = _module_tensor_device(module)
|
||||||
should_stage = origin_device is not None and origin_device != process_device
|
should_stage = (
|
||||||
|
process_device is not None
|
||||||
|
and origin_device is not None
|
||||||
|
and origin_device != process_device
|
||||||
|
)
|
||||||
if should_stage:
|
if should_stage:
|
||||||
module.to(process_device)
|
module.to(process_device)
|
||||||
try:
|
try:
|
||||||
@@ -258,18 +295,27 @@ def _process_quantized_encoder_weights(
|
|||||||
def _require_quantized_encoder_layers(
|
def _require_quantized_encoder_layers(
|
||||||
model: nn.Module,
|
model: nn.Module,
|
||||||
component_name: str,
|
component_name: str,
|
||||||
|
quant_config: QuantizationConfig | None = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
if any(
|
has_quantized_layers = any(
|
||||||
isinstance(module, LinearBase)
|
isinstance(module, LinearBase)
|
||||||
and module.quant_method is not None
|
and module.quant_method is not None
|
||||||
and not isinstance(module.quant_method, UnquantizedLinearMethod)
|
and not isinstance(module.quant_method, UnquantizedLinearMethod)
|
||||||
for module in model.modules()
|
for module in model.modules()
|
||||||
):
|
|
||||||
return
|
|
||||||
raise ComponentCheckpointUnsupportedError(
|
|
||||||
f"The native {type(model).__name__} implementation does not construct "
|
|
||||||
f"quantized linear layers for {component_name!r}"
|
|
||||||
)
|
)
|
||||||
|
if not has_quantized_layers:
|
||||||
|
raise ComponentCheckpointUnsupportedError(
|
||||||
|
f"The native {type(model).__name__} implementation does not construct "
|
||||||
|
f"quantized linear layers for {component_name!r}"
|
||||||
|
)
|
||||||
|
if isinstance(quant_config, (ComfyFp8Config, KitchenInt8Config)):
|
||||||
|
missing = set(quant_config.layer_markers) - set(quant_config.selected)
|
||||||
|
if missing:
|
||||||
|
raise ComponentCheckpointUnsupportedError(
|
||||||
|
f"The native {type(model).__name__} implementation did not consume "
|
||||||
|
f"Comfy quantization markers for {component_name!r}: "
|
||||||
|
f"{sorted(missing)[:5]}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def _checkpoint_bytes(model_path: str) -> int:
|
def _checkpoint_bytes(model_path: str) -> int:
|
||||||
@@ -478,6 +524,12 @@ class TextEncoderLoader(ComponentLoader):
|
|||||||
Callable[[str], bool] | None,
|
Callable[[str], bool] | None,
|
||||||
getattr(model, "should_materialize_checkpoint_weight", None),
|
getattr(model, "should_materialize_checkpoint_weight", None),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def include_checkpoint_weight(name: str) -> bool:
|
||||||
|
return not name.endswith(".comfy_quant") and (
|
||||||
|
key_filter is None or key_filter(name)
|
||||||
|
)
|
||||||
|
|
||||||
primary_weights = TextEncoderLoader.Source(
|
primary_weights = TextEncoderLoader.Source(
|
||||||
model_path,
|
model_path,
|
||||||
prefix="",
|
prefix="",
|
||||||
@@ -487,7 +539,7 @@ class TextEncoderLoader(ComponentLoader):
|
|||||||
yield from self._get_weights_iterator(
|
yield from self._get_weights_iterator(
|
||||||
primary_weights,
|
primary_weights,
|
||||||
to_cpu,
|
to_cpu,
|
||||||
key_filter,
|
include_checkpoint_weight,
|
||||||
)
|
)
|
||||||
|
|
||||||
secondary_weights = cast(
|
secondary_weights = cast(
|
||||||
@@ -498,7 +550,7 @@ class TextEncoderLoader(ComponentLoader):
|
|||||||
yield from self._get_weights_iterator(
|
yield from self._get_weights_iterator(
|
||||||
source,
|
source,
|
||||||
to_cpu,
|
to_cpu,
|
||||||
key_filter,
|
include_checkpoint_weight,
|
||||||
)
|
)
|
||||||
|
|
||||||
def load_customized(
|
def load_customized(
|
||||||
@@ -566,14 +618,23 @@ class TextEncoderLoader(ComponentLoader):
|
|||||||
encoder_index
|
encoder_index
|
||||||
]
|
]
|
||||||
# TODO(will): add support for other dtypes
|
# TODO(will): add support for other dtypes
|
||||||
return self.load_model(
|
try:
|
||||||
component_weights_path,
|
return self.load_model(
|
||||||
encoder_config,
|
component_weights_path,
|
||||||
server_args,
|
encoder_config,
|
||||||
encoder_dtype,
|
server_args,
|
||||||
component_starts_on_cpu=component_starts_on_cpu,
|
encoder_dtype,
|
||||||
component_name=component_name,
|
component_starts_on_cpu=component_starts_on_cpu,
|
||||||
)
|
component_name=component_name,
|
||||||
|
)
|
||||||
|
except ComponentCheckpointUnsupportedError:
|
||||||
|
raise
|
||||||
|
except Exception as error:
|
||||||
|
if encoder_config.quant_config is None:
|
||||||
|
raise
|
||||||
|
raise ComponentCheckpointUnsupportedError(
|
||||||
|
f"Failed to load quantized native {component_name!r}: {error}"
|
||||||
|
) from error
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _extract_encoder_index(component_name: str) -> int:
|
def _extract_encoder_index(component_name: str) -> int:
|
||||||
@@ -689,7 +750,9 @@ class TextEncoderLoader(ComponentLoader):
|
|||||||
model.bind_encoder_tp_group(encoder_tp_group)
|
model.bind_encoder_tp_group(encoder_tp_group)
|
||||||
|
|
||||||
if quant_config is not None:
|
if quant_config is not None:
|
||||||
_require_quantized_encoder_layers(model, component_name)
|
_require_quantized_encoder_layers(
|
||||||
|
model, component_name, quant_config=quant_config
|
||||||
|
)
|
||||||
|
|
||||||
if component_starts_on_cpu and (
|
if component_starts_on_cpu and (
|
||||||
current_platform.is_mps() or _keep_this_checkpoint_mapped(model_path)
|
current_platform.is_mps() or _keep_this_checkpoint_mapped(model_path)
|
||||||
@@ -713,9 +776,15 @@ class TextEncoderLoader(ComponentLoader):
|
|||||||
)
|
)
|
||||||
|
|
||||||
if quant_config is not None:
|
if quant_config is not None:
|
||||||
|
postprocess_device: torch.device | None = local_torch_device
|
||||||
|
if (
|
||||||
|
isinstance(quant_config, KitchenInt8Config)
|
||||||
|
and quant_config.is_checkpoint_int8_serialized
|
||||||
|
):
|
||||||
|
postprocess_device = None
|
||||||
processed_layers = _process_quantized_encoder_weights(
|
processed_layers = _process_quantized_encoder_weights(
|
||||||
model,
|
model,
|
||||||
local_torch_device,
|
postprocess_device,
|
||||||
component_name,
|
component_name,
|
||||||
)
|
)
|
||||||
logger.info(
|
logger.info(
|
||||||
|
|||||||
@@ -1,17 +1,16 @@
|
|||||||
# SPDX-License-Identifier: Apache-2.0
|
# SPDX-License-Identifier: Apache-2.0
|
||||||
"""Checkpoint inspection for MiniMax-H3 transformer overrides."""
|
"""Checkpoint inspection for MiniMax-H3 transformer overrides."""
|
||||||
|
|
||||||
import json
|
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
from safetensors import safe_open
|
from safetensors import safe_open
|
||||||
|
|
||||||
from sglang.multimodal_gen.runtime.layers.quantization.comfy_fp8 import ComfyFp8Config
|
|
||||||
from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config import (
|
from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config import (
|
||||||
QuantizationConfig,
|
QuantizationConfig,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.layers.quantization.configs.kitchen_int8_config import (
|
from sglang.multimodal_gen.runtime.utils.quantization_utils import (
|
||||||
KitchenInt8Config,
|
inspect_comfy_quant_markers,
|
||||||
|
resolve_comfy_checkpoint_quantization,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -24,15 +23,11 @@ def inspect_minimax_h3_safetensors(
|
|||||||
) -> tuple[tuple[int, int] | None, dict[str, dict[str, Any]]]:
|
) -> tuple[tuple[int, int] | None, dict[str, dict[str, Any]]]:
|
||||||
"""Read H3 architecture metadata and Comfy per-layer format markers."""
|
"""Read H3 architecture metadata and Comfy per-layer format markers."""
|
||||||
adaln_curve_shape = None
|
adaln_curve_shape = None
|
||||||
layer_markers: dict[str, dict[str, Any]] = {}
|
layer_markers = inspect_comfy_quant_markers(safetensors_list)
|
||||||
checkpoint_keys: set[str] = set()
|
|
||||||
checkpoint_meta: dict[str, tuple[str, tuple[int, ...]]] = {}
|
|
||||||
fp8_weight_prefixes: set[str] = set()
|
|
||||||
|
|
||||||
for path in safetensors_list:
|
for path in safetensors_list:
|
||||||
with safe_open(path, framework="pt", device="cpu") as checkpoint:
|
with safe_open(path, framework="pt", device="cpu") as checkpoint:
|
||||||
keys = checkpoint.keys()
|
keys = checkpoint.keys()
|
||||||
checkpoint_keys.update(keys)
|
|
||||||
if "adaln_t_table" in keys:
|
if "adaln_t_table" in keys:
|
||||||
shape = tuple(checkpoint.get_slice("adaln_t_table").get_shape())
|
shape = tuple(checkpoint.get_slice("adaln_t_table").get_shape())
|
||||||
if len(shape) != 2 or shape[0] < 2:
|
if len(shape) != 2 or shape[0] < 2:
|
||||||
@@ -47,96 +42,13 @@ def inspect_minimax_h3_safetensors(
|
|||||||
)
|
)
|
||||||
adaln_curve_shape = shape
|
adaln_curve_shape = shape
|
||||||
|
|
||||||
for key in keys:
|
|
||||||
if key.endswith((".weight", ".weight_scale")):
|
|
||||||
tensor_slice = checkpoint.get_slice(key)
|
|
||||||
dtype = tensor_slice.get_dtype()
|
|
||||||
checkpoint_meta[key] = (
|
|
||||||
dtype,
|
|
||||||
tuple(tensor_slice.get_shape()),
|
|
||||||
)
|
|
||||||
if key.endswith(".weight") and dtype == "F8_E4M3":
|
|
||||||
fp8_weight_prefixes.add(key.removesuffix(".weight"))
|
|
||||||
if not key.endswith(".comfy_quant"):
|
|
||||||
continue
|
|
||||||
try:
|
|
||||||
marker = json.loads(checkpoint.get_tensor(key).numpy().tobytes())
|
|
||||||
except (json.JSONDecodeError, UnicodeDecodeError) as exc:
|
|
||||||
raise ValueError(
|
|
||||||
f"Invalid Comfy quantization marker {key!r} in {path}"
|
|
||||||
) from exc
|
|
||||||
if not isinstance(marker, dict):
|
|
||||||
raise ValueError(
|
|
||||||
f"Comfy quantization marker {key!r} must contain a JSON object"
|
|
||||||
)
|
|
||||||
prefix = key.removesuffix(".comfy_quant")
|
|
||||||
previous = layer_markers.get(prefix)
|
|
||||||
if previous is not None and previous != marker:
|
|
||||||
raise ValueError(
|
|
||||||
f"Conflicting Comfy quantization markers for {prefix!r}"
|
|
||||||
)
|
|
||||||
layer_markers[prefix] = marker
|
|
||||||
|
|
||||||
if layer_markers:
|
|
||||||
missing_markers = fp8_weight_prefixes - layer_markers.keys()
|
|
||||||
if missing_markers:
|
|
||||||
raise ValueError(
|
|
||||||
"MiniMax-H3 FP8 weights are missing comfy_quant metadata: "
|
|
||||||
f"{sorted(missing_markers)[:5]}"
|
|
||||||
)
|
|
||||||
|
|
||||||
for prefix, marker in layer_markers.items():
|
|
||||||
marker_format = marker.get("format")
|
|
||||||
required = {f"{prefix}.weight", f"{prefix}.weight_scale"}
|
|
||||||
if marker_format == "float8_e4m3fn" and not marker.get(
|
|
||||||
"full_precision_matrix_mult", False
|
|
||||||
):
|
|
||||||
required.add(f"{prefix}.input_scale")
|
|
||||||
if marker_format not in ("float8_e4m3fn", "int8_tensorwise"):
|
|
||||||
continue
|
|
||||||
missing = required - checkpoint_keys
|
|
||||||
if missing:
|
|
||||||
raise ValueError(
|
|
||||||
f"MiniMax-H3 Comfy layer {prefix!r} is missing checkpoint "
|
|
||||||
f"tensors: {sorted(missing)}"
|
|
||||||
)
|
|
||||||
if marker_format == "int8_tensorwise":
|
|
||||||
weight_dtype, weight_shape = checkpoint_meta[f"{prefix}.weight"]
|
|
||||||
scale_dtype, scale_shape = checkpoint_meta[f"{prefix}.weight_scale"]
|
|
||||||
if weight_dtype != "I8" or scale_dtype != "F32":
|
|
||||||
raise ValueError(
|
|
||||||
f"MiniMax-H3 Comfy INT8 layer {prefix!r} needs I8 weights "
|
|
||||||
f"and F32 scales, got {weight_dtype} and {scale_dtype}"
|
|
||||||
)
|
|
||||||
if len(weight_shape) != 2:
|
|
||||||
raise ValueError(
|
|
||||||
f"MiniMax-H3 Comfy INT8 layer {prefix!r} needs a 2D weight, "
|
|
||||||
f"got {weight_shape}"
|
|
||||||
)
|
|
||||||
expected_scale_shape = (weight_shape[0], 1)
|
|
||||||
if scale_shape != expected_scale_shape:
|
|
||||||
raise ValueError(
|
|
||||||
f"MiniMax-H3 Comfy INT8 layer {prefix!r} needs scale shape "
|
|
||||||
f"{expected_scale_shape}, got {scale_shape}"
|
|
||||||
)
|
|
||||||
|
|
||||||
return adaln_curve_shape, layer_markers
|
return adaln_curve_shape, layer_markers
|
||||||
|
|
||||||
|
|
||||||
def resolve_minimax_h3_checkpoint_quantization(
|
def resolve_minimax_h3_checkpoint_quantization(
|
||||||
layer_markers: dict[str, dict[str, Any]],
|
layer_markers: dict[str, dict[str, Any]],
|
||||||
) -> QuantizationConfig | None:
|
) -> QuantizationConfig | None:
|
||||||
if not layer_markers:
|
return resolve_comfy_checkpoint_quantization(layer_markers)
|
||||||
return None
|
|
||||||
|
|
||||||
formats = sorted({str(marker.get("format")) for marker in layer_markers.values()})
|
|
||||||
if formats == ["int8_tensorwise"]:
|
|
||||||
return KitchenInt8Config(layer_markers=layer_markers)
|
|
||||||
if formats == ["float8_e4m3fn"]:
|
|
||||||
return ComfyFp8Config(layer_markers)
|
|
||||||
raise NotImplementedError(
|
|
||||||
"Unsupported MiniMax-H3 Comfy quantization format(s): " + ", ".join(formats)
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def validate_minimax_h3_checkpoint_variant(
|
def validate_minimax_h3_checkpoint_variant(
|
||||||
|
|||||||
@@ -16,12 +16,23 @@ from sglang.multimodal_gen.configs.models.encoders.minimax_h3_qwen3vl import (
|
|||||||
MiniMaxH3Qwen3VLConfig,
|
MiniMaxH3Qwen3VLConfig,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
|
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
|
||||||
|
from sglang.multimodal_gen.runtime.loader.utils import get_param_names_mapping
|
||||||
from sglang.multimodal_gen.runtime.loader.weight_utils import default_weight_loader
|
from sglang.multimodal_gen.runtime.loader.weight_utils import default_weight_loader
|
||||||
from sglang.multimodal_gen.runtime.models.encoders.base import TextEncoder
|
from sglang.multimodal_gen.runtime.models.encoders.base import TextEncoder
|
||||||
from sglang.multimodal_gen.runtime.models.encoders.qwen3vl import Qwen3VLModel
|
from sglang.multimodal_gen.runtime.models.encoders.qwen3vl import Qwen3VLModel
|
||||||
|
|
||||||
MINIMAX_H3_QWEN3VL_HIDDEN_DIM = 5120
|
MINIMAX_H3_QWEN3VL_HIDDEN_DIM = 5120
|
||||||
_LAYER_WEIGHT_RE = re.compile(r"^model\.language_model\.layers\.(\d+)\.")
|
_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"^visual\.": r"model.visual.",
|
||||||
|
r"^(model\.visual\.blocks\.\d+\.attn\.)qkv\.": r"\1qkv_proj.",
|
||||||
|
}
|
||||||
|
_MAP_CHECKPOINT_NAME = get_param_names_mapping(_PARAM_NAMES_MAPPING)
|
||||||
|
|
||||||
|
|
||||||
|
def _map_checkpoint_name(name: str) -> str:
|
||||||
|
return _MAP_CHECKPOINT_NAME(name)[0]
|
||||||
|
|
||||||
|
|
||||||
def _is_unconsumed_checkpoint_weight(name: str) -> bool:
|
def _is_unconsumed_checkpoint_weight(name: str) -> bool:
|
||||||
@@ -46,9 +57,11 @@ class MiniMaxH3Qwen3VLEncoder(TextEncoder):
|
|||||||
layer_names = [*TextEncoder.layer_names, "model.visual.blocks"]
|
layer_names = [*TextEncoder.layer_names, "model.visual.blocks"]
|
||||||
|
|
||||||
supports_dp_encode = True
|
supports_dp_encode = True
|
||||||
|
param_names_mapping = _PARAM_NAMES_MAPPING
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def should_materialize_checkpoint_weight(name: str) -> bool:
|
def should_materialize_checkpoint_weight(name: str) -> bool:
|
||||||
|
name = _map_checkpoint_name(name)
|
||||||
return (
|
return (
|
||||||
"rotary_emb.inv_freq" not in name
|
"rotary_emb.inv_freq" not in name
|
||||||
and not _is_unconsumed_checkpoint_weight(name)
|
and not _is_unconsumed_checkpoint_weight(name)
|
||||||
@@ -67,6 +80,7 @@ class MiniMaxH3Qwen3VLEncoder(TextEncoder):
|
|||||||
arch,
|
arch,
|
||||||
quant_config=config.quant_config,
|
quant_config=config.quant_config,
|
||||||
use_tensor_parallel=True,
|
use_tensor_parallel=True,
|
||||||
|
prefix="model",
|
||||||
)
|
)
|
||||||
# H3 consumes the unnormalized output immediately after layer 49.
|
# H3 consumes the unnormalized output immediately after layer 49.
|
||||||
self.model.language_model.norm = nn.Identity()
|
self.model.language_model.norm = nn.Identity()
|
||||||
@@ -181,9 +195,10 @@ class MiniMaxH3Qwen3VLEncoder(TextEncoder):
|
|||||||
params = dict(self.named_parameters(remove_duplicate=False))
|
params = dict(self.named_parameters(remove_duplicate=False))
|
||||||
loaded: set[str] = set()
|
loaded: set[str] = set()
|
||||||
for name, loaded_weight in weights:
|
for name, loaded_weight in weights:
|
||||||
|
name = _map_checkpoint_name(name)
|
||||||
if not self.should_materialize_checkpoint_weight(name):
|
if not self.should_materialize_checkpoint_weight(name):
|
||||||
continue
|
continue
|
||||||
param_name = name.replace(".attn.qkv.", ".attn.qkv_proj.")
|
param_name = name
|
||||||
param = params.get(param_name)
|
param = params.get(param_name)
|
||||||
if param is None:
|
if param is None:
|
||||||
raise KeyError(
|
raise KeyError(
|
||||||
|
|||||||
@@ -37,6 +37,7 @@ from sglang.multimodal_gen.runtime.models.encoders.qwen_vl_rope import (
|
|||||||
build_qwen_vl_text_rope,
|
build_qwen_vl_text_rope,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum
|
from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum
|
||||||
|
from sglang.multimodal_gen.runtime.utils.common import add_prefix
|
||||||
from sglang.srt.layers.layernorm import RMSNorm
|
from sglang.srt.layers.layernorm import RMSNorm
|
||||||
|
|
||||||
"""Inference-only Qwen3-VL model compatible with HuggingFace weights."""
|
"""Inference-only Qwen3-VL model compatible with HuggingFace weights."""
|
||||||
@@ -165,6 +166,10 @@ def _make_text_row_linear(
|
|||||||
use_row_parallel = (
|
use_row_parallel = (
|
||||||
use_tensor_parallel and tp_size > 1 and in_features % tp_size == 0
|
use_tensor_parallel and tp_size > 1 and in_features % tp_size == 0
|
||||||
)
|
)
|
||||||
|
if use_row_parallel and quant_config is not None:
|
||||||
|
use_row_parallel = quant_config.supports_input_partition(
|
||||||
|
prefix, in_features // tp_size
|
||||||
|
)
|
||||||
if use_weight_only_fp8:
|
if use_weight_only_fp8:
|
||||||
if use_row_parallel:
|
if use_row_parallel:
|
||||||
return WeightOnlyFP8RowParallelLinear(
|
return WeightOnlyFP8RowParallelLinear(
|
||||||
@@ -492,6 +497,7 @@ class Qwen3VLTextModel(nn.Module):
|
|||||||
quant_config: QuantizationConfig | None = None,
|
quant_config: QuantizationConfig | None = None,
|
||||||
use_weight_only_fp8: bool = False,
|
use_weight_only_fp8: bool = False,
|
||||||
use_tensor_parallel: bool = False,
|
use_tensor_parallel: bool = False,
|
||||||
|
prefix: str = "",
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.config = config
|
self.config = config
|
||||||
@@ -509,7 +515,7 @@ class Qwen3VLTextModel(nn.Module):
|
|||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
use_weight_only_fp8=use_weight_only_fp8,
|
use_weight_only_fp8=use_weight_only_fp8,
|
||||||
use_tensor_parallel=use_tensor_parallel,
|
use_tensor_parallel=use_tensor_parallel,
|
||||||
prefix=f"layers.{layer_idx}",
|
prefix=add_prefix(f"layers.{layer_idx}", prefix),
|
||||||
)
|
)
|
||||||
for layer_idx in range(config.num_hidden_layers)
|
for layer_idx in range(config.num_hidden_layers)
|
||||||
]
|
]
|
||||||
@@ -670,6 +676,7 @@ class Qwen3VLModel(nn.Module):
|
|||||||
*,
|
*,
|
||||||
quant_config: QuantizationConfig | None = None,
|
quant_config: QuantizationConfig | None = None,
|
||||||
use_tensor_parallel: bool = False,
|
use_tensor_parallel: bool = False,
|
||||||
|
prefix: str = "",
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.visual = Qwen3VLVisionTransformer(config.vision_config)
|
self.visual = Qwen3VLVisionTransformer(config.vision_config)
|
||||||
@@ -677,6 +684,7 @@ class Qwen3VLModel(nn.Module):
|
|||||||
config.text_config,
|
config.text_config,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
use_tensor_parallel=use_tensor_parallel,
|
use_tensor_parallel=use_tensor_parallel,
|
||||||
|
prefix=add_prefix("language_model", prefix),
|
||||||
)
|
)
|
||||||
self.rope_deltas = None # cache rope_deltas here
|
self.rope_deltas = None # cache rope_deltas here
|
||||||
self.config = config
|
self.config = config
|
||||||
|
|||||||
@@ -4,7 +4,7 @@ import os
|
|||||||
import re
|
import re
|
||||||
import struct
|
import struct
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any, Dict, List, Optional
|
from typing import Any, Callable, Dict, List, Optional
|
||||||
|
|
||||||
from safetensors import safe_open
|
from safetensors import safe_open
|
||||||
|
|
||||||
@@ -12,6 +12,10 @@ from sglang.multimodal_gen.runtime.layers.quantization import (
|
|||||||
QuantizationConfig,
|
QuantizationConfig,
|
||||||
get_quantization_config,
|
get_quantization_config,
|
||||||
)
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.layers.quantization.comfy_fp8 import ComfyFp8Config
|
||||||
|
from sglang.multimodal_gen.runtime.layers.quantization.configs.kitchen_int8_config import (
|
||||||
|
KitchenInt8Config,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
from sglang.srt.layers.modelopt_utils import canonicalize_modelopt_quant_algo
|
from sglang.srt.layers.modelopt_utils import canonicalize_modelopt_quant_algo
|
||||||
from sglang.srt.model_loader.checkpoint_quantization import (
|
from sglang.srt.model_loader.checkpoint_quantization import (
|
||||||
@@ -21,6 +25,111 @@ from sglang.srt.model_loader.checkpoint_quantization import (
|
|||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
def inspect_comfy_quant_markers(
|
||||||
|
safetensors_list: list[str],
|
||||||
|
param_name_mapper: Callable[[str], str] | None = None,
|
||||||
|
) -> dict[str, dict[str, Any]]:
|
||||||
|
"""Read and validate Comfy's tensor-level quantization markers."""
|
||||||
|
checkpoint_meta: dict[str, tuple[str, tuple[int, ...]]] = {}
|
||||||
|
raw_markers: dict[str, dict[str, Any]] = {}
|
||||||
|
marked_dtype_weight_prefixes: set[str] = set()
|
||||||
|
|
||||||
|
for path in safetensors_list:
|
||||||
|
with safe_open(path, framework="pt", device="cpu") as checkpoint:
|
||||||
|
for key in checkpoint.keys():
|
||||||
|
tensor_slice = checkpoint.get_slice(key)
|
||||||
|
checkpoint_meta[key] = (
|
||||||
|
tensor_slice.get_dtype(),
|
||||||
|
tuple(tensor_slice.get_shape()),
|
||||||
|
)
|
||||||
|
if key.endswith(".weight") and tensor_slice.get_dtype() in (
|
||||||
|
"F8_E4M3",
|
||||||
|
"I8",
|
||||||
|
):
|
||||||
|
marked_dtype_weight_prefixes.add(key.removesuffix(".weight"))
|
||||||
|
if not key.endswith(".comfy_quant"):
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
marker = json.loads(checkpoint.get_tensor(key).numpy().tobytes())
|
||||||
|
except (json.JSONDecodeError, UnicodeDecodeError) as exc:
|
||||||
|
raise ValueError(
|
||||||
|
f"Invalid Comfy quantization marker {key!r} in {path}"
|
||||||
|
) from exc
|
||||||
|
if not isinstance(marker, dict):
|
||||||
|
raise ValueError(
|
||||||
|
f"Comfy quantization marker {key!r} must contain a JSON object"
|
||||||
|
)
|
||||||
|
prefix = key.removesuffix(".comfy_quant")
|
||||||
|
previous = raw_markers.get(prefix)
|
||||||
|
if previous is not None and previous != marker:
|
||||||
|
raise ValueError(
|
||||||
|
f"Conflicting Comfy quantization markers for {prefix!r}"
|
||||||
|
)
|
||||||
|
raw_markers[prefix] = marker
|
||||||
|
|
||||||
|
missing_markers = marked_dtype_weight_prefixes - raw_markers.keys()
|
||||||
|
if missing_markers:
|
||||||
|
raise ValueError(
|
||||||
|
"Quantized weights are missing comfy_quant metadata: "
|
||||||
|
f"{sorted(missing_markers)[:5]}"
|
||||||
|
)
|
||||||
|
|
||||||
|
for prefix, marker in raw_markers.items():
|
||||||
|
marker_format = marker.get("format")
|
||||||
|
required = {f"{prefix}.weight", f"{prefix}.weight_scale"}
|
||||||
|
if marker_format == "float8_e4m3fn" and not marker.get(
|
||||||
|
"full_precision_matrix_mult", False
|
||||||
|
):
|
||||||
|
required.add(f"{prefix}.input_scale")
|
||||||
|
if marker_format not in ("float8_e4m3fn", "int8_tensorwise"):
|
||||||
|
continue
|
||||||
|
missing = required - checkpoint_meta.keys()
|
||||||
|
if missing:
|
||||||
|
raise ValueError(
|
||||||
|
f"Comfy layer {prefix!r} is missing checkpoint tensors: "
|
||||||
|
f"{sorted(missing)}"
|
||||||
|
)
|
||||||
|
if marker_format != "int8_tensorwise":
|
||||||
|
continue
|
||||||
|
weight_dtype, weight_shape = checkpoint_meta[f"{prefix}.weight"]
|
||||||
|
scale_dtype, scale_shape = checkpoint_meta[f"{prefix}.weight_scale"]
|
||||||
|
if weight_dtype != "I8" or scale_dtype != "F32":
|
||||||
|
raise ValueError(
|
||||||
|
f"Comfy INT8 layer {prefix!r} needs I8 weights and F32 scales, "
|
||||||
|
f"got {weight_dtype} and {scale_dtype}"
|
||||||
|
)
|
||||||
|
if len(weight_shape) != 2 or scale_shape != (weight_shape[0], 1):
|
||||||
|
raise ValueError(
|
||||||
|
f"Comfy INT8 layer {prefix!r} has incompatible weight/scale "
|
||||||
|
f"shapes: {weight_shape} and {scale_shape}"
|
||||||
|
)
|
||||||
|
|
||||||
|
mapped_markers: dict[str, dict[str, Any]] = {}
|
||||||
|
for prefix, marker in raw_markers.items():
|
||||||
|
mapped_prefix = param_name_mapper(prefix) if param_name_mapper else prefix
|
||||||
|
if mapped_prefix in mapped_markers:
|
||||||
|
raise ValueError(
|
||||||
|
f"Comfy markers collide after parameter mapping at {mapped_prefix!r}"
|
||||||
|
)
|
||||||
|
mapped_markers[mapped_prefix] = marker
|
||||||
|
return mapped_markers
|
||||||
|
|
||||||
|
|
||||||
|
def resolve_comfy_checkpoint_quantization(
|
||||||
|
layer_markers: dict[str, dict[str, Any]],
|
||||||
|
) -> QuantizationConfig | None:
|
||||||
|
if not layer_markers:
|
||||||
|
return None
|
||||||
|
formats = sorted({str(marker.get("format")) for marker in layer_markers.values()})
|
||||||
|
if formats == ["int8_tensorwise"]:
|
||||||
|
return KitchenInt8Config(layer_markers=layer_markers)
|
||||||
|
if formats == ["float8_e4m3fn"]:
|
||||||
|
return ComfyFp8Config(layer_markers)
|
||||||
|
raise NotImplementedError(
|
||||||
|
"Unsupported Comfy quantization format(s): " + ", ".join(formats)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def normalize_flat_modelopt_quant_config(
|
def normalize_flat_modelopt_quant_config(
|
||||||
quant_cfg: dict[str, Any] | None,
|
quant_cfg: dict[str, Any] | None,
|
||||||
) -> dict[str, Any] | None:
|
) -> dict[str, Any] | None:
|
||||||
|
|||||||
@@ -1,12 +1,18 @@
|
|||||||
|
import json
|
||||||
|
import tempfile
|
||||||
import unittest
|
import unittest
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
from unittest import mock
|
from unittest import mock
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import transformers
|
import transformers
|
||||||
|
from safetensors.torch import save_file
|
||||||
from torch import nn
|
from torch import nn
|
||||||
|
|
||||||
from sglang.multimodal_gen.runtime.layers.linear import LinearBase
|
from sglang.multimodal_gen.runtime.layers.linear import LinearBase
|
||||||
|
from sglang.multimodal_gen.runtime.layers.quantization.configs.kitchen_int8_config import (
|
||||||
|
KitchenInt8Config,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.runtime.layers.quantization.fp8 import Fp8Config
|
from sglang.multimodal_gen.runtime.layers.quantization.fp8 import Fp8Config
|
||||||
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
|
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
|
||||||
ComponentCheckpointUnsupportedError,
|
ComponentCheckpointUnsupportedError,
|
||||||
@@ -135,6 +141,9 @@ class TestMiniMaxH3CheckpointFilter(unittest.TestCase):
|
|||||||
"lm_head.weight": False,
|
"lm_head.weight": False,
|
||||||
"model.language_model.rotary_emb.inv_freq": False,
|
"model.language_model.rotary_emb.inv_freq": False,
|
||||||
"model.visual.blocks.0.attn.qkv.weight": True,
|
"model.visual.blocks.0.attn.qkv.weight": True,
|
||||||
|
"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": True,
|
||||||
"module.model.language_model.layers.63.mlp.down_proj.weight": True,
|
"module.model.language_model.layers.63.mlp.down_proj.weight": True,
|
||||||
}
|
}
|
||||||
@@ -163,6 +172,25 @@ class TestMiniMaxH3CheckpointFilter(unittest.TestCase):
|
|||||||
torch.tensor([1.0, 2.0]),
|
torch.tensor([1.0, 2.0]),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def test_comfy_language_checkpoint_name_maps_to_native_namespace(self):
|
||||||
|
encoder = MiniMaxH3Qwen3VLEncoder.__new__(MiniMaxH3Qwen3VLEncoder)
|
||||||
|
torch.nn.Module.__init__(encoder)
|
||||||
|
encoder.model = torch.nn.Module()
|
||||||
|
encoder.model.language_model = torch.nn.Module()
|
||||||
|
layer = torch.nn.Module()
|
||||||
|
layer.self_attn = torch.nn.Module()
|
||||||
|
layer.self_attn.q_proj = torch.nn.Linear(2, 2, bias=False)
|
||||||
|
encoder.model.language_model.layers = torch.nn.ModuleList([layer])
|
||||||
|
source = torch.tensor([[1.0, 2.0], [3.0, 4.0]])
|
||||||
|
|
||||||
|
loaded = encoder.load_weights(
|
||||||
|
[("model.layers.0.self_attn.q_proj.weight", source)]
|
||||||
|
)
|
||||||
|
|
||||||
|
target_name = "model.language_model.layers.0.self_attn.q_proj.weight"
|
||||||
|
self.assertEqual(loaded, {target_name})
|
||||||
|
torch.testing.assert_close(layer.self_attn.q_proj.weight, source)
|
||||||
|
|
||||||
|
|
||||||
class TestTextEncoderQuantization(unittest.TestCase):
|
class TestTextEncoderQuantization(unittest.TestCase):
|
||||||
def setUp(self):
|
def setUp(self):
|
||||||
@@ -212,6 +240,49 @@ class TestTextEncoderQuantization(unittest.TestCase):
|
|||||||
self.assertIs(model_config.quant_config, self.serialized)
|
self.assertIs(model_config.quant_config, self.serialized)
|
||||||
get_file_quant_config.assert_called_once_with("/weights/encoder.safetensors")
|
get_file_quant_config.assert_called_once_with("/weights/encoder.safetensors")
|
||||||
|
|
||||||
|
def test_comfy_int8_weight_file_configures_native_encoder(self):
|
||||||
|
self.get_quant_config.return_value = None
|
||||||
|
marker = json.dumps(
|
||||||
|
{
|
||||||
|
"format": "int8_tensorwise",
|
||||||
|
"convrot": True,
|
||||||
|
"convrot_groupsize": 256,
|
||||||
|
}
|
||||||
|
).encode()
|
||||||
|
with tempfile.NamedTemporaryFile(suffix=".safetensors") as checkpoint:
|
||||||
|
save_file(
|
||||||
|
{
|
||||||
|
"model.layers.0.self_attn.q_proj.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(
|
||||||
|
list(marker), dtype=torch.uint8
|
||||||
|
),
|
||||||
|
},
|
||||||
|
checkpoint.name,
|
||||||
|
)
|
||||||
|
model_config = SimpleNamespace(quant_config=None)
|
||||||
|
with mock.patch(
|
||||||
|
"sglang.multimodal_gen.runtime.loader.component_loaders."
|
||||||
|
"text_encoder_loader.get_quant_config_from_safetensors_metadata",
|
||||||
|
return_value=None,
|
||||||
|
):
|
||||||
|
_configure_encoder_quantization(
|
||||||
|
model_config,
|
||||||
|
MiniMaxH3Qwen3VLEncoder,
|
||||||
|
{},
|
||||||
|
"/model/text_encoder",
|
||||||
|
checkpoint.name,
|
||||||
|
"text_encoder",
|
||||||
|
)
|
||||||
|
|
||||||
|
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"},
|
||||||
|
)
|
||||||
|
|
||||||
def test_encoder_must_use_native_loader(self):
|
def test_encoder_must_use_native_loader(self):
|
||||||
model_config = SimpleNamespace(quant_config=None)
|
model_config = SimpleNamespace(quant_config=None)
|
||||||
with self.assertRaisesRegex(
|
with self.assertRaisesRegex(
|
||||||
@@ -350,6 +421,25 @@ class TestQuantizedTextEncoderPostprocess(unittest.TestCase):
|
|||||||
):
|
):
|
||||||
_require_quantized_encoder_layers(nn.Linear(2, 2), "text_encoder")
|
_require_quantized_encoder_layers(nn.Linear(2, 2), "text_encoder")
|
||||||
|
|
||||||
|
def test_rejects_unconsumed_comfy_marker(self):
|
||||||
|
config = KitchenInt8Config(
|
||||||
|
layer_markers={
|
||||||
|
"visual.proj": {
|
||||||
|
"format": "int8_tensorwise",
|
||||||
|
"convrot": True,
|
||||||
|
"convrot_groupsize": 256,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
)
|
||||||
|
with self.assertRaisesRegex(
|
||||||
|
ComponentCheckpointUnsupportedError, "did not consume"
|
||||||
|
):
|
||||||
|
_require_quantized_encoder_layers(
|
||||||
|
_QuantizedEncoder(_RecordingQuantMethod()),
|
||||||
|
"text_encoder",
|
||||||
|
quant_config=config,
|
||||||
|
)
|
||||||
|
|
||||||
def test_processes_quantized_layers_without_moving_the_model(self):
|
def test_processes_quantized_layers_without_moving_the_model(self):
|
||||||
quant_method = _RecordingQuantMethod()
|
quant_method = _RecordingQuantMethod()
|
||||||
model = _QuantizedEncoder(quant_method)
|
model = _QuantizedEncoder(quant_method)
|
||||||
|
|||||||
@@ -310,6 +310,8 @@ class TestTransformerQuantHelpers(unittest.TestCase):
|
|||||||
self.assertTrue(config.checkpoint_uses_native_qkv_layout)
|
self.assertTrue(config.checkpoint_uses_native_qkv_layout)
|
||||||
self.assertFalse(KitchenInt8Config().checkpoint_uses_native_qkv_layout)
|
self.assertFalse(KitchenInt8Config().checkpoint_uses_native_qkv_layout)
|
||||||
self.assertFalse(_needs_device_weight_postprocess(config))
|
self.assertFalse(_needs_device_weight_postprocess(config))
|
||||||
|
self.assertTrue(config.supports_input_partition("blocks.0.mlp.fc1", 6400))
|
||||||
|
self.assertFalse(config.supports_input_partition("blocks.0.mlp.fc1", 3200))
|
||||||
|
|
||||||
@patch(
|
@patch(
|
||||||
"sglang.multimodal_gen.runtime.layers.quantization.kitchen_int8."
|
"sglang.multimodal_gen.runtime.layers.quantization.kitchen_int8."
|
||||||
|
|||||||
Reference in New Issue
Block a user