Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
30 changes: 28 additions & 2 deletions comfy/ops.py
Original file line number Diff line number Diff line change
Expand Up @@ -1118,6 +1118,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.
Expand Down Expand Up @@ -1151,12 +1161,19 @@ def pop_scale(name, dtype=None):

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)
Comment thread
coderabbitai[bot] marked this conversation as resolved.
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
Expand Down Expand Up @@ -1587,11 +1604,20 @@ def _load_from_state_dict(self, state_dict, prefix, local_metadata, strict, miss
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)
Comment thread
coderabbitai[bot] marked this conversation as resolved.
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:
Expand Down
152 changes: 152 additions & 0 deletions tests-unit/comfy_quant/test_mixed_precision.py
Original file line number Diff line number Diff line change
Expand Up @@ -228,6 +228,158 @@ def test_error_handling_unknown_format(self):
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)
Expand Down
Loading