Support convrot int4 models. (#14859)

linear_dtype in comfy_quant metadata can be used to set if the int4 op does
the matrix multiplication in int8 or int4, the default is int4 on GPUs that
support it with fallback to int8 for GPUs that don't.
This commit is contained in:
comfyanonymous
2026-07-09 15:57:09 -07:00
committed by GitHub
parent 1ea724339c
commit 73e84d5ec8
4 changed files with 94 additions and 2 deletions

View File

@@ -1104,6 +1104,21 @@ def _load_quantized_module(module, super_load, state_dict, prefix, local_metadat
scales["convrot_groupsize"] = int(
layer_conf.get("convrot_groupsize", params_conf.get("convrot_groupsize", 256))
)
elif module.quant_format == "convrot_w4a4":
scale = pop_scale("weight_scale")
if scale is None:
raise ValueError(f"Missing ConvRot W4A4 weight scale for layer {layer_name}")
params_conf = layer_conf.get("params", {})
if not isinstance(params_conf, dict):
params_conf = {}
scales = {
"scale": scale,
"convrot_groupsize": int(
layer_conf.get("convrot_groupsize", params_conf.get("convrot_groupsize", 256))
),
"quant_group_size": 64,
"linear_dtype": layer_conf.get("linear_dtype", params_conf.get("linear_dtype", "int4")),
}
else:
raise ValueError(f"Unsupported quantization format: {module.quant_format}")
@@ -1150,6 +1165,11 @@ def _quantized_weight_state_dict(module, sd, prefix, extra_quant_conf=None, extr
if module.quant_format == "int8_tensorwise" and getattr(params, "convrot", False):
quant_conf["convrot"] = True
quant_conf["convrot_groupsize"] = getattr(params, "convrot_groupsize", 256)
elif module.quant_format == "convrot_w4a4":
quant_conf["convrot_groupsize"] = getattr(params, "convrot_groupsize", 256)
linear_dtype = getattr(params, "linear_dtype", "int4")
if linear_dtype != "int4":
quant_conf["linear_dtype"] = linear_dtype
if extra_quant_conf:
quant_conf.update(extra_quant_conf)
sd[f"{prefix}comfy_quant"] = torch.tensor(list(json.dumps(quant_conf).encode("utf-8")), dtype=torch.uint8)
@@ -1430,6 +1450,12 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec
}
if hasattr(params, "block_scale"): # NVFP4
kwargs["block_scale"] = params.block_scale[i]
if hasattr(params, "quant_group_size"):
kwargs["quant_group_size"] = params.quant_group_size
if hasattr(params, "convrot_groupsize"):
kwargs["convrot_groupsize"] = params.convrot_groupsize
if hasattr(params, "linear_dtype"):
kwargs["linear_dtype"] = params.linear_dtype
return QuantizedTensor(weight._qdata[i], weight._layout_cls, type(params)(**kwargs))
def state_dict(self, *args, destination=None, prefix="", **kwargs):