Merge pull request #103 from anthonywu/support-hf-models

release 0.6.0 – comprehensive ModelConfig refactor to support compatible HuggingFace dev/schnell models
This commit is contained in:
Filip Strand 2025-01-19 23:13:37 +01:00 committed by GitHub
commit 976573c837
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
13 changed files with 230 additions and 40 deletions

View File

@ -263,9 +263,9 @@ Or, with the correct python environment active, create and run a separate script
from mflux import Flux1, Config from mflux import Flux1, Config
# Load the model # Load the model
flux = Flux1.from_alias( flux = Flux1.from_name(
alias="schnell", # "schnell" or "dev" model_name="schnell", # "schnell" or "dev"
quantize=8, # 4 or 8 quantize=8, # 4 or 8
) )
# Generate an image # Generate an image

View File

@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
[project] [project]
name = "mflux" name = "mflux"
version = "0.5.1" version = "0.6.0"
description = "A MLX port of FLUX based on the Huggingface Diffusers implementation." description = "A MLX port of FLUX based on the Huggingface Diffusers implementation."
readme = "README.md" readme = "README.md"
keywords = ["diffusers", "flux", "mlx"] keywords = ["diffusers", "flux", "mlx"]

View File

@ -1,5 +1,5 @@
from mflux.config.config import Config, ConfigControlnet from mflux.config.config import Config, ConfigControlnet
from mflux.config.model_config import ModelConfig from mflux.config.model_config import ModelConfig, ModelLookup
from mflux.controlnet.flux_controlnet import Flux1Controlnet from mflux.controlnet.flux_controlnet import Flux1Controlnet
from mflux.error.exceptions import StopImageGenerationException from mflux.error.exceptions import StopImageGenerationException
from mflux.flux.flux import Flux1 from mflux.flux.flux import Flux1
@ -11,6 +11,7 @@ __all__ = [
"Config", "Config",
"ConfigControlnet", "ConfigControlnet",
"ModelConfig", "ModelConfig",
"ModelLookup",
"ImageUtil", "ImageUtil",
"StopImageGenerationException", "StopImageGenerationException",
] ]

View File

