Qwen-Image-Layered-MRP-MLX/src/mflux/config/model_config.py
2025-04-24 09:08:21 +02:00

158 lines
4.9 KiB
Python

from functools import lru_cache
from typing import Literal
from mflux.error.error import InvalidBaseModel, ModelConfigError
class ModelConfig:
def __init__(
self,
alias: str | None,
model_name: str,
base_model: str | None,
num_train_steps: int,
max_sequence_length: int,
supports_guidance: bool,
requires_sigma_shift: bool,
priority: int,
):
self.alias = alias
self.model_name = model_name
self.base_model = base_model
self.num_train_steps = num_train_steps
self.max_sequence_length = max_sequence_length
self.supports_guidance = supports_guidance
self.requires_sigma_shift = requires_sigma_shift
self.priority = priority
@staticmethod
@lru_cache
def dev() -> "ModelConfig":
return AVAILABLE_MODELS["dev"]
@staticmethod
@lru_cache
def dev_fill() -> "ModelConfig":
return AVAILABLE_MODELS["dev-fill"]
@staticmethod
@lru_cache
def dev_depth() -> "ModelConfig":
return AVAILABLE_MODELS["dev-depth"]
@staticmethod
@lru_cache
def dev_redux() -> "ModelConfig":
return AVAILABLE_MODELS["dev-redux"]
@staticmethod
@lru_cache
def schnell() -> "ModelConfig":
return AVAILABLE_MODELS["schnell"]
def x_embedder_input_dim(self) -> int:
if self.alias == "dev-fill":
return 384
if self.alias == "dev-depth":
return 128
else:
return 64
@staticmethod
def from_name(
model_name: str,
base_model: Literal["dev", "schnell", "dev-fill"] | None = None,
) -> "ModelConfig":
# 0. Get all base models (where base_model is None) sorted by priority
base_models = sorted(
[model for model in AVAILABLE_MODELS.values() if model.base_model is None], key=lambda x: x.priority
)
# 1. Check if model_name matches any base model's alias or full name
for base in base_models:
if model_name in (base.alias, base.model_name):
return base
# 2. Validate explicit base_model
allowed_names = []
for base in base_models:
allowed_names.extend([base.alias, base.model_name])
if base_model and base_model not in allowed_names:
raise InvalidBaseModel(f"Invalid base_model. Choose one of {allowed_names}")
# 3. Determine the base model (explicit or inferred)
if base_model:
# Find by explicit base_model name
default_base = next((b for b in base_models if base_model in (b.alias, b.model_name)), None)
else:
# Infer from model_name substring (priority order via sorted base_models)
default_base = next((b for b in base_models if b.alias and b.alias in model_name), None)
if not default_base:
raise ModelConfigError(f"Cannot infer base_model from {model_name}")
# 4. Construct the config
return ModelConfig(
alias=default_base.alias,
model_name=model_name,
base_model=default_base.model_name,
num_train_steps=default_base.num_train_steps,
max_sequence_length=default_base.max_sequence_length,
supports_guidance=default_base.supports_guidance,
requires_sigma_shift=default_base.requires_sigma_shift,
priority=default_base.priority,
)
AVAILABLE_MODELS = {
"schnell": ModelConfig(
alias="schnell",
model_name="black-forest-labs/FLUX.1-schnell",
base_model=None,
num_train_steps=1000,
max_sequence_length=256,
supports_guidance=False,
requires_sigma_shift=False,
priority=2,
),
"dev": ModelConfig(
alias="dev",
model_name="black-forest-labs/FLUX.1-dev",
base_model=None,
num_train_steps=1000,
max_sequence_length=512,
supports_guidance=True,
requires_sigma_shift=True,
priority=1,
),
"dev-fill": ModelConfig(
alias="dev-fill",
model_name="black-forest-labs/FLUX.1-Fill-dev",
base_model=None,
num_train_steps=1000,
max_sequence_length=512,
supports_guidance=True,
requires_sigma_shift=True,
priority=0,
),
"dev-depth": ModelConfig(
alias="dev-depth",
model_name="black-forest-labs/FLUX.1-Depth-dev",
base_model=None,
num_train_steps=1000,
max_sequence_length=512,
supports_guidance=True,
requires_sigma_shift=True,
priority=4,
),
"dev-redux": ModelConfig(
alias="dev-redux",
model_name="black-forest-labs/FLUX.1-Redux-dev",
base_model=None,
num_train_steps=1000,
max_sequence_length=512,
supports_guidance=True,
requires_sigma_shift=True,
priority=3,
),
}