diff --git a/modules/modelSampler/AnimaSampler.py b/modules/modelSampler/AnimaSampler.py index d5de18ac1..41872c7cc 100644 --- a/modules/modelSampler/AnimaSampler.py +++ b/modules/modelSampler/AnimaSampler.py @@ -48,6 +48,7 @@ def __sample_base( diffusion_steps: int, cfg_scale: float, noise_scheduler: NoiseScheduler, + override_shift: float | None = None, on_update_progress: Callable[[int, int], None] = lambda _, __: None, ) -> ModelSamplerOutput: with self.model.autocast_context: @@ -58,6 +59,8 @@ def __sample_base( generator.manual_seed(seed) noise_scheduler = copy.deepcopy(self.model.noise_scheduler) + if override_shift is not None: + noise_scheduler.set_shift(override_shift) transformer = self.model.transformer vae = self.model.vae @@ -151,6 +154,7 @@ def sample( diffusion_steps=sample_config.diffusion_steps, cfg_scale=sample_config.cfg_scale, noise_scheduler=sample_config.noise_scheduler, + override_shift=sample_config.override_shift, on_update_progress=on_update_progress, ) diff --git a/modules/modelSampler/ChromaSampler.py b/modules/modelSampler/ChromaSampler.py index 23ef4aee5..b188e98be 100644 --- a/modules/modelSampler/ChromaSampler.py +++ b/modules/modelSampler/ChromaSampler.py @@ -45,6 +45,7 @@ def __sample_base( diffusion_steps: int, cfg_scale: float, noise_scheduler: NoiseScheduler, + override_shift: float | None = None, text_encoder_layer_skip: int = 0, on_update_progress: Callable[[int, int], None] = lambda _, __: None, ) -> ModelSamplerOutput: @@ -56,6 +57,8 @@ def __sample_base( generator.manual_seed(seed) noise_scheduler = copy.deepcopy(self.model.noise_scheduler) + if override_shift is not None: + noise_scheduler.set_shift(override_shift) image_processor = self.pipeline.image_processor transformer = self.pipeline.transformer vae = self.pipeline.vae @@ -170,6 +173,7 @@ def sample( diffusion_steps=sample_config.diffusion_steps, cfg_scale=sample_config.cfg_scale, noise_scheduler=sample_config.noise_scheduler, + override_shift=sample_config.override_shift, text_encoder_layer_skip=sample_config.text_encoder_1_layer_skip, on_update_progress=on_update_progress, ) diff --git a/modules/modelSampler/ErnieSampler.py b/modules/modelSampler/ErnieSampler.py index a9cb57e0a..8be78fa0c 100644 --- a/modules/modelSampler/ErnieSampler.py +++ b/modules/modelSampler/ErnieSampler.py @@ -46,6 +46,7 @@ def __sample_base( diffusion_steps: int, cfg_scale: float, noise_scheduler: NoiseScheduler, + override_shift: float | None = None, on_update_progress: Callable[[int, int], None] = lambda _, __: None, ) -> ModelSamplerOutput: with self.model.autocast_context: @@ -56,6 +57,8 @@ def __sample_base( generator.manual_seed(seed) noise_scheduler = copy.deepcopy(self.model.noise_scheduler) + if override_shift is not None: + noise_scheduler.set_shift(override_shift) vae = self.pipeline.vae vae_scale_factor = 8 @@ -145,6 +148,7 @@ def sample( diffusion_steps=sample_config.diffusion_steps, cfg_scale=sample_config.cfg_scale, noise_scheduler=sample_config.noise_scheduler, + override_shift=sample_config.override_shift, on_update_progress=on_update_progress, ) diff --git a/modules/modelSampler/Flux2Sampler.py b/modules/modelSampler/Flux2Sampler.py index 7ecbd5c83..a89e2f397 100644 --- a/modules/modelSampler/Flux2Sampler.py +++ b/modules/modelSampler/Flux2Sampler.py @@ -1,5 +1,6 @@ import copy import inspect +import math from collections.abc import Callable from modules.model.Flux2Model import Flux2Model @@ -48,6 +49,7 @@ def __sample_base( diffusion_steps: int, cfg_scale: float, noise_scheduler: NoiseScheduler, + override_shift: float | None = None, text_encoder_sequence_length: int | None = None, on_update_progress: Callable[[int, int], None] = lambda _, __: None, ) -> ModelSamplerOutput: @@ -90,7 +92,8 @@ def __sample_base( latent_image = self.model.pack_latents(latent_image) image_seq_len = latent_image.shape[1] - mu = compute_empirical_mu(image_seq_len, diffusion_steps) + # the override is a shift factor, the same quantity the other flow-matching samplers pass as log(shift) + mu = math.log(override_shift) if override_shift else compute_empirical_mu(image_seq_len, diffusion_steps) # prepare timesteps #TODO for other models, too? This is different than with sigmas=None @@ -170,6 +173,7 @@ def sample( diffusion_steps=sample_config.diffusion_steps, cfg_scale=sample_config.cfg_scale, noise_scheduler=sample_config.noise_scheduler, + override_shift=sample_config.override_shift, text_encoder_sequence_length=sample_config.text_encoder_1_sequence_length, on_update_progress=on_update_progress, ) diff --git a/modules/modelSampler/FluxSampler.py b/modules/modelSampler/FluxSampler.py index fc2b2e0c8..5bf8e1db5 100644 --- a/modules/modelSampler/FluxSampler.py +++ b/modules/modelSampler/FluxSampler.py @@ -50,6 +50,7 @@ def __sample_base( diffusion_steps: int, cfg_scale: float, noise_scheduler: NoiseScheduler, + override_shift: float | None = None, text_encoder_1_layer_skip: int = 0, text_encoder_2_layer_skip: int = 0, text_encoder_2_sequence_length: int | None = None, @@ -97,7 +98,7 @@ def __sample_base( self.model.train_dtype.torch_dtype() ) - shift = self.model.calculate_timestep_shift(latent_image.shape[-2], latent_image.shape[-1]) + shift = override_shift or self.model.calculate_timestep_shift(latent_image.shape[-2], latent_image.shape[-1]) latent_image = self.model.pack_latents(latent_image) # prepare timesteps @@ -189,6 +190,7 @@ def __sample_inpainting( diffusion_steps: int, cfg_scale: float, noise_scheduler: NoiseScheduler, + override_shift: float | None = None, sample_inpainting: bool = False, base_image_path: str = "", mask_image_path: str = "", @@ -313,7 +315,7 @@ def __sample_inpainting( self.model.train_dtype.torch_dtype() ) - shift = self.model.calculate_timestep_shift(latent_image.shape[-2], latent_image.shape[-1]) + shift = override_shift or self.model.calculate_timestep_shift(latent_image.shape[-2], latent_image.shape[-1]) latent_image = self.model.pack_latents(latent_image) noise_scheduler.set_timesteps(diffusion_steps, device=self.train_device, mu=math.log(shift)) timesteps = noise_scheduler.timesteps @@ -401,6 +403,7 @@ def sample( diffusion_steps=sample_config.diffusion_steps, cfg_scale=sample_config.cfg_scale, noise_scheduler=sample_config.noise_scheduler, + override_shift=sample_config.override_shift, sample_inpainting=sample_config.sample_inpainting, base_image_path=sample_config.base_image_path, mask_image_path=sample_config.mask_image_path, @@ -421,6 +424,7 @@ def sample( diffusion_steps=sample_config.diffusion_steps, cfg_scale=sample_config.cfg_scale, noise_scheduler=sample_config.noise_scheduler, + override_shift=sample_config.override_shift, text_encoder_1_layer_skip=sample_config.text_encoder_1_layer_skip, text_encoder_2_layer_skip=sample_config.text_encoder_2_layer_skip, text_encoder_2_sequence_length=sample_config.text_encoder_2_sequence_length, diff --git a/modules/modelSampler/HunyuanVideoSampler.py b/modules/modelSampler/HunyuanVideoSampler.py index 10b22bfc9..f7049e6b1 100644 --- a/modules/modelSampler/HunyuanVideoSampler.py +++ b/modules/modelSampler/HunyuanVideoSampler.py @@ -47,6 +47,7 @@ def __sample_base( diffusion_steps: int, cfg_scale: float, noise_scheduler: NoiseScheduler, + override_shift: float | None = None, text_encoder_1_layer_skip: int = 0, text_encoder_2_layer_skip: int = 0, transformer_attention_mask: bool = False, @@ -60,6 +61,8 @@ def __sample_base( generator.manual_seed(seed) noise_scheduler = copy.deepcopy(self.model.noise_scheduler) + if override_shift is not None: + noise_scheduler.set_shift(override_shift) video_processor = self.pipeline.video_processor transformer = self.pipeline.transformer vae = self.pipeline.vae @@ -185,6 +188,7 @@ def sample( diffusion_steps=sample_config.diffusion_steps, cfg_scale=sample_config.cfg_scale, noise_scheduler=sample_config.noise_scheduler, + override_shift=sample_config.override_shift, text_encoder_1_layer_skip=sample_config.text_encoder_1_layer_skip, text_encoder_2_layer_skip=sample_config.text_encoder_2_layer_skip, transformer_attention_mask=sample_config.transformer_attention_mask, diff --git a/modules/modelSampler/Krea2Sampler.py b/modules/modelSampler/Krea2Sampler.py index b83205a93..205caa435 100644 --- a/modules/modelSampler/Krea2Sampler.py +++ b/modules/modelSampler/Krea2Sampler.py @@ -48,6 +48,7 @@ def __sample_base( diffusion_steps: int, cfg_scale: float, noise_scheduler: NoiseScheduler, + override_shift: float | None = None, on_update_progress: Callable[[int, int], None] = lambda _, __: None, ) -> ModelSamplerOutput: with self.model.autocast_context: @@ -83,7 +84,7 @@ def __sample_base( dtype=torch.float32, ) - shift = self.model.calculate_timestep_shift(latent_image.shape[-2], latent_image.shape[-1]) + shift = override_shift or self.model.calculate_timestep_shift(latent_image.shape[-2], latent_image.shape[-1]) latent_image = self.model.pack_latents(latent_image) noise_scheduler.set_timesteps(diffusion_steps, device=self.train_device, mu=math.log(shift)) @@ -160,6 +161,7 @@ def sample( diffusion_steps=sample_config.diffusion_steps, cfg_scale=sample_config.cfg_scale, noise_scheduler=sample_config.noise_scheduler, + override_shift=sample_config.override_shift, on_update_progress=on_update_progress, ) diff --git a/modules/modelSampler/QwenSampler.py b/modules/modelSampler/QwenSampler.py index c18eece7e..8d8bd9c91 100644 --- a/modules/modelSampler/QwenSampler.py +++ b/modules/modelSampler/QwenSampler.py @@ -46,6 +46,7 @@ def __sample_base( diffusion_steps: int, cfg_scale: float, noise_scheduler: NoiseScheduler, + override_shift: float | None = None, on_update_progress: Callable[[int, int], None] = lambda _, __: None, ) -> ModelSamplerOutput: with self.model.autocast_context: @@ -82,7 +83,7 @@ def __sample_base( dtype=torch.float32, ) - shift = self.model.calculate_timestep_shift(latent_image.shape[-2], latent_image.shape[-1]) + shift = override_shift or self.model.calculate_timestep_shift(latent_image.shape[-2], latent_image.shape[-1]) latent_image = self.model.pack_latents(latent_image) noise_scheduler.set_timesteps(diffusion_steps, device=self.train_device, mu=math.log(shift)) @@ -170,6 +171,7 @@ def sample( diffusion_steps=sample_config.diffusion_steps, cfg_scale=sample_config.cfg_scale, noise_scheduler=sample_config.noise_scheduler, + override_shift=sample_config.override_shift, on_update_progress=on_update_progress, ) diff --git a/modules/modelSampler/ZImageSampler.py b/modules/modelSampler/ZImageSampler.py index 0e001df2a..638f8a4a5 100644 --- a/modules/modelSampler/ZImageSampler.py +++ b/modules/modelSampler/ZImageSampler.py @@ -45,6 +45,7 @@ def __sample_base( diffusion_steps: int, cfg_scale: float, noise_scheduler: NoiseScheduler, + override_shift: float | None = None, on_update_progress: Callable[[int, int], None] = lambda _, __: None, ) -> ModelSamplerOutput: with self.model.autocast_context: @@ -55,6 +56,8 @@ def __sample_base( generator.manual_seed(seed) noise_scheduler = copy.deepcopy(self.model.noise_scheduler) + if override_shift is not None: + noise_scheduler.set_shift(override_shift) image_processor = self.pipeline.image_processor transformer = self.pipeline.transformer vae = self.pipeline.vae @@ -146,6 +149,7 @@ def sample( diffusion_steps=sample_config.diffusion_steps, cfg_scale=sample_config.cfg_scale, noise_scheduler=sample_config.noise_scheduler, + override_shift=sample_config.override_shift, on_update_progress=on_update_progress, ) diff --git a/modules/ui/BaseSampleFrameView.py b/modules/ui/BaseSampleFrameView.py index eacd8c0ac..da888c6f6 100644 --- a/modules/ui/BaseSampleFrameView.py +++ b/modules/ui/BaseSampleFrameView.py @@ -71,25 +71,32 @@ def build_content(self, top_frame, bottom_frame, ui_state, controller, include_p self.components.label(bottom_frame, 4, 0, "steps:") self.components.entry(bottom_frame, 4, 1, ui_state, "diffusion_steps") + if controller.model_type.supports_sample_shift_override(): + self.components.label(bottom_frame, 5, 0, "override shift:", + tooltip="Timestep shift for this sample, as the multiplicative factor " + "(not mu). Empty uses the model's own shift, which is either " + "derived from the token count or fixed by the scheduler config.") + self.components.entry(bottom_frame, 5, 1, ui_state, "override_shift") + # inpainting if is_inpainting_model: - self.components.label(bottom_frame, 5, 0, "inpainting:", + self.components.label(bottom_frame, 6, 0, "inpainting:", tooltip="Enables inpainting sampling. Only available when sampling from an inpainting model.") - self.components.switch(bottom_frame, 5, 1, ui_state, "sample_inpainting") + self.components.switch(bottom_frame, 6, 1, ui_state, "sample_inpainting") # base image path - self.components.label(bottom_frame, 6, 0, "base image path:", + self.components.label(bottom_frame, 7, 0, "base image path:", tooltip="The base image used when inpainting.") - self.components.path_entry(bottom_frame, 6, 1, ui_state, "base_image_path", + self.components.path_entry(bottom_frame, 7, 1, ui_state, "base_image_path", mode="file", allow_model_files=False, allow_image_files=True, ) # mask image path - self.components.label(bottom_frame, 6, 2, "mask image path:", + self.components.label(bottom_frame, 7, 2, "mask image path:", tooltip="The mask used when inpainting.") - self.components.path_entry(bottom_frame, 6, 3, ui_state, "mask_image_path", + self.components.path_entry(bottom_frame, 7, 3, ui_state, "mask_image_path", mode="file", allow_model_files=False, allow_image_files=True, diff --git a/modules/util/config/SampleConfig.py b/modules/util/config/SampleConfig.py index 9b2b2c0b1..29b6cd769 100644 --- a/modules/util/config/SampleConfig.py +++ b/modules/util/config/SampleConfig.py @@ -17,6 +17,7 @@ def _get_model_defaults(model_type) -> dict: "cfg_scale": 7.0, "noise_scheduler": NoiseScheduler.DDIM, "negative_prompt": "", + "override_shift": None, } if model_type is None: @@ -170,6 +171,7 @@ class SampleConfig(BaseConfig): diffusion_steps: int cfg_scale: float noise_scheduler: NoiseScheduler + override_shift: float | None text_encoder_1_layer_skip: int text_encoder_1_sequence_length: int | None @@ -214,6 +216,7 @@ def default_values(model_type=None): data.append(("diffusion_steps", defaults["diffusion_steps"], int, False)) data.append(("cfg_scale", defaults["cfg_scale"], float, False)) data.append(("noise_scheduler", defaults["noise_scheduler"], NoiseScheduler, False)) + data.append(("override_shift", defaults["override_shift"], float, True)) data.append(("text_encoder_1_layer_skip", 0, int, False)) data.append(("text_encoder_1_sequence_length", None, int, True)) diff --git a/modules/util/enum/ModelType.py b/modules/util/enum/ModelType.py index 8820892e8..17de7d3fa 100644 --- a/modules/util/enum/ModelType.py +++ b/modules/util/enum/ModelType.py @@ -129,6 +129,16 @@ def is_ernie(self): def is_ideogram(self): return self == ModelType.IDEOGRAM_4 + def supports_sample_shift_override(self) -> bool: + return self.is_flux() \ + or self.is_chroma() \ + or self.is_qwen() \ + or self.is_anima() \ + or self.is_krea2() \ + or self.is_hunyuan_video() \ + or self.is_z_image() \ + or self.is_ernie() + def supports_negative_prompt(self) -> bool: # asymmetric dual-network CFG models drive the negative branch from a frozen unconditional network (or an # empty prompt), not a user-supplied negative prompt