Small updates
This commit is contained in:
parent
344ac90c74
commit
21a98b1369
@ -264,8 +264,8 @@ from mflux import Flux1, Config
|
|||||||
|
|
||||||
# Load the model
|
# Load the model
|
||||||
flux = Flux1.from_name(
|
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
|
||||||
|
|||||||
@ -1,3 +1,4 @@
|
|||||||
|
import warnings
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import Literal
|
from typing import Literal
|
||||||
|
|
||||||
@ -22,14 +23,6 @@ class ModelConfig:
|
|||||||
supports_guidance: bool
|
supports_guidance: bool
|
||||||
base_model: str | None
|
base_model: str | None
|
||||||
|
|
||||||
@property
|
|
||||||
def alias(self):
|
|
||||||
# maintain compatibility with < 0.4.0 behavior
|
|
||||||
# where alias is the name of an official model
|
|
||||||
if self.model_name.startswith("black-forest-labs/FLUX.1-"):
|
|
||||||
return self.model_name[len("black-forest-labs/FLUX.1-") :].lower()
|
|
||||||
return None
|
|
||||||
|
|
||||||
|
|
||||||
DefaultModelConfigs = {
|
DefaultModelConfigs = {
|
||||||
"dev": ModelConfig(
|
"dev": ModelConfig(
|
||||||
@ -50,18 +43,27 @@ DefaultModelConfigs = {
|
|||||||
|
|
||||||
|
|
||||||
class ModelLookup:
|
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_name(
|
def from_name(
|
||||||
alias: str,
|
model_name: str,
|
||||||
base_model: Literal["dev", "schnell"] | None = None,
|
base_model: Literal["dev", "schnell"] | None = None,
|
||||||
) -> ModelConfig:
|
) -> ModelConfig:
|
||||||
if alias in DefaultModelConfigs:
|
if model_name in DefaultModelConfigs:
|
||||||
return DefaultModelConfigs[alias]
|
return DefaultModelConfigs[model_name]
|
||||||
|
|
||||||
if all(["dev" not in alias, "schnell" not in alias, base_model is None]):
|
if all(["dev" not in model_name, "schnell" not in model_name, base_model is None]):
|
||||||
raise ModelConfigError(
|
raise ModelConfigError(
|
||||||
"Cannot infer base model and max_sequence_length "
|
"Cannot infer base model and max_sequence_length "
|
||||||
f"from model reference: {alias!r}. "
|
f"from model reference: {model_name!r}. "
|
||||||
"Please specify --base-model [dev | schnell]"
|
"Please specify --base-model [dev | schnell]"
|
||||||
)
|
)
|
||||||
|
|
||||||
@ -69,10 +71,10 @@ class ModelLookup:
|
|||||||
raise InvalidBaseModel("As of this version, mflux only recognizes base models dev or schnell")
|
raise InvalidBaseModel("As of this version, mflux only recognizes base models dev or schnell")
|
||||||
|
|
||||||
if base_model is None:
|
if base_model is None:
|
||||||
# infer base model on apparent model namming
|
# infer base model on apparent model naming
|
||||||
if "dev" in alias:
|
if "dev" in model_name:
|
||||||
base_model = "dev"
|
base_model = "dev"
|
||||||
elif "schnell" in alias:
|
elif "schnell" in model_name:
|
||||||
base_model = "schnell"
|
base_model = "schnell"
|
||||||
|
|
||||||
if base_model == "dev":
|
if base_model == "dev":
|
||||||
@ -83,16 +85,14 @@ class ModelLookup:
|
|||||||
max_sequence_length = KNOWN_SEQUENCE_LENGTH_BY_BASE_MODEL["schnell"]
|
max_sequence_length = KNOWN_SEQUENCE_LENGTH_BY_BASE_MODEL["schnell"]
|
||||||
|
|
||||||
return ModelConfig(
|
return ModelConfig(
|
||||||
alias, # actually this arg is model_name
|
model_name=model_name,
|
||||||
DEFAULT_TRAIN_STEPS,
|
num_train_steps=DEFAULT_TRAIN_STEPS,
|
||||||
max_sequence_length,
|
max_sequence_length=max_sequence_length,
|
||||||
supports_guidance,
|
supports_guidance=supports_guidance,
|
||||||
base_model,
|
base_model=base_model,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
# maintain old `ModelConfig.from_alias` function name for backwards compatibility in user code and docs
|
|
||||||
ModelConfig.from_alias = ModelLookup.from_name
|
|
||||||
# keep these class members to be backwards compatible with < 0.5.0 ModelConfig Enum implementation
|
# keep these class members to be backwards compatible with < 0.5.0 ModelConfig Enum implementation
|
||||||
ModelConfig.FLUX1_DEV = DefaultModelConfigs["dev"]
|
ModelConfig.FLUX1_DEV = DefaultModelConfigs["dev"]
|
||||||
ModelConfig.FLUX1_SCHNELL = DefaultModelConfigs["schnell"]
|
ModelConfig.FLUX1_SCHNELL = DefaultModelConfigs["schnell"]
|
||||||
|
|||||||
@ -1,3 +1,4 @@
|
|||||||
|
import warnings
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
import mlx.core as mx
|
import mlx.core as mx
|
||||||
@ -135,16 +136,22 @@ class Flux1(nn.Module):
|
|||||||
config=config,
|
config=config,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
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
|
@staticmethod
|
||||||
def from_name(model_name: str, quantize: int | None = None) -> "Flux1":
|
def from_name(model_name: str, quantize: int | None = None) -> "Flux1":
|
||||||
return Flux1(
|
return Flux1(
|
||||||
model_config=ModelLookup.from_name(model_name),
|
model_config=ModelLookup.from_name(model_name=model_name, base_model=None),
|
||||||
quantize=quantize,
|
quantize=quantize,
|
||||||
)
|
)
|
||||||
|
|
||||||
# maintain old `from_alias` function name for backwards compatibility in user code and docs
|
|
||||||
from_alias = from_name
|
|
||||||
|
|
||||||
def save_model(self, base_path: str) -> None:
|
def save_model(self, base_path: str) -> None:
|
||||||
ModelSaver.save_model(self, self.bits, base_path)
|
ModelSaver.save_model(self, self.bits, base_path)
|
||||||
|
|
||||||
|
|||||||
@ -17,7 +17,7 @@ def main():
|
|||||||
|
|
||||||
# Load the model
|
# Load the model
|
||||||
flux = Flux1(
|
flux = Flux1(
|
||||||
model_config=ModelLookup.from_name(args.model, base_model=args.base_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,
|
||||||
|
|||||||
@ -16,7 +16,7 @@ def main():
|
|||||||
|
|
||||||
# Load the model
|
# Load the model
|
||||||
flux = Flux1Controlnet(
|
flux = Flux1Controlnet(
|
||||||
model_config=ModelLookup.from_name(args.model, base_model=args.base_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,
|
||||||
|
|||||||
@ -57,9 +57,7 @@ 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,
|
||||||
if self.model_config.alias is not None
|
|
||||||
else self.model_config.model_name,
|
|
||||||
"base_model": str(self.model_config.base_model),
|
"base_model": str(self.model_config.base_model),
|
||||||
"seed": self.seed,
|
"seed": self.seed,
|
||||||
"steps": self.steps,
|
"steps": self.steps,
|
||||||
|
|||||||
@ -1,20 +1,6 @@
|
|||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from mflux.config.model_config import InvalidBaseModel, ModelConfig, ModelConfigError, ModelLookup
|
from mflux.config.model_config import InvalidBaseModel, ModelConfigError, ModelLookup
|
||||||
|
|
||||||
|
|
||||||
def test_from_alias_function_redirect():
|
|
||||||
# backwards compatibility for when user follows older docs
|
|
||||||
# but is using a newer mflux version >= 0.5
|
|
||||||
assert ModelConfig.from_alias == ModelLookup.from_name
|
|
||||||
|
|
||||||
|
|
||||||
def test_model_config_class_members_and_alias():
|
|
||||||
# these FLUX1_* class members and the alias attribute
|
|
||||||
# existed as members of ModelConfig when it was an Enum
|
|
||||||
# keep them around for backwards compatibility
|
|
||||||
assert ModelConfig.FLUX1_DEV.alias == "dev"
|
|
||||||
assert ModelConfig.FLUX1_SCHNELL.alias == "schnell"
|
|
||||||
|
|
||||||
|
|
||||||
def test_bfl_dev():
|
def test_bfl_dev():
|
||||||
@ -24,6 +10,13 @@ def test_bfl_dev():
|
|||||||
assert model_attrs.supports_guidance is True
|
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():
|
def test_bfl_schnell():
|
||||||
model_attrs = ModelLookup.from_name("schnell")
|
model_attrs = ModelLookup.from_name("schnell")
|
||||||
assert model_attrs.model_name.startswith("black-forest-labs/")
|
assert model_attrs.model_name.startswith("black-forest-labs/")
|
||||||
@ -31,6 +24,13 @@ def test_bfl_schnell():
|
|||||||
assert model_attrs.supports_guidance is False
|
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():
|
def test_community_dev_implicit_base_model():
|
||||||
model_attrs = ModelLookup.from_name("acme-lab/some-awesome-dev-model")
|
model_attrs = ModelLookup.from_name("acme-lab/some-awesome-dev-model")
|
||||||
assert model_attrs.max_sequence_length == 512
|
assert model_attrs.max_sequence_length == 512
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user