From 362ec0d41290a6ea452070ec67023a7794da1980 Mon Sep 17 00:00:00 2001 From: Disty0 Date: Tue, 12 Aug 2025 22:07:52 +0300 Subject: [PATCH] Fix Chroma quantization --- pipelines/generic.py | 8 ++++---- pipelines/model_chroma.py | 2 +- 2 files changed, 5 insertions(+), 5 deletions(-) diff --git a/pipelines/generic.py b/pipelines/generic.py index 4983d060a..90f9d034e 100644 --- a/pipelines/generic.py +++ b/pipelines/generic.py @@ -8,10 +8,10 @@ from modules import shared, devices, errors, sd_models, model_quant debug = os.environ.get('SD_LOAD_DEBUG', None) is not None -def load_transformer(repo_id, cls_name, load_config={}, subfolder="transformer", allow_quant=True, variant=None, dtype=None): +def load_transformer(repo_id, cls_name, load_config={}, subfolder="transformer", allow_quant=True, variant=None, dtype=None, modules_to_not_convert=[]): transformer = None try: - load_args, quant_args = model_quant.get_dit_args(load_config, module='Model', device_map=True, allow_quant=allow_quant) + load_args, quant_args = model_quant.get_dit_args(load_config, module='Model', device_map=True, allow_quant=allow_quant, modules_to_not_convert=modules_to_not_convert) quant_type = model_quant.get_quant_type(quant_args) dtype = dtype or devices.dtype @@ -71,10 +71,10 @@ def load_transformer(repo_id, cls_name, load_config={}, subfolder="transformer", return transformer -def load_text_encoder(repo_id, cls_name, load_config={}, subfolder="text_encoder", allow_quant=True, allow_shared=True, variant=None, dtype=None): +def load_text_encoder(repo_id, cls_name, load_config={}, subfolder="text_encoder", allow_quant=True, allow_shared=True, variant=None, dtype=None, modules_to_not_convert=[]): text_encoder = None try: - load_args, quant_args = model_quant.get_dit_args(load_config, module='TE', device_map=True, allow_quant=allow_quant) + load_args, quant_args = model_quant.get_dit_args(load_config, module='TE', device_map=True, allow_quant=allow_quant, modules_to_not_convert=modules_to_not_convert) quant_type = model_quant.get_quant_type(quant_args) dtype = dtype or devices.dtype diff --git a/pipelines/model_chroma.py b/pipelines/model_chroma.py index eb94bfed4..b5f60d380 100644 --- a/pipelines/model_chroma.py +++ b/pipelines/model_chroma.py @@ -11,7 +11,7 @@ def load_chroma(checkpoint_info, diffusers_load_config={}): load_args, _quant_args = model_quant.get_dit_args(diffusers_load_config, allow_quant=False) shared.log.debug(f'Load model: type=Chroma repo="{repo_id}" config={diffusers_load_config} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args}') - transformer = generic.load_transformer(repo_id, cls_name=diffusers.ChromaTransformer2DModel, load_config=diffusers_load_config) + transformer = generic.load_transformer(repo_id, cls_name=diffusers.ChromaTransformer2DModel, load_config=diffusers_load_config, modules_to_not_convert=["distilled_guidance_layer"]) text_encoder = generic.load_text_encoder(repo_id, cls_name=transformers.T5EncoderModel, load_config=diffusers_load_config) pipe = diffusers.ChromaPipeline.from_pretrained(