@ -1,27 +1,98 @@
from enum import Enum import warnings
from dataclasses import dataclass
from typing import Literal
DEFAULT_TRAIN_STEPS = 1000
KNOWN_SEQUENCE_LENGTH_BY_BASE_MODEL = {"dev": 512, "schnell": 256}
class ModelConfig(Enum): class ModelConfigError(ValueError):
FLUX1_DEV = ("black-forest-labs/FLUX.1-dev", "dev", 1000, 512) """User error in model config."""
FLUX1_SCHNELL = ("black-forest-labs/FLUX.1-schnell", "schnell", 1000, 256)
def __init__(
self, class InvalidBaseModel(ModelConfigError):
model_name: str, """Invalid base model, cannot infer model properties."""
alias: str,
num_train_steps: int,
max_sequence_length: int, @dataclass
): class ModelConfig:
self.alias = alias model_name: str
self.model_name = model_name num_train_steps: int
self.num_train_steps = num_train_steps max_sequence_length: int
self.max_sequence_length = max_sequence_length supports_guidance: bool
base_model: str | None
DefaultModelConfigs = {
"dev": ModelConfig(
model_name="black-forest-labs/FLUX.1-dev",
num_train_steps=DEFAULT_TRAIN_STEPS,
max_sequence_length=KNOWN_SEQUENCE_LENGTH_BY_BASE_MODEL["dev"],
supports_guidance=True,
base_model=None,
),
"schnell": ModelConfig(
model_name="black-forest-labs/FLUX.1-schnell",
num_train_steps=DEFAULT_TRAIN_STEPS,
max_sequence_length=KNOWN_SEQUENCE_LENGTH_BY_BASE_MODEL["schnell"],
supports_guidance=False,
base_model=None,
),
}
class ModelLookup:
@staticmethod
def from_alias(alias: str) -> ModelConfig:
warnings.warn(
"from_alias is deprecated and will be removed in a future release. Please use from_name instead.",
DeprecationWarning,
stacklevel=2,
)
return ModelLookup.from_name(model_name=alias, base_model=None)
@staticmethod @staticmethod
def from_alias(alias: str) -> "ModelConfig": def from_name(
try: model_name: str,
for model in ModelConfig: base_model: Literal["dev", "schnell"] | None = None,
if model.alias == alias: ) -> ModelConfig:
return model if model_name in DefaultModelConfigs:
except KeyError: return DefaultModelConfigs[model_name]
raise ValueError(f"'{alias}' is not a valid model")
if all(["dev" not in model_name, "schnell" not in model_name, base_model is None]):
raise ModelConfigError(
"Cannot infer base model and max_sequence_length "
f"from model reference: {model_name!r}. "
"Please specify --base-model [dev | schnell]"
)
if base_model is not None and base_model not in ["dev", "schnell"]:
raise InvalidBaseModel("As of this version, mflux only recognizes base models dev or schnell")
if base_model is None:
# infer base model on apparent model naming
if "dev" in model_name:
base_model = "dev"
elif "schnell" in model_name:
base_model = "schnell"
if base_model == "dev":
supports_guidance = True
max_sequence_length = KNOWN_SEQUENCE_LENGTH_BY_BASE_MODEL["dev"]
elif base_model == "schnell":
supports_guidance = False
max_sequence_length = KNOWN_SEQUENCE_LENGTH_BY_BASE_MODEL["schnell"]
return ModelConfig(
model_name=model_name,
num_train_steps=DEFAULT_TRAIN_STEPS,
max_sequence_length=max_sequence_length,
supports_guidance=supports_guidance,
base_model=base_model,
)
# keep these class members to be backwards compatible with < 0.5.0 ModelConfig Enum implementation
ModelConfig.FLUX1_DEV = DefaultModelConfigs["dev"]
ModelConfig.FLUX1_SCHNELL = DefaultModelConfigs["schnell"]

View File

@ -1,3 +1,4 @@
import warnings
from pathlib import Path from pathlib import Path
import mlx.core as mx import mlx.core as mx
@ -5,7 +6,7 @@ from mlx import nn
from tqdm import tqdm from tqdm import tqdm
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, ModelLookup
from mflux.config.runtime_config import RuntimeConfig from mflux.config.runtime_config import RuntimeConfig
from mflux.error.exceptions import StopImageGenerationException from mflux.error.exceptions import StopImageGenerationException
from mflux.latent_creator.latent_creator import LatentCreator from mflux.latent_creator.latent_creator import LatentCreator
@ -137,8 +138,17 @@ class Flux1(nn.Module):
@staticmethod @staticmethod
def from_alias(alias: str, quantize: int | None = None) -> "Flux1": def from_alias(alias: str, quantize: int | None = None) -> "Flux1":
warnings.warn(
"from_alias is deprecated and will be removed in a future release. Please use from_name instead.",
DeprecationWarning,
stacklevel=2,
)
return Flux1.from_name(model_name=alias, quantize=quantize)
@staticmethod
def from_name(model_name: str, quantize: int | None = None) -> "Flux1":
return Flux1( return Flux1(
model_config=ModelConfig.from_alias(alias), model_config=ModelLookup.from_name(model_name=model_name, base_model=None),
quantize=quantize, quantize=quantize,
) )

View File

