This commit is contained in:
chelsealong 2026-08-16 13:23:01 +08:00 committed by GitHub
commit 3362586725
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
2 changed files with 180 additions and 2 deletions

View File

@ -1108,6 +1108,16 @@ def _quantized_apply(module, fn, recurse=True):
return module
# Some quantizers write a comfy_quant payload with a weight_scale but no explicit
# "format" (e.g. the MiniMax H3 nvfp4 AWQ checkpoint). Infer the format from the
# on-disk weight dtype in that case, matching the storage dtypes each format uses.
_QUANT_FORMAT_BY_WEIGHT_DTYPE = {
torch.int8: "int8_tensorwise",
torch.float8_e4m3fn: "float8_e4m3fn",
torch.uint8: "nvfp4",
}
def _load_quantized_module(module, super_load, state_dict, prefix, local_metadata, strict,
missing_keys, unexpected_keys, error_msgs, load_extra_params=False):
"""Shared _load_from_state_dict body for quantized-weight modules.
@ -1141,12 +1151,19 @@ def _load_quantized_module(module, super_load, state_dict, prefix, local_metadat
layer_conf = state_dict.pop(f"{prefix}comfy_quant", None)
if layer_conf is not None:
layer_conf = json.loads(layer_conf.numpy().tobytes())
raw_conf = layer_conf.numpy().tobytes()
# Some quantizers mark unquantized layers with an all-NUL comfy_quant
# placeholder instead of omitting it; treat that the same as absent.
layer_conf = json.loads(raw_conf) if raw_conf.strip(b"\x00") else None
if layer_conf is None:
module.quant_format = None
module.layout_type = None
module.weight = torch.nn.Parameter(weight.to(device=device, dtype=compute_dtype), requires_grad=False)
else:
module.quant_format = layer_conf.get("format", None)
if module.quant_format is None and f"{prefix}weight_scale" in state_dict:
module.quant_format = _QUANT_FORMAT_BY_WEIGHT_DTYPE.get(weight.dtype)
module._full_precision_mm_config = layer_conf.get("full_precision_matrix_mult", False)
if not module._full_precision_mm:
module._full_precision_mm = module._full_precision_mm_config
@ -1566,11 +1583,20 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec
weight_key = f"{prefix}weight"
layer_conf = state_dict.pop(f"{prefix}comfy_quant", None)
if layer_conf is not None:
layer_conf = json.loads(layer_conf.numpy().tobytes())
raw_conf = layer_conf.numpy().tobytes()
# Some quantizers mark unquantized layers with an all-NUL comfy_quant
# placeholder instead of omitting it; treat that the same as absent.
layer_conf = json.loads(raw_conf) if raw_conf.strip(b"\x00") else None
# Only fp8 and int8_tensorwise support per-row dequant via index select.
# Block-scaled formats (NVFP4, MXFP8) can't do per-row lookup efficiently.
quant_format = layer_conf.get("format") if layer_conf is not None else None
if quant_format is None and layer_conf is not None and f"{prefix}weight_scale" in state_dict:
_stored_weight = state_dict.get(weight_key)
if _stored_weight is not None:
quant_format = _QUANT_FORMAT_BY_WEIGHT_DTYPE.get(_stored_weight.dtype)
if quant_format == "nvfp4":
raise ValueError(f"NVFP4 embedding format is unsupported for layer {prefix.rstrip('.')}")
manually_loaded_keys = []
if quant_format in ("float8_e4m3fn", "float8_e5m2", "int8_tensorwise") and weight_key in state_dict:

View File

@ -228,6 +228,158 @@ class TestMixedPrecisionOps(unittest.TestCase):
with self.assertRaises(KeyError):
model.load_state_dict(state_dict, strict=False)
def test_all_nul_comfy_quant_marker_loads_as_unquantized(self):
"""Some quantizers mark unquantized layers with an all-NUL comfy_quant
placeholder instead of omitting the key; it must load as plain weight,
not crash decoding it as JSON or raise for a missing format."""
state_dict = {
"layer1.weight": torch.randn(20, 10, dtype=torch.bfloat16),
"layer1.bias": torch.randn(20, dtype=torch.bfloat16),
"layer1.comfy_quant": torch.zeros(29, dtype=torch.uint8),
"layer2.weight": torch.randn(30, 20, dtype=torch.bfloat16),
"layer2.bias": torch.randn(30, dtype=torch.bfloat16),
"layer3.weight": torch.randn(40, 30, dtype=torch.bfloat16),
"layer3.bias": torch.randn(40, dtype=torch.bfloat16),
}
model = SimpleModel(operations=ops.mixed_precision_ops({}))
model.load_state_dict(state_dict, strict=False)
self.assertNotIsInstance(model.layer1.weight, QuantizedTensor)
for layer in [model.layer1, model.layer2, model.layer3]:
layer.weight_function = []
layer.bias_function = []
input_tensor = torch.randn(5, 10, dtype=torch.bfloat16)
output = model(input_tensor)
self.assertEqual(output.shape, (5, 40))
def test_formatless_scaled_comfy_quant_infers_format_from_dtype(self):
"""Some quantizers write a comfy_quant payload with a weight_scale but no
"format" key (e.g. a q_proj-style layer in a MiniMax H3 nvfp4 AWQ
checkpoint). The loader must infer the format from the on-disk weight
dtype instead of raising "Unknown quantization format"."""
state_dict = {
"layer1.weight": torch.randint(-128, 127, (20, 10), dtype=torch.int8),
"layer1.comfy_quant": torch.tensor(list(json.dumps({}).encode("utf-8")), dtype=torch.uint8),
"layer1.weight_scale": torch.ones(20),
"layer1.bias": torch.randn(20, dtype=torch.bfloat16),
"layer2.weight": torch.randn(30, 20, dtype=torch.bfloat16),
"layer2.bias": torch.randn(30, dtype=torch.bfloat16),
"layer3.weight": torch.randn(40, 30, dtype=torch.bfloat16),
"layer3.bias": torch.randn(40, dtype=torch.bfloat16),
}
model = SimpleModel(operations=ops.mixed_precision_ops({}))
model.load_state_dict(state_dict, strict=False)
self.assertIsInstance(model.layer1.weight, QuantizedTensor)
self.assertEqual(model.layer1.quant_format, "int8_tensorwise")
for layer in [model.layer1, model.layer2, model.layer3]:
layer.weight_function = []
layer.bias_function = []
input_tensor = torch.randn(5, 10, dtype=torch.bfloat16)
output = model(input_tensor)
self.assertEqual(output.shape, (5, 40))
def test_formatless_scaled_comfy_quant_embedding_infers_format_from_dtype(self):
"""Same formatless-but-scaled scenario, but for the Embedding load path,
which has its own inline comfy_quant handling separate from
_load_quantized_module."""
operations = ops.mixed_precision_ops({})
class EmbModel(torch.nn.Module):
def __init__(self):
super().__init__()
self.emb = operations.Embedding(100, 20, device="cpu", dtype=torch.bfloat16)
state_dict = {
"emb.weight": torch.randint(-128, 127, (100, 20), dtype=torch.int8),
"emb.comfy_quant": torch.tensor(list(json.dumps({}).encode("utf-8")), dtype=torch.uint8),
"emb.weight_scale": torch.ones(100),
}
model = EmbModel()
model.load_state_dict(state_dict, strict=False)
self.assertIsInstance(model.emb.weight, QuantizedTensor)
self.assertEqual(model.emb.quant_format, "int8_tensorwise")
def test_reload_unquantized_resets_stale_quant_state(self):
"""A module that previously loaded a quantized checkpoint must clear
quant_format/layout_type when reloaded with an unquantized checkpoint,
so forward doesn't take the stale quantized path against what is now
a plain Parameter."""
layer_quant_config = {
"layer1": {
"format": "float8_e4m3fn",
"params": {}
}
}
fp8_weight = torch.randn(20, 10, dtype=torch.float32).to(torch.float8_e4m3fn)
state_dict1 = {
"layer1.weight": fp8_weight,
"layer1.bias": torch.randn(20, dtype=torch.bfloat16),
"layer1.weight_scale": torch.tensor(2.0, dtype=torch.float32),
"layer2.weight": torch.randn(30, 20, dtype=torch.bfloat16),
"layer2.bias": torch.randn(30, dtype=torch.bfloat16),
"layer3.weight": torch.randn(40, 30, dtype=torch.bfloat16),
"layer3.bias": torch.randn(40, dtype=torch.bfloat16),
}
state_dict1, _ = comfy.utils.convert_old_quants(state_dict1, metadata={"_quantization_metadata": json.dumps({"layers": layer_quant_config})})
model = SimpleModel(operations=ops.mixed_precision_ops({}))
model.load_state_dict(state_dict1, strict=False)
self.assertIsInstance(model.layer1.weight, QuantizedTensor)
# Reload layer1 with a plain (unquantized) weight, no comfy_quant key.
state_dict2 = {
"layer1.weight": torch.randn(20, 10, dtype=torch.bfloat16),
"layer1.bias": torch.randn(20, dtype=torch.bfloat16),
"layer2.weight": torch.randn(30, 20, dtype=torch.bfloat16),
"layer2.bias": torch.randn(30, dtype=torch.bfloat16),
"layer3.weight": torch.randn(40, 30, dtype=torch.bfloat16),
"layer3.bias": torch.randn(40, dtype=torch.bfloat16),
}
model.load_state_dict(state_dict2, strict=False)
self.assertNotIsInstance(model.layer1.weight, QuantizedTensor)
self.assertIsNone(model.layer1.quant_format)
self.assertIsNone(model.layer1.layout_type)
for layer in [model.layer1, model.layer2, model.layer3]:
layer.weight_function = []
layer.bias_function = []
input_tensor = torch.randn(5, 10, dtype=torch.bfloat16)
output = model(input_tensor)
self.assertEqual(output.shape, (5, 40))
def test_formatless_scaled_comfy_quant_embedding_rejects_nvfp4(self):
"""A formatless comfy_quant payload that infers nvfp4 from a uint8
weight dtype must raise, since the embedding load path has no
per-row dequant support for NVFP4; it must not silently load the
raw quantized bytes as an ordinary embedding weight."""
operations = ops.mixed_precision_ops({})
class EmbModel(torch.nn.Module):
def __init__(self):
super().__init__()
self.emb = operations.Embedding(100, 20, device="cpu", dtype=torch.bfloat16)
state_dict = {
"emb.weight": torch.randint(0, 255, (100, 20), dtype=torch.uint8),
"emb.comfy_quant": torch.tensor(list(json.dumps({}).encode("utf-8")), dtype=torch.uint8),
"emb.weight_scale": torch.ones(100),
}
model = EmbModel()
with self.assertRaises(ValueError):
model.load_state_dict(state_dict, strict=False)
def test_int8_convrot_metadata_loads_into_params(self):
"""ConvRot metadata must reach TensorWiseINT8Layout params."""
torch.manual_seed(123)