Scheduler interface implementation. Implement pre-existing linear scheduler only. (#258)

Co-authored-by: Anthony Wu <pls-file-gh-issue@users.noreply.github.com>
Co-authored-by: filipstrand <strand.filip@gmail.com>
This commit is contained in:
Anthony Wu 2025-10-12 03:20:40 -07:00 committed by GitHub
parent 1278bafc53
commit 6947a23d6c
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
36 changed files with 394 additions and 54 deletions

View File

@ -65,6 +65,7 @@ classifiers = [
[project.optional-dependencies] [project.optional-dependencies]
dev = [ dev = [
"matplotlib>3.10,<4.0",
"pytest>=8.3.0,<9.0", "pytest>=8.3.0,<9.0",
"pytest-timer>=1.0,<2.0", "pytest-timer>=1.0,<2.0",
"mlx==0.27.1", # Used ONLY during test runs to ensure deterministic test results "mlx==0.27.1", # Used ONLY during test runs to ensure deterministic test results
@ -173,7 +174,7 @@ disable_error_code = [
"union-attr", "union-attr",
"return", "return",
"return-value", "return-value",
"var-annotated" "var-annotated",
] ]
error_summary = true error_summary = true
ignore_missing_imports = true ignore_missing_imports = true

View File

@ -22,6 +22,7 @@ class Config:
redux_image_strengths: list[float] | None = None, redux_image_strengths: list[float] | None = None,
masked_image_path: Path | None = None, masked_image_path: Path | None = None,
controlnet_strength: float | None = None, controlnet_strength: float | None = None,
scheduler: str = "linear",
): ):
if width % 16 != 0 or height % 16 != 0: if width % 16 != 0 or height % 16 != 0:
log.warning("Width and height should be multiples of 16. Rounding down.") log.warning("Width and height should be multiples of 16. Rounding down.")
@ -36,3 +37,4 @@ class Config:
self.redux_image_strengths = redux_image_strengths self.redux_image_strengths = redux_image_strengths
self.masked_image_path = masked_image_path self.masked_image_path = masked_image_path
self.controlnet_strength = controlnet_strength self.controlnet_strength = controlnet_strength
self.scheduler_str = scheduler

View File

@ -2,10 +2,11 @@ import logging
from pathlib import Path from pathlib import Path
import mlx.core as mx import mlx.core as mx
import numpy as np
from mflux.config.config import Config from mflux.config.config import Config
from mflux.config.model_config import ModelConfig from mflux.config.model_config import ModelConfig
from mflux.schedulers import SCHEDULER_REGISTRY, try_import_external_scheduler
from mflux.schedulers.linear_scheduler import LinearScheduler
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@ -18,7 +19,7 @@ class RuntimeConfig:
): ):
self.config = config self.config = config
self.model_config = model_config self.model_config = model_config
self.sigmas = self._create_sigmas(config, model_config) self._scheduler = None
@property @property
def height(self) -> int: def height(self) -> int:
@ -96,27 +97,20 @@ class RuntimeConfig:
return None return None
@staticmethod @property
def _create_sigmas(config: Config, model_config: ModelConfig) -> mx.array: def scheduler(self):
sigmas = RuntimeConfig._create_sigmas_values(config.num_inference_steps) if self._scheduler is not None:
if model_config.requires_sigma_shift: return self._scheduler
sigmas = RuntimeConfig._shift_sigmas(sigmas=sigmas, width=config.width, height=config.height)
return sigmas
@staticmethod if self.config.scheduler_str == "linear":
def _create_sigmas_values(num_inference_steps: int) -> mx.array: self._scheduler = LinearScheduler(self)
sigmas = np.linspace(1.0, 1 / num_inference_steps, num_inference_steps) elif (registered_scheduler := SCHEDULER_REGISTRY.get(self.config.scheduler_str, None)) is not None:
sigmas = mx.array(sigmas).astype(mx.float32) self._scheduler = registered_scheduler(self)
return mx.concatenate([sigmas, mx.zeros(1)]) elif "." in self.config.scheduler_str:
# this raises ValueError if scheduler is not importable
scheduler_cls = try_import_external_scheduler(self.config.scheduler_str)
self._scheduler = scheduler_cls(self)
else:
raise NotImplementedError(f"The scheduler {self.config.scheduler_str!r} is not implemented by mflux.")
@staticmethod return self._scheduler
def _shift_sigmas(sigmas: mx.array, width: int, height: int) -> mx.array:
y1 = 0.5
x1 = 256
m = (1.15 - y1) / (4096 - x1)
b = y1 - m * x1
mu = m * width * height / 256 + b
mu = mx.array(mu)
shifted_sigmas = mx.exp(mu) / (mx.exp(mu) + (1 / sigmas - 1))
shifted_sigmas[-1] = 0
return shifted_sigmas