@ -1,7 +1,7 @@
import time import time
from pathlib import Path from pathlib import Path
from mflux import Config, Flux1, ModelConfig, StopImageGenerationException from mflux import Config, Flux1, ModelLookup, StopImageGenerationException
from mflux.ui.cli.parsers import CommandLineParser from mflux.ui.cli.parsers import CommandLineParser
@ -17,7 +17,7 @@ def main():
# Load the model # Load the model
flux = Flux1( flux = Flux1(
model_config=ModelConfig.from_alias(args.model), model_config=ModelLookup.from_name(model_name=args.model, base_model=args.base_model),
quantize=args.quantize, quantize=args.quantize,
local_path=args.path, local_path=args.path,
lora_paths=args.lora_paths, lora_paths=args.lora_paths,

View File

@ -1,7 +1,7 @@
import time import time
from pathlib import Path from pathlib import Path
from mflux import ConfigControlnet, Flux1Controlnet, ModelConfig, StopImageGenerationException from mflux import ConfigControlnet, Flux1Controlnet, ModelLookup, StopImageGenerationException
from mflux.ui.cli.parsers import CommandLineParser from mflux.ui.cli.parsers import CommandLineParser
@ -16,7 +16,7 @@ def main():
# Load the model # Load the model
flux = Flux1Controlnet( flux = Flux1Controlnet(
model_config=ModelConfig.from_alias(args.model), model_config=ModelLookup.from_name(model_name=args.model, base_model=args.base_model),
quantize=args.quantize, quantize=args.quantize,
local_path=args.path, local_path=args.path,
lora_paths=args.lora_paths, lora_paths=args.lora_paths,

View File

@ -14,7 +14,7 @@ class TimeTextEmbed(nn.Module):
def __init__(self, model_config: ModelConfig): def __init__(self, model_config: ModelConfig):
super().__init__() super().__init__()
self.text_embedder = TextEmbedder() self.text_embedder = TextEmbedder()
self.guidance_embedder = GuidanceEmbedder() if model_config == ModelConfig.FLUX1_DEV else None self.guidance_embedder = GuidanceEmbedder() if model_config.supports_guidance else None
self.timestep_embedder = TimestepEmbedder() self.timestep_embedder = TimestepEmbedder()
def __call__( def __call__(

View File

@ -57,10 +57,11 @@ class GeneratedImage:
# mflux_version is used by future metadata readers # mflux_version is used by future metadata readers
# to determine supportability of metadata-derived workflows # to determine supportability of metadata-derived workflows
"mflux_version": GeneratedImage.get_version(), "mflux_version": GeneratedImage.get_version(),
"model": str(self.model_config.alias), "model": self.model_config.model_name,
"base_model": str(self.model_config.base_model),
"seed": self.seed, "seed": self.seed,
"steps": self.steps, "steps": self.steps,
"guidance": self.guidance if ModelConfig.FLUX1_DEV else None, # only the dev model supports guidance "guidance": self.guidance if self.model_config.supports_guidance else None,
"precision": str(self.precision), "precision": str(self.precision),
"quantize": self.quantization, "quantize": self.quantization,
"generation_time_seconds": round(self.generation_time, 2), "generation_time_seconds": round(self.generation_time, 2),

View File

@ -1,4 +1,4 @@
from mflux import Flux1, ModelConfig from mflux import Flux1, ModelLookup
from mflux.ui.cli.parsers import CommandLineParser from mflux.ui.cli.parsers import CommandLineParser
@ -11,7 +11,7 @@ def main():
print(f"Saving model {args.model} with quantization level {args.quantize}\n") print(f"Saving model {args.model} with quantization level {args.quantize}\n")
flux = Flux1( flux = Flux1(
model_config=ModelConfig.from_alias(args.model), model_config=ModelLookup.from_name(args.model, base_model=args.base_model),
quantize=args.quantize, quantize=args.quantize,
lora_paths=args.lora_paths, lora_paths=args.lora_paths,
lora_scales=args.lora_scales, lora_scales=args.lora_scales,

View File

@ -6,6 +6,21 @@ from pathlib import Path
from mflux.ui import defaults as ui_defaults from mflux.ui import defaults as ui_defaults
class ModelSpecAction(argparse.Action):
def __call__(self, parser, namespace, values, option_string=None):
if values in ["dev", "schnell"]:
setattr(namespace, self.dest, values)
return
if values.count("/") != 1:
raise argparse.ArgumentError(
self, ('Value must be either "dev", "schnell", or "' f'in format "org/model". Got: {values}')
)
# If we got here, values contains exactly one slash
setattr(namespace, self.dest, values)
# fmt: off # fmt: off
class CommandLineParser(argparse.ArgumentParser): class CommandLineParser(argparse.ArgumentParser):
@ -19,12 +34,12 @@ class CommandLineParser(argparse.ArgumentParser):
def add_model_arguments(self, path_type: t.Literal["load", "save"] = "load", require_model_arg: bool = True) -> None: def add_model_arguments(self, path_type: t.Literal["load", "save"] = "load", require_model_arg: bool = True) -> None:
self.add_argument("--model", "-m", type=str, required=require_model_arg, choices=ui_defaults.MODEL_CHOICES, help=f"The model to use ({' or '.join(ui_defaults.MODEL_CHOICES)}).") self.add_argument("--model", "-m", type=str, required=require_model_arg, action=ModelSpecAction, help=f"The model to use ({' or '.join(ui_defaults.MODEL_CHOICES)} or a compatible huggingface repo_id org/model).")
if path_type == "load": if path_type == "load":
self.add_argument("--path", type=str, default=None, help="Local path for loading a model from disk") self.add_argument("--path", type=str, default=None, help="Local path for loading a model from disk")
else: else:
self.add_argument("--path", type=str, required=True, help="Local path for saving a model to disk.") self.add_argument("--path", type=str, required=True, help="Local path for saving a model to disk.")
self.add_argument("--base-model", type=str, required=False, choices=ui_defaults.MODEL_CHOICES, help="When using a third-party huggingface model, explicitly specify whether the base model is dev or schnell")
self.add_argument("--quantize", "-q", type=int, choices=ui_defaults.QUANTIZE_CHOICES, default=None, help=f"Quantize the model ({' or '.join(map(str, ui_defaults.QUANTIZE_CHOICES))}, Default is None)") self.add_argument("--quantize", "-q", type=int, choices=ui_defaults.QUANTIZE_CHOICES, default=None, help=f"Quantize the model ({' or '.join(map(str, ui_defaults.QUANTIZE_CHOICES))}, Default is None)")
def add_lora_arguments(self) -> None: def add_lora_arguments(self) -> None:
@ -93,6 +108,9 @@ class CommandLineParser(argparse.ArgumentParser):
# when not provided by CLI flag, find it in the config file # when not provided by CLI flag, find it in the config file
namespace.model = prior_gen_metadata.get("model", None) namespace.model = prior_gen_metadata.get("model", None)
if namespace.base_model is None:
namespace.base_model = prior_gen_metadata.get("base_model", None)
if namespace.prompt is None: if namespace.prompt is None:
namespace.prompt = prior_gen_metadata.get("prompt", None) namespace.prompt = prior_gen_metadata.get("prompt", None)

View File

@ -93,10 +93,12 @@ def test_model_arg_not_in_file(mflux_generate_parser, mflux_generate_minimal_arg
with patch('sys.argv', mflux_generate_minimal_argv + ['--model', 'dev', '--config-from-metadata', metadata_file.as_posix()]): # fmt: off with patch('sys.argv', mflux_generate_minimal_argv + ['--model', 'dev', '--config-from-metadata', metadata_file.as_posix()]): # fmt: off
args = mflux_generate_parser.parse_args() args = mflux_generate_parser.parse_args()
assert args.model == "dev" assert args.model == "dev"
assert args.base_model is None
# test value read from flag # test value read from flag
with patch('sys.argv', mflux_generate_minimal_argv + ['--model', 'schnell', '--config-from-metadata', metadata_file.as_posix()]): # fmt: off with patch('sys.argv', mflux_generate_minimal_argv + ['--model', 'schnell', '--config-from-metadata', metadata_file.as_posix()]): # fmt: off
args = mflux_generate_parser.parse_args() args = mflux_generate_parser.parse_args()
assert args.model == "schnell" assert args.model == "schnell"
assert args.base_model is None
def test_model_arg_in_file(mflux_generate_parser, mflux_generate_minimal_argv, base_metadata_dict, temp_dir): def test_model_arg_in_file(mflux_generate_parser, mflux_generate_minimal_argv, base_metadata_dict, temp_dir):
@ -114,6 +116,25 @@ def test_model_arg_in_file(mflux_generate_parser, mflux_generate_minimal_argv, b
assert args.model == "schnell" assert args.model == "schnell"
def test_base_model_arg_in_file(mflux_generate_parser, mflux_generate_minimal_argv, base_metadata_dict, temp_dir):
metadata_file = temp_dir / "model.json"
with metadata_file.open("wt") as m:
base_metadata_dict["model"] = "some-lab/some-model"
base_metadata_dict["base_model"] = "dev"
json.dump(base_metadata_dict, m, indent=4)
# test value read from file
with patch('sys.argv', mflux_generate_minimal_argv + ['--config-from-metadata', metadata_file.as_posix()]): # fmt: off
args = mflux_generate_parser.parse_args()
assert args.model == "some-lab/some-model"
assert args.base_model == "dev"
# test value read from flag, overrides value from file
with patch('sys.argv', mflux_generate_minimal_argv + ['--base-model', 'schnell', '--config-from-metadata', metadata_file.as_posix()]): # fmt: off
args = mflux_generate_parser.parse_args()
assert args.model == "some-lab/some-model"
# override metadata base model with CLI --base-model
assert args.base_model == "schnell"
def test_prompt_arg(mflux_generate_parser, mflux_generate_minimal_argv, base_metadata_dict, temp_dir): def test_prompt_arg(mflux_generate_parser, mflux_generate_minimal_argv, base_metadata_dict, temp_dir):
metadata_file = temp_dir / "prompt.json" metadata_file = temp_dir / "prompt.json"
file_prompt = "origin of the universe" file_prompt = "origin of the universe"

View File

@ -0,0 +1,68 @@
import pytest
from mflux.config.model_config import InvalidBaseModel, ModelConfigError, ModelLookup
def test_bfl_dev():
model_attrs = ModelLookup.from_name("dev")
assert model_attrs.model_name.startswith("black-forest-labs/")
assert model_attrs.max_sequence_length == 512
assert model_attrs.supports_guidance is True
def test_bfl_dev_from_alias():
model_attrs = ModelLookup.from_alias("dev")
assert model_attrs.model_name.startswith("black-forest-labs/")
assert model_attrs.max_sequence_length == 512
assert model_attrs.supports_guidance is True
def test_bfl_schnell():
model_attrs = ModelLookup.from_name("schnell")
assert model_attrs.model_name.startswith("black-forest-labs/")
assert model_attrs.max_sequence_length == 256
assert model_attrs.supports_guidance is False
def test_bfl_schnell_from_alias():
model_attrs = ModelLookup.from_alias("schnell")
assert model_attrs.model_name.startswith("black-forest-labs/")
assert model_attrs.max_sequence_length == 256
assert model_attrs.supports_guidance is False
def test_community_dev_implicit_base_model():
model_attrs = ModelLookup.from_name("acme-lab/some-awesome-dev-model")
assert model_attrs.max_sequence_length == 512
assert model_attrs.supports_guidance is True
def test_community_schnell_implicit_base_model():
model_attrs = ModelLookup.from_name("acme-lab/some-quick-schnell-model")
assert model_attrs.max_sequence_length == 256
assert model_attrs.supports_guidance is False
def test_community_dev_explicit_base_model():
model_attrs = ModelLookup.from_name("acme-lab/some-awesome-model", base_model="dev")
assert model_attrs.max_sequence_length == 512
assert model_attrs.supports_guidance is True
def test_community_schnell_explicit_base_model():
model_attrs = ModelLookup.from_name("acme-lab/some-awesome-model", base_model="schnell")
assert model_attrs.max_sequence_length == 256
assert model_attrs.supports_guidance is False
def test_model_config_error():
assert pytest.raises(ModelConfigError, ModelLookup.from_name, "acme-lab/some-model-who-knows-what-its-based-on")
def test_invalid_base_model_error():
assert pytest.raises(
InvalidBaseModel,
ModelLookup.from_name,
"acme-lab/some-model-who-knows-what-its-based-on",
base_model="something_unknown",
)