Fix quantization issues

This commit is contained in:
Jaret Burkett
2026-08-28 10:48:34 -06:00
parent 85a6880643
commit 92df289931
7 changed files with 470 additions and 19 deletions

View File

@@ -91,6 +91,13 @@ class OstrisModelMixin:
# precision instead
aitk_cast_on_load: bool = True
# pre-quantized (comfy marker) checkpoints normally load their
# NON-quantized tensors at stored precision (the mix can be deliberate,
# e.g. ltx2.5's fp32 tables). Classes whose forward assumes one uniform
# dtype (wan: fp16 conv biases next to fp32 tables in the fp8 files) set
# this True to cast those float tensors to the load dtype instead
aitk_cast_quantized_load: bool = False
# ---- state set by the loader / quantizer ----
aitk_is_quantized: bool = False
aitk_qtype: Optional[str] = None
@@ -461,7 +468,22 @@ class OstrisModelMixin:
state_dict, num_quantized = import_comfy_quantized_layers(
model, state_dict, orig_dtype=dtype
)
if cls.aitk_cast_quantized_load:
# uniform-dtype models: the leftover (non-quantized) float
# tensors follow the load dtype, like the non-marker path
for key, value in state_dict.items():
if value.is_floating_point():
state_dict[key] = value.to(dtype=dtype)
cls._load_state_dict_with_quantized(model, state_dict)
if cls.aitk_cast_quantized_load:
# the importer assigns quantized-layer biases directly at
# stored dtype; they must follow too (a stray fp16 bias also
# makes ModelMixin.dtype — the pipelines' cast target — lie)
from toolkit.util.ostris_quant import OstrisLinear
for m in model.modules():
if isinstance(m, OstrisLinear) and m.bias is not None:
m.bias.data = m.bias.data.to(dtype=dtype)
model.aitk_is_quantized = True
elif cls.aitk_cast_on_load:
for key, value in state_dict.items():

View File

@@ -76,6 +76,12 @@ class WanTransformer3DModel(DiffusersWanTransformer3DModel, OstrisModelMixin):
},
}
# the comfy fp8/scaled_fp8 wan files mix fp16 conv biases with fp32
# tables; wan's forward assumes one uniform dtype, so cast the
# non-quantized tensors to the load dtype (the old holder did this with a
# blanket .to(device, dtype) after load)
aitk_cast_quantized_load = True
@classmethod
def get_transformer_block_names(cls):
return ["blocks"]

View File

@@ -34,6 +34,9 @@ class UMT5TextEncoder(UMT5EncoderModel, OstrisTransformersMixin):
aitk_subfolder = "text_encoder"
aitk_tokenizer_subfolder = "tokenizer"
aitk_config_repo = "ai-toolkit/umt5_xxl_encoder"
# comfy fp8 umt5 files store non-quantized tensors in a dtype mix; T5
# assumes a uniform dtype, so follow the load dtype
aitk_cast_quantized_load = True
aitk_comfy_repo = "Comfy-Org/Wan_2.1_ComfyUI_repackaged"
# comfy umt5 files already use the transformers key layout; the

View File

@@ -1319,6 +1319,19 @@ class ConvRotInt8Quantizer(OstrisQuantizer):
f"(needs in divisible by 16, out by 8, and a power-of-4 block >= 16)"
)
return False
if d < 128 and module.out_features % 16 != 0:
# cublasLt has no int8 kernel for K < 128 with N % 16 != 0
# (torch._int_mm raises CUBLAS_STATUS_NOT_SUPPORTED), e.g.
# omnigen2's 64 -> 2520 x_embedder
key = (d, module.out_features)
if key not in _skip_warned:
_skip_warned.add(key)
print_acc(
f"ConvRot: skipping linears with in_features={d}, "
f"out_features={module.out_features} (int8 gemm needs out "
f"divisible by 16 when in < 128)"
)
return False
return True
def quantize_(self, module: torch.nn.Linear, weight_fp32: torch.Tensor) -> None: