Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions modules/modelSampler/AnimaSampler.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand All @@ -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
Expand Down Expand Up @@ -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,
)

Expand Down
4 changes: 4 additions & 0 deletions modules/modelSampler/ChromaSampler.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand All @@ -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
Expand Down Expand Up @@ -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,
)
Expand Down
4 changes: 4 additions & 0 deletions modules/modelSampler/ErnieSampler.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand All @@ -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
Expand Down Expand Up @@ -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,
)

Expand Down
6 changes: 5 additions & 1 deletion modules/modelSampler/Flux2Sampler.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
import copy
import inspect
import math
from collections.abc import Callable

from modules.model.Flux2Model import Flux2Model
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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,
)
Expand Down
8 changes: 6 additions & 2 deletions modules/modelSampler/FluxSampler.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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 = "",
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand All @@ -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,
Expand Down
4 changes: 4 additions & 0 deletions modules/modelSampler/HunyuanVideoSampler.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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
Expand Down Expand Up @@ -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,
Expand Down
4 changes: 3 additions & 1 deletion modules/modelSampler/Krea2Sampler.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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))
Expand Down Expand Up @@ -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,
)

Expand Down
4 changes: 3 additions & 1 deletion modules/modelSampler/QwenSampler.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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))
Expand Down Expand Up @@ -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,
)

Expand Down
4 changes: 4 additions & 0 deletions modules/modelSampler/ZImageSampler.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand All @@ -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
Expand Down Expand Up @@ -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,
)

Expand Down
19 changes: 13 additions & 6 deletions modules/ui/BaseSampleFrameView.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
3 changes: 3 additions & 0 deletions modules/util/config/SampleConfig.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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))
Expand Down
10 changes: 10 additions & 0 deletions modules/util/enum/ModelType.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down