Fix GGUF dequantization shape error on MPS (#403)
Skip GGUF quantized buffers in _force_nadit_precision - these must remain in packed format for on-the-fly dequantization during inference.
This commit is contained in:
@@ -826,8 +826,11 @@ class CompatibleDiT(torch.nn.Module):
|
||||
param.data = param.data.to(target_dtype)
|
||||
converted_count += 1
|
||||
|
||||
# Also convert buffers
|
||||
# Also convert buffers (skip GGUF quantized buffers - they have tensor_type attribute)
|
||||
for name, buffer in self.dit_model.named_buffers():
|
||||
# Skip GGUF quantized buffers - these must stay in packed format for on-the-fly dequantization
|
||||
if hasattr(buffer, 'tensor_type'):
|
||||
continue
|
||||
if buffer.dtype != target_dtype:
|
||||
if buffer.device.type == "mps":
|
||||
temp_cpu = buffer.data.to("cpu")
|
||||
|
||||
Reference in New Issue
Block a user