View File

@ -51,6 +51,7 @@ def main():
guidance=args.guidance, guidance=args.guidance,
image_path=args.image_path, image_path=args.image_path,
image_strength=args.image_strength, image_strength=args.image_strength,
scheduler=args.scheduler,
), ),
) )
# 4. Save the image # 4. Save the image

View File

@ -48,6 +48,7 @@ def main():
width=args.width, width=args.width,
guidance=args.guidance, guidance=args.guidance,
controlnet_strength=args.controlnet_strength, controlnet_strength=args.controlnet_strength,
scheduler=args.scheduler,
), ),
) )

View File

@ -46,6 +46,7 @@ def main():
guidance=args.guidance, guidance=args.guidance,
image_path=args.image_path, image_path=args.image_path,
depth_image_path=args.depth_image_path, depth_image_path=args.depth_image_path,
scheduler=args.scheduler,
), ),
) )

View File

@ -46,6 +46,7 @@ def main():
guidance=args.guidance, guidance=args.guidance,
image_path=args.image_path, image_path=args.image_path,
masked_image_path=args.masked_image_path, masked_image_path=args.masked_image_path,
scheduler=args.scheduler,
), ),
) )

View File

@ -61,6 +61,7 @@ def main():
guidance=args.guidance, guidance=args.guidance,
image_path=args.person_image, image_path=args.person_image,
masked_image_path=args.person_mask, masked_image_path=args.person_mask,
scheduler=args.scheduler,
), ),
) )

View File

@ -56,6 +56,7 @@ def main():
width=args.width, width=args.width,
guidance=args.guidance, guidance=args.guidance,
image_path=args.image_path, image_path=args.image_path,
scheduler=args.scheduler,
), ),
) )
# 4. Save the image # 4. Save the image

View File

@ -61,6 +61,7 @@ def main():
height=height, height=height,
width=width, width=width,
guidance=args.guidance, guidance=args.guidance,
scheduler=args.scheduler,
), ),
) )

View File

@ -47,6 +47,7 @@ def main():
width=args.width, width=args.width,
guidance=args.guidance, guidance=args.guidance,
image_path=args.image_path, image_path=args.image_path,
scheduler=args.scheduler,
), ),
) )

View File

@ -56,6 +56,7 @@ def main():
guidance=args.guidance, guidance=args.guidance,
redux_image_paths=args.redux_image_paths, redux_image_paths=args.redux_image_paths,
redux_image_strengths=redux_image_strengths, redux_image_strengths=redux_image_strengths,
scheduler=args.scheduler,
), ),
) )

View File

@ -150,7 +150,7 @@ class Transformer(nn.Module):
time_text_embed: TimeTextEmbed, time_text_embed: TimeTextEmbed,
config: RuntimeConfig, config: RuntimeConfig,
) -> mx.array: ) -> mx.array:
time_step = config.sigmas[t] * config.num_train_steps time_step = config.scheduler.sigmas[t] * config.num_train_steps
time_step = mx.broadcast_to(time_step, (1,)).astype(config.precision) time_step = mx.broadcast_to(time_step, (1,)).astype(config.precision)
guidance = mx.broadcast_to(config.guidance * config.num_train_steps, (1,)).astype(config.precision) guidance = mx.broadcast_to(config.guidance * config.num_train_steps, (1,)).astype(config.precision)
text_embeddings = time_text_embed(time_step, pooled_prompt_embeds, guidance) text_embeddings = time_text_embed(time_step, pooled_prompt_embeds, guidance)

View File

