diff --git a/modules/intel/ipex/diffusers.py b/modules/intel/ipex/diffusers.py index d3487fefd..1f0391295 100644 --- a/modules/intel/ipex/diffusers.py +++ b/modules/intel/ipex/diffusers.py @@ -26,7 +26,7 @@ class FluxPosEmbed(torch.nn.Module): n_axes = ids.shape[-1] cos_out = [] sin_out = [] - pos = ids.float() + pos = ids.to(dtype=torch.float32) for i in range(n_axes): cos, sin = diffusers.models.embeddings.get_1d_rotary_pos_embed( self.axes_dim[i], @@ -99,12 +99,12 @@ def apply_rotary_emb(x, freqs_cis, use_real: bool = True, use_real_unbind_dim: i else: raise ValueError(f"`use_real_unbind_dim={use_real_unbind_dim}` but should be -1 or -2.") - out = (x.float() * cos + x_rotated.float() * sin).to(x.dtype) + out = (x.to(dtype=torch.float32) * cos + x_rotated.to(dtype=torch.float32) * sin).to(x.dtype) return out else: # used for lumina # force cpu with Alchemist - x_rotated = torch.view_as_complex(x.to("cpu").float().reshape(*x.shape[:-1], -1, 2)) + x_rotated = torch.view_as_complex(x.to("cpu").to(dtype=torch.float32).reshape(*x.shape[:-1], -1, 2)) freqs_cis = freqs_cis.to("cpu").unsqueeze(2) x_out = torch.view_as_real(x_rotated * freqs_cis).flatten(3) return x_out.type_as(x).to(x.device) @@ -122,5 +122,6 @@ def ipex_diffusers(device_supports_fp64=False): diffusers.models.embeddings.apply_rotary_emb = apply_rotary_emb diffusers.models.transformers.transformer_flux.FluxPosEmbed = FluxPosEmbed diffusers.models.transformers.transformer_lumina2.apply_rotary_emb = apply_rotary_emb - diffusers.models.controlnets.controlnet_flux.FluxPosEmbed = FluxPosEmbed diffusers.models.transformers.transformer_hidream_image.rope = hidream_rope + diffusers.models.transformers.transformer_chroma.FluxPosEmbed = FluxPosEmbed + diffusers.models.controlnets.controlnet_flux.FluxPosEmbed = FluxPosEmbed diff --git a/modules/model_chroma.py b/modules/model_chroma.py index b722837fc..f694917f4 100644 --- a/modules/model_chroma.py +++ b/modules/model_chroma.py @@ -115,10 +115,10 @@ def load_quants(kwargs, pretrained_model_name_or_path, cache_dir, allow_quant): if 'transformer' not in kwargs and model_quant.check_nunchaku('Transformer'): raise NotImplementedError('Nunchaku does not support Chroma Model yet. See https://github.com/mit-han-lab/nunchaku/issues/167') elif 'transformer' not in kwargs and model_quant.check_quant('Transformer'): - quant_args = model_quant.create_config(allow=allow_quant, module='Transformer') + quant_args = model_quant.create_config(allow=allow_quant, module='Transformer', modules_to_not_convert=["distilled_guidance_layer"]) if quant_args: if os.path.isfile(pretrained_model_name_or_path): - kwargs['transformer'] = diffusers.ChromaTransformer2DModel.from_single_file(pretrained_model_name_or_path, cache_dir=cache_dir, torch_dtype=devices.dtype, **quant_args) + kwargs['transformer'] = diffusers.ChromaTransformer2DModel.from_single_file(pretrained_model_name_or_path, cache_dir=cache_dir, torch_dtype=devices.dtype, **quant_args) pass else: kwargs['transformer'] = diffusers.ChromaTransformer2DModel.from_pretrained(pretrained_model_name_or_path, subfolder="transformer", cache_dir=cache_dir, torch_dtype=devices.dtype, **quant_args) @@ -179,7 +179,7 @@ def load_transformer(file_path): # triggered by opts.sd_unet change transformer, _text_encoder = load_chroma_nf4(file_path, prequantized=False) if transformer is not None: return transformer - quant_args = model_quant.create_config(module='Transformer') + quant_args = model_quant.create_config(module='Transformer', modules_to_not_convert=["distilled_guidance_layer"]) if quant_args: shared.log.info(f'Load module: type=UNet/Transformer file="{file_path}" offload={shared.opts.diffusers_offload_mode} quant=torchao dtype={devices.dtype}') transformer = diffusers.ChromaTransformer2DModel.from_single_file(file_path, **diffusers_load_config, **quant_args) @@ -318,7 +318,7 @@ def load_chroma(checkpoint_info, diffusers_load_config): # triggered by opts.sd_ allow_quant = 'gguf' not in (sd_unet.loaded_unet or '') and (prequantized is None or prequantized == 'none') if (fn is None) or (not os.path.exists(fn) or os.path.isdir(fn)): kwargs = load_quants(kwargs, repo_id or fn, cache_dir=shared.opts.diffusers_dir, allow_quant=allow_quant) - # kwargs = model_quant.create_config(kwargs, allow_quant) + # kwargs = model_quant.create_config(kwargs, allow_quant, modules_to_not_convert=["distilled_guidance_layer"]) if fn.endswith('.safetensors') and os.path.isfile(fn): pipe = diffusers.ChromaPipeline.from_single_file(fn, cache_dir=shared.opts.diffusers_dir, **kwargs, **diffusers_load_config) else: diff --git a/modules/model_chroma_nf4.py b/modules/model_chroma_nf4.py index c72d624ab..fb02ebe05 100644 --- a/modules/model_chroma_nf4.py +++ b/modules/model_chroma_nf4.py @@ -1,5 +1,5 @@ """ -Copied from: https://github.com/huggingface/diffusers/issues/9165 +Copied from: https://github.com/huggingface/diffusers/issues/9165 + adjusted for Chroma by skipping distilled guidance layer """ diff --git a/modules/model_quant.py b/modules/model_quant.py index 8b4dcb8ad..55cdd8c3c 100644 --- a/modules/model_quant.py +++ b/modules/model_quant.py @@ -433,7 +433,7 @@ def optimum_quanto_model(model, op=None, sd_model=None, weights=None, activation from modules import devices, shared quanto = load_quanto('Quantize model: type=Optimum Quanto') global quant_last_model_name, quant_last_model_device # pylint: disable=global-statement - if sd_model is not None and "Flux" in sd_model.__class__.__name__ or "Chroma" in sd_model.__class__.__name__: # LayerNorm is not supported + if sd_model is not None and ("Flux" in sd_model.__class__.__name__ or "Chroma" in sd_model.__class__.__name__): # LayerNorm is not supported exclude_list = ["transformer_blocks.*.norm1.norm", "transformer_blocks.*.norm2", "transformer_blocks.*.norm1_context.norm", "transformer_blocks.*.norm2_context", "single_transformer_blocks.*.norm.norm", "norm_out.norm"] if "Chroma" in sd_model.__class__.__name__: # we ignore the distilled guidance layer because it degrades quality too much diff --git a/modules/processing_vae.py b/modules/processing_vae.py index 6feaec0cd..b1af54d69 100644 --- a/modules/processing_vae.py +++ b/modules/processing_vae.py @@ -298,9 +298,9 @@ def vae_decode(latents, model, output_type='np', vae_type='Full', width=None, he latent_num_frames = (frames - 1) // model.vae_temporal_compression_ratio + 1 latents = model._unpack_latents(latents.unsqueeze(0), latent_num_frames, height // 32, width // 32, model.transformer_spatial_patch_size, model.transformer_temporal_patch_size) # pylint: disable=protected-access latents = model._denormalize_latents(latents, model.vae.latents_mean, model.vae.latents_std, model.vae.config.scaling_factor) # pylint: disable=protected-access - if hasattr(model, '_unpack_latents') and hasattr(model, "vae_scale_factor") and width is not None and height is not None: # FLUX + if hasattr(model, '_unpack_latents') and hasattr(model, "vae_scale_factor") and width is not None and height is not None and latents.ndim == 3: # FLUX latents = model._unpack_latents(latents, height, width, model.vae_scale_factor) # pylint: disable=protected-access - if len(latents.shape) == 3: # lost a batch dim in hires + if latents.ndim == 3: # lost a batch dim in hires latents = latents.unsqueeze(0) if latents.shape[-1] <= 4: # not a latent, likely an image decoded = latents.float().cpu().numpy() diff --git a/modules/sd_detect.py b/modules/sd_detect.py index 414348b16..8f10bb31d 100644 --- a/modules/sd_detect.py +++ b/modules/sd_detect.py @@ -94,7 +94,7 @@ def detect_pipeline(f: str, op: str = 'model', warning=True, quiet=False): guess = 'HiDream' if 'chroma' in f.lower(): guess = 'Chroma' - if 'flux' in f.lower() or 'flex.1' in f.lower() or 'lodestones' in f.lower(): + if 'flux' in f.lower() or 'flex.1' in f.lower(): guess = 'FLUX' if size > 11000 and size < 16000: warn(f'Model detected as FLUX UNET model, but attempting to load a base model: {op}={f} size={size} MB')