Fix quantization issues
This commit is contained in:
@@ -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():
|
||||
|
||||
@@ -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"]
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user