@ -66,7 +66,7 @@ class Flux1Concept(nn.Module):
img2img=Img2Img( img2img=Img2Img(
vae=self.vae, vae=self.vae,
image_path=config.image_path, image_path=config.image_path,
sigmas=config.sigmas, sigmas=config.scheduler.sigmas,
init_time_step=config.init_time_step, init_time_step=config.init_time_step,
), ),
) )
@ -102,6 +102,9 @@ class Flux1Concept(nn.Module):
attention_data = GenerationAttentionData() attention_data = GenerationAttentionData()
for t in time_steps: for t in time_steps:
try: try:
# Scale model input if needed by the scheduler
latents = config.scheduler.scale_model_input(latents, t)
# 4.t Predict the noise # 4.t Predict the noise
noise, attention = self.transformer( noise, attention = self.transformer(
t=t, t=t,
@ -115,8 +118,11 @@ class Flux1Concept(nn.Module):
attention_data.append(attention) attention_data.append(attention)
# 5.t Take one denoise step # 5.t Take one denoise step
dt = config.sigmas[t + 1] - config.sigmas[t] latents = config.scheduler.step(
latents += noise * dt model_output=noise,
timestep=t,
sample=latents,
)
# (Optional) Call subscribers in-loop # (Optional) Call subscribers in-loop
Callbacks.in_loop( Callbacks.in_loop(

View File

@ -78,7 +78,7 @@ class Flux1ConceptFromImage(nn.Module):
latents = LatentCreator.add_noise_by_interpolation( latents = LatentCreator.add_noise_by_interpolation(
clean=ArrayUtil.pack_latents(latents=encoded_image, height=config.height, width=config.width), clean=ArrayUtil.pack_latents(latents=encoded_image, height=config.height, width=config.width),
noise=static_noise, noise=static_noise,
sigma=config.sigmas[config.init_time_step], sigma=config.scheduler.sigmas[config.init_time_step],
) )
# 2. Encode the main prompt # 2. Encode the main prompt
@ -112,6 +112,9 @@ class Flux1ConceptFromImage(nn.Module):
attention_data = GenerationAttentionData() attention_data = GenerationAttentionData()
for t in time_steps: for t in time_steps:
try: try:
# Scale model input if needed by the scheduler
latents = config.scheduler.scale_model_input(latents, t)
# 4.t Predict the noise (we don't use the noise, only the attention) # 4.t Predict the noise (we don't use the noise, only the attention)
_, attention = self.transformer( _, attention = self.transformer(
t=t, t=t,
@ -128,7 +131,7 @@ class Flux1ConceptFromImage(nn.Module):
latents = LatentCreator.add_noise_by_interpolation( latents = LatentCreator.add_noise_by_interpolation(
clean=ArrayUtil.pack_latents(latents=encoded_image, height=config.height, width=config.width), clean=ArrayUtil.pack_latents(latents=encoded_image, height=config.height, width=config.width),
noise=static_noise, noise=static_noise,
sigma=config.sigmas[t + 1], sigma=config.scheduler.sigmas[t + 1],
) )
# (Optional) Call subscribers in-loop # (Optional) Call subscribers in-loop

View File

@ -96,6 +96,9 @@ class Flux1Controlnet(nn.Module):
for t in time_steps: for t in time_steps:
try: try:
# Scale model input if needed by the scheduler
latents = config.scheduler.scale_model_input(latents, t)
# 4.t Compute controlnet samples # 4.t Compute controlnet samples
controlnet_block_samples, controlnet_single_block_samples = self.transformer_controlnet( controlnet_block_samples, controlnet_single_block_samples = self.transformer_controlnet(
t=t, t=t,
@ -118,8 +121,11 @@ class Flux1Controlnet(nn.Module):
) )
# 6.t Take one denoise step # 6.t Take one denoise step
dt = config.sigmas[t + 1] - config.sigmas[t] latents = config.scheduler.step(
latents += noise * dt model_output=noise,
timestep=t,
sample=latents,
)
# (Optional) Call subscribers in-loop # (Optional) Call subscribers in-loop
Callbacks.in_loop( Callbacks.in_loop(

View File

@ -92,6 +92,9 @@ class Flux1Depth(nn.Module):
for t in time_steps: for t in time_steps:
try: try:
# Scale model input if needed by the scheduler
latents = config.scheduler.scale_model_input(latents, t)
# 4.t Concatenate the updated latents with the (static) depth map # 4.t Concatenate the updated latents with the (static) depth map
hidden_states = mx.concatenate([latents, static_depth_map], axis=-1) hidden_states = mx.concatenate([latents, static_depth_map], axis=-1)
@ -105,8 +108,11 @@ class Flux1Depth(nn.Module):
) )
# 6.t Take one denoise step # 6.t Take one denoise step
dt = config.sigmas[t + 1] - config.sigmas[t] latents = config.scheduler.step(
latents += noise * dt model_output=noise,
timestep=t,
sample=latents,
)
# (Optional) Call subscribers in-loop # (Optional) Call subscribers in-loop
Callbacks.in_loop( Callbacks.in_loop(

View File

@ -49,7 +49,7 @@ class DreamBoothLoss:
latents_t = LatentCreator.add_noise_by_interpolation( latents_t = LatentCreator.add_noise_by_interpolation(
clean=clean_image, clean=clean_image,
noise=pure_noise, noise=pure_noise,
sigma=config.sigmas[t], sigma=config.scheduler.sigmas[t],
) )
# Predict the noise from timestep t # Predict the noise from timestep t

View File

@ -89,6 +89,9 @@ class Flux1Fill(nn.Module):
for t in time_steps: for t in time_steps:
try: try:
# Scale model input if needed by the scheduler
latents = config.scheduler.scale_model_input(latents, t)
# 4.t Concatenate the updated latents with the static masked latents # 4.t Concatenate the updated latents with the static masked latents
hidden_states = mx.concatenate([latents, static_masked_latents], axis=-1) hidden_states = mx.concatenate([latents, static_masked_latents], axis=-1)
@ -102,8 +105,11 @@ class Flux1Fill(nn.Module):
) )
# 6.t Take one denoise step # 6.t Take one denoise step
dt = config.sigmas[t + 1] - config.sigmas[t] latents = config.scheduler.step(
latents += noise * dt model_output=noise,
timestep=t,
sample=latents,
)
# (Optional) Call subscribers in-loop # (Optional) Call subscribers in-loop
Callbacks.in_loop( Callbacks.in_loop(

View File

@ -89,6 +89,9 @@ class Flux1InContextDev(nn.Module):
for t in time_steps: for t in time_steps:
try: try:
# Scale model input if needed by the scheduler
latents = config.scheduler.scale_model_input(latents, t)
# 4.t Predict the noise # 4.t Predict the noise
noise = self.transformer( noise = self.transformer(
t=t, t=t,
@ -99,8 +102,11 @@ class Flux1InContextDev(nn.Module):
) )
# 5.t Take one denoise step and update latents # 5.t Take one denoise step and update latents
dt = config.sigmas[t + 1] - config.sigmas[t] latents = config.scheduler.step(
latents += noise * dt model_output=noise,
timestep=t,
sample=latents,
)
# 6.t Override the left-hand side of latents by linearly interpolating between latents and static noise # 6.t Override the left-hand side of latents by linearly interpolating between latents and static noise
latents = Flux1InContextDev._update_latents( latents = Flux1InContextDev._update_latents(
@ -109,6 +115,7 @@ class Flux1InContextDev(nn.Module):
latents=latents, latents=latents,
encoded_image=encoded_image, encoded_image=encoded_image,
static_noise=static_noise, static_noise=static_noise,
sigmas=config.scheduler.sigmas,
) )
# (Optional) Call subscribers in-loop # (Optional) Call subscribers in-loop
@ -180,6 +187,7 @@ class Flux1InContextDev(nn.Module):
latents: mx.array, latents: mx.array,
encoded_image: mx.array, encoded_image: mx.array,
static_noise: mx.array, static_noise: mx.array,
sigmas: mx.array,
) -> mx.array: ) -> mx.array:
# 1. Unpack the latents # 1. Unpack the latents
unpacked = ArrayUtil.unpack_latents(latents=latents, height=config.height, width=config.width) unpacked = ArrayUtil.unpack_latents(latents=latents, height=config.height, width=config.width)
@ -192,7 +200,7 @@ class Flux1InContextDev(nn.Module):
unpacked[:, :, :, 0:latent_width] = LatentCreator.add_noise_by_interpolation( unpacked[:, :, :, 0:latent_width] = LatentCreator.add_noise_by_interpolation(
clean=encoded_image[:, :, :, 0:latent_width], clean=encoded_image[:, :, :, 0:latent_width],
noise=unpacked_static_noise[:, :, :, 0:latent_width], noise=unpacked_static_noise[:, :, :, 0:latent_width],
sigma=config.sigmas[t + 1], sigma=sigmas[t + 1],
) )
# 4. Repack the latents # 4. Repack the latents

View File

@ -99,6 +99,9 @@ class Flux1InContextFill(nn.Module):
for t in time_steps: for t in time_steps:
try: try:
# Scale model input if needed by the scheduler
latents = config.scheduler.scale_model_input(latents, t)
# 4.t Concatenate the updated latents with the static masked latents # 4.t Concatenate the updated latents with the static masked latents
hidden_states = mx.concatenate([latents, static_masked_latents], axis=-1) hidden_states = mx.concatenate([latents, static_masked_latents], axis=-1)
@ -112,8 +115,11 @@ class Flux1InContextFill(nn.Module):
) )
# 6.t Take one denoise step # 6.t Take one denoise step
dt = config.sigmas[t + 1] - config.sigmas[t] latents = config.scheduler.step(
latents += noise * dt model_output=noise,
timestep=t,
sample=latents,
)
# (Optional) Call subscribers in-loop # (Optional) Call subscribers in-loop
Callbacks.in_loop( Callbacks.in_loop(

View File

@ -88,6 +88,9 @@ class Flux1Kontext(nn.Module):
for t in time_steps: for t in time_steps:
try: try:
# Scale model input if needed by the scheduler
latents = config.scheduler.scale_model_input(latents, t)
# 4.t Concatenate the updated latents with the static image latents # 4.t Concatenate the updated latents with the static image latents
hidden_states = mx.concatenate([latents, static_image_latents], axis=1) hidden_states = mx.concatenate([latents, static_image_latents], axis=1)
@ -105,8 +108,11 @@ class Flux1Kontext(nn.Module):
noise = noise[:, : latents.shape[1]] noise = noise[:, : latents.shape[1]]
# 7.t Take one denoise step # 7.t Take one denoise step
dt = config.sigmas[t + 1] - config.sigmas[t] latents = config.scheduler.step(
latents += noise * dt model_output=noise,
timestep=t,
sample=latents,
)
# (Optional) Call subscribers in-loop # (Optional) Call subscribers in-loop
Callbacks.in_loop( Callbacks.in_loop(

View File

@ -92,6 +92,9 @@ class Flux1Redux(nn.Module):
for t in time_steps: for t in time_steps:
try: try:
# Scale model input if needed by the scheduler
latents = runtime_config.scheduler.scale_model_input(latents, t)
# 3.t Predict the noise # 3.t Predict the noise
noise = self.transformer( noise = self.transformer(
t=t, t=t,
@ -102,8 +105,11 @@ class Flux1Redux(nn.Module):
) )
# 4.t Take one denoise step # 4.t Take one denoise step
dt = runtime_config.sigmas[t + 1] - runtime_config.sigmas[t] latents = runtime_config.scheduler.step(
latents += noise * dt model_output=noise,
timestep=t,
sample=latents,
)
# (Optional) Call subscribers in-loop # (Optional) Call subscribers in-loop
Callbacks.in_loop( Callbacks.in_loop(

View File

@ -63,7 +63,7 @@ class Flux1(nn.Module):
img2img=Img2Img( img2img=Img2Img(
vae=self.vae, vae=self.vae,
image_path=config.image_path, image_path=config.image_path,
sigmas=config.sigmas, sigmas=config.scheduler.sigmas,
init_time_step=config.init_time_step, init_time_step=config.init_time_step,
), ),
) )
@ -88,6 +88,9 @@ class Flux1(nn.Module):
for t in time_steps: for t in time_steps:
try: try:
# Scale model input if needed by the scheduler
latents = config.scheduler.scale_model_input(latents, t)
# 3.t Predict the noise # 3.t Predict the noise
noise = self.transformer( noise = self.transformer(
t=t, t=t,
@ -98,8 +101,11 @@ class Flux1(nn.Module):
) )
# 4.t Take one denoise step # 4.t Take one denoise step
dt = config.sigmas[t + 1] - config.sigmas[t] latents = config.scheduler.step(
latents += noise * dt model_output=noise,
timestep=t,
sample=latents,
)
# (Optional) Call subscribers in-loop # (Optional) Call subscribers in-loop
Callbacks.in_loop( Callbacks.in_loop(

View File

@ -89,7 +89,7 @@ class QwenTransformer(nn.Module):
time_text_embed: QwenTimeTextEmbed, time_text_embed: QwenTimeTextEmbed,
config: RuntimeConfig, config: RuntimeConfig,
) -> mx.array: ) -> mx.array:
time_step = config.sigmas[t] time_step = config.scheduler.sigmas[t]
batch = hidden_states.shape[0] batch = hidden_states.shape[0]
timestep = mx.array(np.full((batch,), time_step, dtype=np.float32)) timestep = mx.array(np.full((batch,), time_step, dtype=np.float32))
text_embeddings = time_text_embed(timestep, hidden_states) text_embeddings = time_text_embed(timestep, hidden_states)

View File

@ -67,7 +67,7 @@ class QwenImage(nn.Module):
width=runtime_config.width, width=runtime_config.width,
img2img=Img2Img( img2img=Img2Img(
vae=self.vae, vae=self.vae,
sigmas=runtime_config.sigmas, sigmas=runtime_config.scheduler.sigmas,
init_time_step=runtime_config.init_time_step, init_time_step=runtime_config.init_time_step,
image_path=runtime_config.image_path, image_path=runtime_config.image_path,
), ),
@ -92,6 +92,9 @@ class QwenImage(nn.Module):
for t in time_steps: for t in time_steps:
try: try:
# Scale model input if needed by the scheduler
latents = runtime_config.scheduler.scale_model_input(latents, t)
# 3. Predict the noise # 3. Predict the noise
noise = self.transformer( noise = self.transformer(
t=t, t=t,
@ -110,8 +113,11 @@ class QwenImage(nn.Module):
guided_noise = QwenImage._compute_guided_noise(noise, noise_negative, runtime_config.guidance) guided_noise = QwenImage._compute_guided_noise(noise, noise_negative, runtime_config.guidance)
# 4.t Take one denoise step # 4.t Take one denoise step
dt = runtime_config.sigmas[t + 1] - runtime_config.sigmas[t] latents = runtime_config.scheduler.step(
latents = latents + guided_noise * dt model_output=guided_noise,
timestep=t,
sample=latents,
)
# (Optional) Call subscribers in-loop # (Optional) Call subscribers in-loop
Callbacks.in_loop( Callbacks.in_loop(

View File

@ -0,0 +1,68 @@
from .linear_scheduler import LinearScheduler
__all__ = [
"LinearScheduler",
]
class SchedulerModuleNotFound(ValueError): ...
class SchedulerClassNotFound(ValueError): ...
class InvalidSchedulerType(TypeError): ...
SCHEDULER_REGISTRY = {
"linear": LinearScheduler,
"LinearScheduler": LinearScheduler,
}
def register_contrib(scheduler_object, scheduler_name=None):
if scheduler_name is None:
scheduler_name = scheduler_object.__name__
SCHEDULER_REGISTRY[scheduler_name] = scheduler_object
def try_import_external_scheduler(scheduler_object_path: str):
import importlib
import inspect
from .base_scheduler import BaseScheduler
try:
last_dot_index = scheduler_object_path.rfind(".")
if last_dot_index < 0:
raise SchedulerModuleNotFound(
f"Invalid scheduler path format: {scheduler_object_path!r}. "
"Expected format: some_library.some_package.maybe_sub_package.YourScheduler"
)
module_name_str = scheduler_object_path[:last_dot_index]
scheduler_class_name = scheduler_object_path[last_dot_index + 1 :]
module = importlib.import_module(module_name_str)
except ImportError:
raise SchedulerModuleNotFound(scheduler_object_path)
try:
# Step 2: Get the object from the module using its string name
SchedulerClass = getattr(module, scheduler_class_name)
except AttributeError:
raise SchedulerClassNotFound(scheduler_object_path)
# Step 3: Validate that it's a class and inherits from BaseScheduler
if not inspect.isclass(SchedulerClass):
raise InvalidSchedulerType(
f"{scheduler_object_path!r} is not a class. " f"Schedulers must be classes inheriting from BaseScheduler."
)
if not issubclass(SchedulerClass, BaseScheduler):
raise InvalidSchedulerType(
f"{scheduler_object_path!r} does not inherit from BaseScheduler. "
f"All schedulers must inherit from mflux.schedulers.BaseScheduler."
)
return SchedulerClass

View File

@ -0,0 +1,26 @@
from abc import ABC, abstractmethod
import mlx.core as mx
class BaseScheduler(ABC):
"""
Abstract base class for all schedulers.
"""
@property
@abstractmethod
def sigmas(self) -> mx.array:
"""
The sigma schedule for the diffusion process.
"""
...
@abstractmethod
def step(self, model_output: mx.array, timestep: int, sample: mx.array, **kwargs) -> mx.array: ...
def scale_model_input(self, latents: mx.array, t: int) -> mx.array:
"""
Scale the denoising model input. By default, no scaling applied.
"""
return latents

View File

@ -0,0 +1,49 @@
from typing import TYPE_CHECKING
import mlx.core as mx
if TYPE_CHECKING:
from mflux.config.runtime_config import RuntimeConfig
from mflux.schedulers.base_scheduler import BaseScheduler
class LinearScheduler(BaseScheduler):
"""
Linear scheduler - the default/classic scheduler used in mflux.
Creates a linear schedule from 1.0 to 1/num_steps.
"""
def __init__(self, runtime_config: "RuntimeConfig"):
self.runtime_config = runtime_config
self._sigmas = self._get_sigmas()
@property
def sigmas(self) -> mx.array:
return self._sigmas
def _get_sigmas(self) -> mx.array:
model_config = self.runtime_config.model_config
sigmas = mx.linspace(
1.0,
1.0 / self.runtime_config.num_inference_steps,
self.runtime_config.num_inference_steps,
)
sigmas = mx.array(sigmas).astype(mx.float32)
sigmas = mx.concatenate([sigmas, mx.zeros(1)])
if model_config.requires_sigma_shift:
y1 = 0.5
x1 = 256
m = (1.15 - y1) / (4096 - x1)
b = y1 - m * x1
mu = m * self.runtime_config.width * self.runtime_config.height / 256 + b
mu = mx.array(mu)
shifted_sigmas = mx.exp(mu) / (mx.exp(mu) + (1 / sigmas - 1))
shifted_sigmas[-1] = 0
return shifted_sigmas
else:
return sigmas
def step(self, model_output: mx.array, timestep: int, sample: mx.array, **kwargs) -> mx.array:
dt = self._sigmas[timestep + 1] - self._sigmas[timestep]
return sample + model_output * dt

View File

@ -108,6 +108,7 @@ class CommandLineParser(argparse.ArgumentParser):
self.add_argument("--negative-prompt", type=str, default="", help="The negative prompt to guide what the model should not generate.") self.add_argument("--negative-prompt", type=str, default="", help="The negative prompt to guide what the model should not generate.")
self.add_argument("--seed", type=int, default=None, nargs='+', help="Specify 1+ Entropy Seeds (Default is 1 time-based random-seed)") self.add_argument("--seed", type=int, default=None, nargs='+', help="Specify 1+ Entropy Seeds (Default is 1 time-based random-seed)")
self.add_argument("--auto-seeds", type=int, default=-1, help="Auto generate N Entropy Seeds (random ints between 0 and 1 billion") self.add_argument("--auto-seeds", type=int, default=-1, help="Auto generate N Entropy Seeds (random ints between 0 and 1 billion")
self.add_argument("--scheduler", type=str, default="linear", help="Choose from implemented schedulers (linear only for now). Or bring your own: 'your_package.some_module.FooScheduler'")
self._add_image_generator_common_arguments(supports_dimension_scale_factor=supports_dimension_scale_factor) self._add_image_generator_common_arguments(supports_dimension_scale_factor=supports_dimension_scale_factor)
if supports_metadata_config: if supports_metadata_config:
self.add_metadata_config() self.add_metadata_config()

View File

@ -674,6 +674,21 @@ def test_fill_default_guidance():
assert args.guidance == ui_defaults.DEFAULT_DEV_FILL_GUIDANCE assert args.guidance == ui_defaults.DEFAULT_DEV_FILL_GUIDANCE
def test_scheduler_by_name(mflux_generate_parser, mflux_generate_minimal_model_argv):
with patch("sys.argv", mflux_generate_minimal_model_argv + ["--scheduler", "linear"]):
args = mflux_generate_parser.parse_args()
# we just want the arg to be accepted
assert args.scheduler == "linear"
with patch("sys.argv", mflux_generate_minimal_model_argv + ["--scheduler", "foobar"]):
# let the CLI will accept "foobar" - later RuntimeConfig will validate
args = mflux_generate_parser.parse_args()
with patch("sys.argv", mflux_generate_minimal_model_argv + ["--scheduler", "mflux.someplugin.SchedulerBaz"]):
# let the CLI will accept external import paths - later RuntimeConfig will validate
args = mflux_generate_parser.parse_args()
def test_save_depth_args(mflux_save_depth_parser, mflux_save_depth_minimal_argv): def test_save_depth_args(mflux_save_depth_parser, mflux_save_depth_minimal_argv):
# Test required arguments # Test required arguments
with patch("sys.argv", mflux_save_depth_minimal_argv): with patch("sys.argv", mflux_save_depth_minimal_argv):

Binary file not shown.

Before

Width:  |  Height:  |  Size: 46 KiB

After

Width:  |  Height:  |  Size: 46 KiB

View File

@ -0,0 +1,12 @@
import mlx.core as mx
import pytest
from mflux.schedulers.base_scheduler import BaseScheduler
def test_base_scheduler_is_abstract():
"""
Test that BaseScheduler cannot be instantiated directly.
"""
with pytest.raises(TypeError):
BaseScheduler()

View File

@ -0,0 +1,79 @@
import io
import mlx.core as mx
import numpy as np
import pytest
from mflux.config.config import Config
from mflux.config.model_config import ModelConfig
from mflux.config.runtime_config import RuntimeConfig
from mflux.schedulers import try_import_external_scheduler
from mflux.schedulers.linear_scheduler import LinearScheduler
def test_linear_scheduler_import_by_name():
assert try_import_external_scheduler("mflux.schedulers.linear_scheduler.LinearScheduler") == LinearScheduler
@pytest.fixture
def test_runtime_config():
return RuntimeConfig(
Config(
# these are the only attributes relevant to schedulers
num_inference_steps=14,
width=1024,
height=1024,
scheduler="linear",
),
ModelConfig.dev(), # requires_sigma_shift=True
)
def test_linear_scheduler_initialization(test_runtime_config):
"""
Test the initialization of the LinearScheduler.
"""
scheduler = LinearScheduler(runtime_config=test_runtime_config)
assert scheduler.sigmas is not None
assert isinstance(scheduler.sigmas, mx.array)
assert len(scheduler.sigmas) > 0
def test_linear_scheduler_sigmas_property_no_shift(test_runtime_config):
"""
Test the sigmas property of the LinearScheduler without sigma shift.
"""
scheduler = LinearScheduler(runtime_config=test_runtime_config)
expected_sigmas_from_mflux_0_9_0 = mx.array(
np.load(
io.BytesIO(
bytes.fromhex(
# see: https://gist.github.com/anthonywu/2832147ff5f5f50c81df4d13152d2bed
"934e554d505901003e007b276465736372273a20273c6634272c2027666f727472616e5f6f72646572273a20547275652c20277368617065273a202831352c20297d20202020200a0000803fb7e9793fd92a733f7aa66b3fab38633f30b4593f59df4e3f4f6f423f4b01343f1b10233fc3e30e3f5cedec3e0a90b03e4025483e00000000"
)
)
)
)
assert mx.allclose(scheduler.sigmas, expected_sigmas_from_mflux_0_9_0)
assert scheduler.sigmas.shape == (test_runtime_config.num_inference_steps + 1,)
def test_linear_scheduler_sigmas_property_with_shift(test_runtime_config):
"""
Test the sigmas property of the LinearScheduler with sigma shift.
"""
test_runtime_config.model_config = ModelConfig.schnell() # requires_sigma_shift=True
scheduler = LinearScheduler(runtime_config=test_runtime_config)
expected_sigmas_from_mflux_0_9_0 = mx.array(
np.load(
io.BytesIO(
bytes.fromhex(
# see: https://gist.github.com/anthonywu/2832147ff5f5f50c81df4d13152d2bed
"934e554d505901003e007b276465736372273a20273c6634272c2027666f727472616e5f6f72646572273a20547275652c20277368617065273a202831352c20297d20202020200a0000803fdbb66d3fb76d5b3f9224493f6edb363f4992243f2549123f0000003fb76ddb3e6edbb63e2549923eb76d5b3e2549123e2549923d00000000"
)
)
)
)
assert mx.allclose(scheduler.sigmas, expected_sigmas_from_mflux_0_9_0)
assert scheduler.sigmas.shape == (test_runtime_config.num_inference_steps + 1,)

View File

@ -0,0 +1,17 @@
import pytest
import mflux.schedulers
def test_scheduler_by_path():
mflux.schedulers.try_import_external_scheduler("mflux.schedulers.linear_scheduler.LinearScheduler")
def test_scheduler_bad_module():
with pytest.raises(mflux.schedulers.SchedulerModuleNotFound):
mflux.schedulers.try_import_external_scheduler("someone.other.project.BarScheduler")
def test_scheduler_bad_classname():
with pytest.raises(mflux.schedulers.SchedulerClassNotFound):
mflux.schedulers.try_import_external_scheduler("mflux.schedulers.linear_scheduler.FooBarScheduler")

2
uv.lock generated
View File

@ -786,6 +786,7 @@ dependencies = [
[package.optional-dependencies] [package.optional-dependencies]
dev = [ dev = [
{ name = "matplotlib" },
{ name = "mlx" }, { name = "mlx" },
{ name = "pytest" }, { name = "pytest" },
{ name = "pytest-timer" }, { name = "pytest-timer" },
@ -797,6 +798,7 @@ requires-dist = [
{ name = "filelock", specifier = ">=3.18.0" }, { name = "filelock", specifier = ">=3.18.0" },
{ name = "huggingface-hub", specifier = ">=0.24.5,<1.0" }, { name = "huggingface-hub", specifier = ">=0.24.5,<1.0" },
{ name = "matplotlib", specifier = ">=3.9.2,<4.0" }, { name = "matplotlib", specifier = ">=3.9.2,<4.0" },
{ name = "matplotlib", marker = "extra == 'dev'", specifier = ">3.10,<4.0" },
{ name = "mlx", specifier = ">=0.27.0,<0.28.0" }, { name = "mlx", specifier = ">=0.27.0,<0.28.0" },
{ name = "mlx", marker = "extra == 'dev'", specifier = "==0.27.1" }, { name = "mlx", marker = "extra == 'dev'", specifier = "==0.27.1" },
{ name = "numpy", specifier = ">=2.0.1,<3.0" }, { name = "numpy", specifier = ">=2.0.1,<3.0" },