diff --git a/src/diffusers/quantizers/gguf/utils.py b/src/diffusers/quantizers/gguf/utils.py index e9fd8db82f9a..b6d395eeda06 100644 --- a/src/diffusers/quantizers/gguf/utils.py +++ b/src/diffusers/quantizers/gguf/utils.py @@ -548,6 +548,16 @@ def dequantize_gguf_tensor(tensor): class GGUFParameter(torch.nn.Parameter): def __new__(cls, data, requires_grad=False, quant_type=None): data = data if data is not None else torch.empty(0) + if quant_type is None: + # Offloading rebuilds parameters as `param_cls(new_value, requires_grad=...)` without + # forwarding `quant_type` (see `accelerate.utils.set_module_tensor_to_device`), so + # inherit it from the tensor being wrapped instead of failing with `KeyError: None`. + quant_type = getattr(data, "quant_type", None) + if quant_type not in GGML_QUANT_SIZES: + raise ValueError( + f"`GGUFParameter` expects a valid `quant_type`, but got {quant_type}, and it could not " + "be inferred from the tensor being wrapped." + ) self = torch.Tensor._make_subclass(cls, data, requires_grad) self.quant_type = quant_type block_size, type_size = GGML_QUANT_SIZES[quant_type] diff --git a/tests/pipelines/testing_utils/quantization.py b/tests/pipelines/testing_utils/quantization.py index 813d1228647e..b7d7687f5da4 100644 --- a/tests/pipelines/testing_utils/quantization.py +++ b/tests/pipelines/testing_utils/quantization.py @@ -1380,6 +1380,30 @@ def test_pipeline_inference(self): max_diff = numpy_cosine_similarity_distance(expected_slice, output_slice) assert max_diff < 1e-4 + def test_pipeline_inference_sequential_cpu_offload(self): + r""" + Sequential CPU offload rebuilds every parameter through `param_cls(new_value, ...)`, which + used to drop `GGUFParameter.quant_type` and fail with `KeyError: None`. Like the TorchAO + equivalent this only checks that inference runs. + """ + quantization_config = GGUFQuantizationConfig(compute_dtype=self.torch_dtype) + transformer = self.model_cls.from_single_file( + self.ckpt_path, quantization_config=quantization_config, torch_dtype=self.torch_dtype + ) + pipe = FluxPipeline.from_pretrained( + "black-forest-labs/FLUX.1-dev", transformer=transformer, torch_dtype=self.torch_dtype + ) + pipe.enable_sequential_cpu_offload() + + output = pipe( + prompt="a cat holding a sign that says hello", + num_inference_steps=2, + generator=torch.Generator("cpu").manual_seed(0), + output_type="np", + ).images[0] + + assert output.shape == (1024, 1024, 3) + class TestSD35LargeGGUFPipeline(GGUFPipelineTests): ckpt_path = "https://huggingface.co/city96/stable-diffusion-3.5-large-gguf/blob/main/sd3.5_large-Q4_0.gguf"