Add support for Z-image Turbo & major improvements for weight loading and parameter resolution (#284)
This commit is contained in:
parent
eb83175680
commit
37d202d0df
37
.github/workflows/tests.yml
vendored
Normal file
37
.github/workflows/tests.yml
vendored
Normal file
@ -0,0 +1,37 @@
|
|||||||
|
name: Fast Tests
|
||||||
|
|
||||||
|
on:
|
||||||
|
pull_request:
|
||||||
|
branches: [main]
|
||||||
|
push:
|
||||||
|
branches: [main]
|
||||||
|
|
||||||
|
jobs:
|
||||||
|
fast-tests:
|
||||||
|
runs-on: macos-14 # Apple Silicon runner (M1)
|
||||||
|
|
||||||
|
steps:
|
||||||
|
- name: Checkout code
|
||||||
|
uses: actions/checkout@v4
|
||||||
|
|
||||||
|
- name: Set up Python
|
||||||
|
uses: actions/setup-python@v5
|
||||||
|
with:
|
||||||
|
python-version: "3.11"
|
||||||
|
|
||||||
|
- name: Install uv
|
||||||
|
uses: astral-sh/setup-uv@v4
|
||||||
|
with:
|
||||||
|
version: "latest"
|
||||||
|
|
||||||
|
- name: Install dependencies
|
||||||
|
run: |
|
||||||
|
uv venv
|
||||||
|
uv pip install -e .
|
||||||
|
uv pip install pytest
|
||||||
|
|
||||||
|
- name: Run fast tests
|
||||||
|
run: |
|
||||||
|
source .venv/bin/activate
|
||||||
|
python -m pytest -m fast -v
|
||||||
|
|
||||||
103
CHANGELOG.md
103
CHANGELOG.md
@ -5,6 +5,109 @@ All notable changes to this project will be documented in this file.
|
|||||||
The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.0.0/),
|
The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.0.0/),
|
||||||
and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html).
|
and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html).
|
||||||
|
|
||||||
|
## [0.13.0] - 2025-12-03
|
||||||
|
|
||||||
|
# MFLUX v.0.13.0 Release Notes
|
||||||
|
|
||||||
|
### 🎨 New Model Support
|
||||||
|
|
||||||
|
- **Z-Image Turbo Support**: Added support for [Z-Image Turbo](https://huggingface.co/Tongyi-MAI/Z-Image-Turbo), a fast distilled Z-Image variant optimized for speed
|
||||||
|
- **New command**: `mflux-generate-z-image-turbo` for rapid image generation (with LoRA support, img2img, and quantization)
|
||||||
|
|
||||||
|
### ✨ New Features
|
||||||
|
|
||||||
|
- **FIBO VLM Quantization Support**: The FIBO VLM commands (`mflux-fibo-inspire`, `mflux-fibo-refine`) now support quantization via the `-q` flag (3, 4, 5, 6, or 8-bit)
|
||||||
|
|
||||||
|
- **Unified `--model` argument**: The `--model` flag now accepts local paths, HuggingFace repos, or predefined model names
|
||||||
|
- Local paths: `--model /Users/me/models/fibo-4bit` or `--model ~/my-model`
|
||||||
|
- HuggingFace repos: `--model briaai/Fibo-mlx-4bit`
|
||||||
|
- Predefined names: `--model dev`, `--model schnell`, `--model fibo`
|
||||||
|
- This mirrors how LoRA paths work for a consistent UX
|
||||||
|
|
||||||
|
- **Scale Factor Dimensions for Img2Img**: Generalized the scale factor feature (e.g., `2x`, `0.5x`, `auto`) from upscaling to all img2img commands
|
||||||
|
- Specify output dimensions relative to input image: `--width 2x --height 2x`
|
||||||
|
- Use `auto` to match input image dimensions: `--width auto --height auto`
|
||||||
|
- Mix scale factors with absolute values: `--width 2x --height 512`
|
||||||
|
- Supported in: `mflux-generate`, `mflux-generate-z-image-turbo`, `mflux-generate-fibo`, `mflux-generate-kontext`, `mflux-generate-qwen`
|
||||||
|
- **DimensionResolver utility**: New `DimensionResolver.resolve()` for consistent dimension handling across commands
|
||||||
|
|
||||||
|
### 🔧 Architecture Improvements
|
||||||
|
|
||||||
|
- **Unified Resolution System**: New `resolution/` module for consistent parameter resolution across all models
|
||||||
|
- `PathResolution`: Resolves model paths from local paths, HuggingFace repos, or predefined names
|
||||||
|
- `LoRAResolution`: Handles LoRA path resolution from all supported formats
|
||||||
|
- `ConfigResolution`: Centralizes configuration resolution logic
|
||||||
|
- `QuantizationResolution`: Determines quantization from saved models or CLI args
|
||||||
|
- **Unified Weight Loading System**: Complete rewrite of weight handling with declarative mappings
|
||||||
|
- New `WeightLoader` with single `load(model_path)` interface
|
||||||
|
- `WeightDefinition` classes define model structure per model family
|
||||||
|
- `WeightMapping` declarative mappings replace imperative weight handlers
|
||||||
|
- Removed all per-model `weight_handler_*.py` files in favor of unified system
|
||||||
|
- **Unified Tokenizer System**: New common tokenizer module
|
||||||
|
- `TokenizerLoader.load_all()` with unified `model_path` interface
|
||||||
|
- Removed model-specific tokenizer handlers (`clip_tokenizer.py`, `t5_tokenizer.py`, etc.)
|
||||||
|
- **Unified LoRA API**: Simplified LoRA loading to a single `lora_paths` parameter
|
||||||
|
- All LoRA formats now resolved through `LoRALibrary.resolve_paths()`:
|
||||||
|
- Local paths: `/path/to/lora.safetensors`
|
||||||
|
- Registry names: `my-lora` (from `LORA_LIBRARY_PATH`)
|
||||||
|
- HuggingFace repos: `author/model`
|
||||||
|
- **New**: HuggingFace collections: `repo_id:filename.safetensors`
|
||||||
|
- Simplified model initialization: just pass `lora_paths` and everything resolves automatically
|
||||||
|
- **Unified Latent Creator Interface**: Standardized `unpack_latents(latents, height, width)` signature across all model families
|
||||||
|
- `FluxLatentCreator`, `ZImageLatentCreator`, `FiboLatentCreator`, and `QwenLatentCreator` now share the same interface
|
||||||
|
- Moved `FIBO._unpack_latents` to `FiboLatentCreator.unpack_latents` for consistency
|
||||||
|
- **StepwiseHandler Refactor**: Fixed `StepwiseHandler` to work with all model types by accepting a `latent_creator` parameter
|
||||||
|
- Previously hardcoded to `FluxLatentCreator`, now model-agnostic
|
||||||
|
- Each command passes its appropriate latent creator to `CallbackManager.register_callbacks()`
|
||||||
|
- **CLI Reorganization**: Moved CLI entry points to model-specific directories (e.g., `mflux/models/flux/cli/`)
|
||||||
|
|
||||||
|
### 🔄 Breaking Changes
|
||||||
|
|
||||||
|
- **Simplified `generate_image()` API** (programmatic users only):
|
||||||
|
- Removed `Config` class - parameters are now passed directly to `generate_image()`
|
||||||
|
- Removed `RuntimeConfig` class - internal complexity eliminated
|
||||||
|
- Added `Flux1` export to main `mflux` module for cleaner imports
|
||||||
|
- **LoRA API simplified** (programmatic users only):
|
||||||
|
- Removed `lora_names` and `lora_repo_id` parameters from all model classes (`Flux1`, `QwenImage`, `QwenImageEdit`, etc.)
|
||||||
|
- Removed `--lora-name` and `--lora-repo-id` CLI arguments
|
||||||
|
- Removed `LoRAHuggingFaceDownloader` class
|
||||||
|
|
||||||
|
### 🔄 Breaking Changes (CLI)
|
||||||
|
|
||||||
|
- **`--path` flag removed**: The deprecated `--path` flag for loading models has been removed. Use `--model` instead for local paths, HuggingFace repos, or predefined model names.
|
||||||
|
|
||||||
|
### 🐛 Bug Fixes
|
||||||
|
|
||||||
|
- **`--model` flag not working**: Fixed bug where the `--model` argument wasn't being used for loading models from HuggingFace or local paths. All CLI commands now correctly use `--model` for model path resolution.
|
||||||
|
- **Model Saving Index File**: Fixed issue where locally saved models (via `mflux-save`) would fail to load when uploaded to HuggingFace, due to missing `model.safetensors.index.json`. The model saver now generates this index file alongside the safetensor shards, ensuring compatibility with both mflux and standard HuggingFace loading paths. (see [#285](https://github.com/filipstrand/mflux/issues/285))
|
||||||
|
|
||||||
|
### 🧪 Test Infrastructure
|
||||||
|
|
||||||
|
- **Test markers**: Added `fast` and `slow` pytest markers to categorize tests
|
||||||
|
- Fast tests: Unit tests that don't generate images (parsers, schedulers, resolution, utilities)
|
||||||
|
- Slow tests: Integration tests that generate actual images and compare to references
|
||||||
|
- **New Makefile targets**:
|
||||||
|
- `make test-fast` - Run fast tests only (quick feedback during development)
|
||||||
|
- `make test-slow` - Run slow tests only (image generation tests)
|
||||||
|
- `make test` - Run all tests (unchanged)
|
||||||
|
- Run specific test categories: `pytest -m fast` or `pytest -m slow`
|
||||||
|
- **GitHub Actions CI**: Fast tests now run automatically on PRs and pushes to main
|
||||||
|
|
||||||
|
### 🔧 Internal Changes
|
||||||
|
|
||||||
|
- Simplified `WeightLoader.load()` to take a single `model_path` parameter instead of separate `repo_id` and `local_path`
|
||||||
|
- Simplified `TokenizerLoader.load_all()` with the same unified `model_path` interface
|
||||||
|
- Renamed `local_path` parameter to `model_path` in all model constructors for clarity
|
||||||
|
- Removed `quantization_util.py` - quantization now handled through `QuantizationResolution`
|
||||||
|
- Removed `lora_huggingface_downloader.py` - downloading integrated into `LoRAResolution`
|
||||||
|
- Added comprehensive test coverage for resolution modules
|
||||||
|
|
||||||
|
### 👩💻 Contributors
|
||||||
|
|
||||||
|
- **Filip Strand (@filipstrand)**: Z-Image Turbo support, architecture improvements, core development
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
## [0.12.1] - 2025-11-27
|
## [0.12.1] - 2025-11-27
|
||||||
|
|
||||||
### 🐛 Bug Fixes
|
### 🐛 Bug Fixes
|
||||||
|
|||||||
19
Makefile
19
Makefile
@ -97,6 +97,21 @@ test: ensure-pytest
|
|||||||
$(PYTHON) -m pytest
|
$(PYTHON) -m pytest
|
||||||
# ✅ Tests completed
|
# ✅ Tests completed
|
||||||
|
|
||||||
|
# Run fast tests only (no image generation)
|
||||||
|
.PHONY: test-fast
|
||||||
|
test-fast: ensure-pytest
|
||||||
|
# 🏗️ Running fast tests (no image generation)...
|
||||||
|
$(PYTHON) -m pytest -m fast
|
||||||
|
# ✅ Fast tests completed
|
||||||
|
|
||||||
|
# Run slow tests only (image generation tests)
|
||||||
|
.PHONY: test-slow
|
||||||
|
test-slow: ensure-pytest
|
||||||
|
# 🏗️ Running slow tests (image generation)...
|
||||||
|
uv pip install mlx==0.29.2 # Install pinned MLX version specifically for testing
|
||||||
|
$(PYTHON) -m pytest -m slow
|
||||||
|
# ✅ Slow tests completed
|
||||||
|
|
||||||
|
|
||||||
# Run uv build and check dist sizes for optimized user installs
|
# Run uv build and check dist sizes for optimized user installs
|
||||||
.PHONY: build
|
.PHONY: build
|
||||||
@ -129,7 +144,9 @@ help:
|
|||||||
@echo " make lint - Run ruff python linter"
|
@echo " make lint - Run ruff python linter"
|
||||||
@echo " make format - Run ruff code formatter"
|
@echo " make format - Run ruff code formatter"
|
||||||
@echo " make check - Run linters auto fixes *and* style formatter via pre-commit hook"
|
@echo " make check - Run linters auto fixes *and* style formatter via pre-commit hook"
|
||||||
@echo " make test - Run tests"
|
@echo " make test - Run all tests"
|
||||||
|
@echo " make test-fast - Run fast tests only (no image generation)"
|
||||||
|
@echo " make test-slow - Run slow tests only (image generation)"
|
||||||
@echo " make build - Build distribution packages and check sizes"
|
@echo " make build - Build distribution packages and check sizes"
|
||||||
@echo " make clean - Remove the virtual environment"
|
@echo " make clean - Remove the virtual environment"
|
||||||
@echo " make help - Show this help message"
|
@echo " make help - Show this help message"
|
||||||
|
|||||||
133
README.md
133
README.md
@ -4,7 +4,7 @@
|
|||||||
|
|
||||||
### About
|
### About
|
||||||
|
|
||||||
Run the powerful [FLUX](https://blackforestlabs.ai/#get-flux), [Qwen Image](https://github.com/QwenLM/Qwen-Image) and [FIBO](https://huggingface.co/briaai/FIBO) models locally on your Mac!
|
Run the powerful [FLUX](https://blackforestlabs.ai/#get-flux), [Qwen Image](https://github.com/QwenLM/Qwen-Image), [FIBO](https://huggingface.co/briaai/FIBO), and [Z-Image](https://huggingface.co/Tongyi-MAI/Z-Image-Turbo) models locally on your Mac!
|
||||||
|
|
||||||
### Table of contents
|
### Table of contents
|
||||||
|
|
||||||
@ -17,13 +17,14 @@ Run the powerful [FLUX](https://blackforestlabs.ai/#get-flux), [Qwen Image](http
|
|||||||
- [⏱️ Image generation speed (updated)](#%EF%B8%8F-image-generation-speed-updated)
|
- [⏱️ Image generation speed (updated)](#%EF%B8%8F-image-generation-speed-updated)
|
||||||
- [↔️ Equivalent to Diffusers implementation](#%EF%B8%8F-equivalent-to-diffusers-implementation)
|
- [↔️ Equivalent to Diffusers implementation](#%EF%B8%8F-equivalent-to-diffusers-implementation)
|
||||||
- [🗜️ Quantization](#%EF%B8%8F-quantization)
|
- [🗜️ Quantization](#%EF%B8%8F-quantization)
|
||||||
- [💽 Running a non-quantized model directly from disk](#-running-a-non-quantized-model-directly-from-disk)
|
- [💽 Running a model directly from disk](#-running-a-model-directly-from-disk)
|
||||||
- [🌐 Third-Party HuggingFace Model Support](#-third-party-huggingface-model-support)
|
- [🌐 Third-Party HuggingFace Model Support](#-third-party-huggingface-model-support)
|
||||||
- [🎨 Image-to-Image](#-image-to-image)
|
- [🎨 Image-to-Image](#-image-to-image)
|
||||||
- [🦙 Qwen Models](#-qwen-models)
|
- [🦙 Qwen Models](#-qwen-models)
|
||||||
* [🖼️ Qwen Image](#%EF%B8%8F-qwen-image)
|
* [🖼️ Qwen Image](#%EF%B8%8F-qwen-image)
|
||||||
* [✏️ Qwen Image Edit](#%EF%B8%8F-qwen-image-edit)
|
* [✏️ Qwen Image Edit](#%EF%B8%8F-qwen-image-edit)
|
||||||
- [🌀 FIBO](#-fibo)
|
- [🌀 FIBO](#-fibo)
|
||||||
|
- [⚡ Z-Image](#-z-image)
|
||||||
- [🔌 LoRA](#-lora)
|
- [🔌 LoRA](#-lora)
|
||||||
- [🎭 In-Context Generation](#-in-context-generation)
|
- [🎭 In-Context Generation](#-in-context-generation)
|
||||||
* [📸 Kontext](#-kontext)
|
* [📸 Kontext](#-kontext)
|
||||||
@ -51,7 +52,7 @@ Run the powerful [FLUX](https://blackforestlabs.ai/#get-flux), [Qwen Image](http
|
|||||||
|
|
||||||
### Philosophy
|
### Philosophy
|
||||||
|
|
||||||
MFLUX is a line-by-line port of the FLUX, Qwen and Bria models implementations in the [Huggingface Diffusers](https://github.com/huggingface/diffusers) and [Huggingface Transformers](https://github.com/huggingface/transformers) libraries to [Apple MLX](https://github.com/ml-explore/mlx).
|
MFLUX is a line-by-line port of the FLUX, Qwen, Bria and Z-Image models implementations in the [Huggingface Diffusers](https://github.com/huggingface/diffusers) and [Huggingface Transformers](https://github.com/huggingface/transformers) libraries to [Apple MLX](https://github.com/ml-explore/mlx).
|
||||||
MFLUX is purposefully kept minimal and explicit - Network architectures are hardcoded and no config files are used
|
MFLUX is purposefully kept minimal and explicit - Network architectures are hardcoded and no config files are used
|
||||||
except for the tokenizers. The aim is to have a tiny codebase with the single purpose of expressing these models
|
except for the tokenizers. The aim is to have a tiny codebase with the single purpose of expressing these models
|
||||||
(thereby avoiding too many abstractions). While MFLUX priorities readability over generality and performance, [it can still be quite fast](#%EF%B8%8F-image-generation-speed-updated), [and even faster quantized](#%EF%B8%8F-quantization).
|
(thereby avoiding too many abstractions). While MFLUX priorities readability over generality and performance, [it can still be quite fast](#%EF%B8%8F-image-generation-speed-updated), [and even faster quantized](#%EF%B8%8F-quantization).
|
||||||
@ -171,8 +172,7 @@ This is useful for integrating MFLUX into shell scripts or dynamically generatin
|
|||||||
Alternatively, you can use MFLUX directly in Python:
|
Alternatively, you can use MFLUX directly in Python:
|
||||||
|
|
||||||
```python
|
```python
|
||||||
from mflux.flux.flux import Flux1
|
from mflux import Flux1
|
||||||
from mflux.config.config import Config
|
|
||||||
|
|
||||||
# Load the model
|
# Load the model
|
||||||
flux = Flux1.from_name(
|
flux = Flux1.from_name(
|
||||||
@ -184,11 +184,9 @@ flux = Flux1.from_name(
|
|||||||
image = flux.generate_image(
|
image = flux.generate_image(
|
||||||
seed=2,
|
seed=2,
|
||||||
prompt="Luxury food photograph",
|
prompt="Luxury food photograph",
|
||||||
config=Config(
|
num_inference_steps=2, # "schnell" works well with 2-4 steps, "dev" and "krea-dev" work well with 20-25 steps
|
||||||
num_inference_steps=2, # "schnell" works well with 2-4 steps, "dev" and "krea-dev" work well with 20-25 steps
|
height=1024,
|
||||||
height=1024,
|
width=1024,
|
||||||
width=1024,
|
|
||||||
)
|
|
||||||
)
|
)
|
||||||
|
|
||||||
image.save(path="image.png")
|
image.save(path="image.png")
|
||||||
@ -200,7 +198,7 @@ For more advanced Python usage and additional configuration options, you can exp
|
|||||||
|
|
||||||
*By default, mflux caches files in `~/Library/Caches/mflux/`. The Hugging Face model files themselves are cached separately in the Hugging Face cache directory (e.g., `~/.cache/huggingface/`).*
|
*By default, mflux caches files in `~/Library/Caches/mflux/`. The Hugging Face model files themselves are cached separately in the Hugging Face cache directory (e.g., `~/.cache/huggingface/`).*
|
||||||
|
|
||||||
*To change the mflux cache location, set the `MFLUX_CACHE_DIR` environment variable. To change the Hugging Face cache location, you can modify the `HF_HOME` environment variable. For more details on Hugging Face cache settings, please refer to the [Hugging Face documentation](https://huggingface.co/docs/huggingface_hub/en/package_reference/environment_variables)*.
|
*To change the mflux cache location, set the `MFLUX_CACHE_DIR` environment variable. To change the Hugging Face cache location, you can modify the `HF_HOME` environment variable (e.g. to `HF_HOME=/Volumes/T7/.cache/huggingface`). For more details on Hugging Face cache settings, please refer to the [Hugging Face documentation](https://huggingface.co/docs/huggingface_hub/en/package_reference/environment_variables)*.
|
||||||
|
|
||||||
🔒 [FLUX.1-dev currently requires granted access to its Huggingface repo. For troubleshooting, see the issue tracker](https://github.com/filipstrand/mflux/issues/14) 🔒
|
🔒 [FLUX.1-dev currently requires granted access to its Huggingface repo. For troubleshooting, see the issue tracker](https://github.com/filipstrand/mflux/issues/14) 🔒
|
||||||
|
|
||||||
@ -236,9 +234,12 @@ mflux-generate \
|
|||||||
|
|
||||||
- **`--prompt`** (required, `str`): Text description of the image to generate. Use `-` to read the prompt from stdin (e.g., `echo "A beautiful sunset" | mflux-generate --prompt -`).
|
- **`--prompt`** (required, `str`): Text description of the image to generate. Use `-` to read the prompt from stdin (e.g., `echo "A beautiful sunset" | mflux-generate --prompt -`).
|
||||||
|
|
||||||
- **`--model`** or **`-m`** (required, `str`): Model to use for generation. Can be one of the official Flux models (`"schnell"`, `"dev"`, or `"krea-dev"`) or a HuggingFace repository ID for a compatible third-party model (e.g., `"Freepik/flux.1-lite-8B-alpha"`). For Qwen models, use `mflux-generate-qwen` instead.
|
- **`--model`** or **`-m`** (required, `str`): Model to use for generation. Accepts:
|
||||||
|
- Predefined names: `"schnell"`, `"dev"`, `"krea-dev"`, `"fibo"`, `"z-image-turbo"`
|
||||||
|
- HuggingFace repos: `"Freepik/flux.1-lite-8B-alpha"`, `"briaai/Fibo-mlx-4bit"`
|
||||||
|
- Local paths: `"/Users/me/models/my-model"`, `"~/my-model"`
|
||||||
|
|
||||||
- **`--base-model`** (optional, `str`, default: `None`): Specifies which base architecture a third-party model is derived from (`"schnell"`, `"dev"`, or `"krea-dev"`). Required when using third-party models from HuggingFace.
|
- **`--base-model`** (optional, `str`, default: `None`): Specifies which base architecture a third-party model is derived from (`"schnell"`, `"dev"`, or `"krea-dev"`). Required when using third-party models from HuggingFace or local paths.
|
||||||
|
|
||||||
- **`--output`** (optional, `str`, default: `"image.png"`): Output image filename. If `--seed` or `--auto-seeds` establishes multiple seed values, the output filename will automatically be modified to include the seed value (e.g., `image_seed_42.png`).
|
- **`--output`** (optional, `str`, default: `"image.png"`): Output image filename. If `--seed` or `--auto-seeds` establishes multiple seed values, the output filename will automatically be modified to include the seed value (e.g., `image_seed_42.png`).
|
||||||
|
|
||||||
@ -254,11 +255,15 @@ mflux-generate \
|
|||||||
|
|
||||||
- **`--guidance`** (optional, `float`, default: `3.5`): Guidance scale (only used for `"dev"` and `"krea-dev"` models).
|
- **`--guidance`** (optional, `float`, default: `3.5`): Guidance scale (only used for `"dev"` and `"krea-dev"` models).
|
||||||
|
|
||||||
- **`--path`** (optional, `str`, default: `None`): Path to a local model on disk.
|
- **`--path`** (optional, `str`, default: `None`): **[DEPRECATED: use `--model` instead]** Path to a local model on disk.
|
||||||
|
|
||||||
- **`--quantize`** or **`-q`** (optional, `int`, default: `None`): [Quantization](#%EF%B8%8F-quantization) (choose between `3`, `4`, `5`, `6`, or `8` bits).
|
- **`--quantize`** or **`-q`** (optional, `int`, default: `None`): [Quantization](#%EF%B8%8F-quantization) (choose between `3`, `4`, `5`, `6`, or `8` bits).
|
||||||
|
|
||||||
- **`--lora-paths`** (optional, `[str]`, default: `None`): The paths to the [LoRA](#-LoRA) weights.
|
- **`--lora-paths`** (optional, `[str]`, default: `None`): The paths to the [LoRA](#-LoRA) weights. Supports multiple formats:
|
||||||
|
- Local files: `/path/to/lora.safetensors`
|
||||||
|
- HuggingFace repos: `author/model` (auto-downloads)
|
||||||
|
- HuggingFace collections: `repo_id:filename.safetensors` (downloads specific file)
|
||||||
|
- Registry names: `my-lora` (via `LORA_LIBRARY_PATH`)
|
||||||
|
|
||||||
- **`--lora-scales`** (optional, `[float]`, default: `None`): The scale for each respective [LoRA](#-LoRA) (will default to `1.0` if not specified and only one LoRA weight is loaded.)
|
- **`--lora-scales`** (optional, `[float]`, default: `None`): The scale for each respective [LoRA](#-LoRA) (will default to `1.0` if not specified and only one LoRA weight is loaded.)
|
||||||
|
|
||||||
@ -274,10 +279,6 @@ mflux-generate \
|
|||||||
|
|
||||||
- **`--battery-percentage-stop-limit`** or **`-B`** (optional, `int`, default: `5`): On Mac laptops powered by battery, automatically stops image generation when battery percentage reaches this threshold. Prevents your Mac from shutting down and becoming unresponsive during long generation sessions.
|
- **`--battery-percentage-stop-limit`** or **`-B`** (optional, `int`, default: `5`): On Mac laptops powered by battery, automatically stops image generation when battery percentage reaches this threshold. Prevents your Mac from shutting down and becoming unresponsive during long generation sessions.
|
||||||
|
|
||||||
- **`--lora-name`** (optional, `str`, default: `None`): The name of the LoRA to download from Hugging Face.
|
|
||||||
|
|
||||||
- **`--lora-repo-id`** (optional, `str`, default: `"ali-vilab/In-Context-LoRA"`): The Hugging Face repository ID for LoRAs.
|
|
||||||
|
|
||||||
- **`--stepwise-image-output-dir`** (optional, `str`, default: `None`): [EXPERIMENTAL] Output directory to write step-wise images and their final composite image to. This feature may change in future versions. When specified, MFLUX will save an image for each denoising step, allowing you to visualize the generation process from noise to final image.
|
- **`--stepwise-image-output-dir`** (optional, `str`, default: `None`): [EXPERIMENTAL] Output directory to write step-wise images and their final composite image to. This feature may change in future versions. When specified, MFLUX will save an image for each denoising step, allowing you to visualize the generation process from noise to final image.
|
||||||
|
|
||||||
- **`--vae-tiling`** (optional, flag): Enable VAE tiling to reduce memory usage during the decoding phase. This splits the image into smaller chunks for processing, which can prevent out-of-memory errors when generating high-resolution images. Note that this optimization may occasionally produce a subtle seam in the middle of the image, but it's often worth the tradeoff for being able to generate images that would otherwise cause your system to run out of memory.
|
- **`--vae-tiling`** (optional, flag): Enable VAE tiling to reduce memory usage during the decoding phase. This splits the image into smaller chunks for processing, which can prevent out-of-memory errors when generating high-resolution images. Note that this optimization may occasionally produce a subtle seam in the middle of the image, but it's often worth the tradeoff for being able to generate images that would otherwise cause your system to run out of memory.
|
||||||
@ -701,7 +702,7 @@ However, if we were to import a fixed instance of this latent array saved from t
|
|||||||
The images below illustrate this equivalence.
|
The images below illustrate this equivalence.
|
||||||
In all cases the Schnell model was run for 2 time steps.
|
In all cases the Schnell model was run for 2 time steps.
|
||||||
The Diffusers implementation ran in CPU mode.
|
The Diffusers implementation ran in CPU mode.
|
||||||
The precision for MFLUX can be set in the [Config](src/mflux/config/config.py) class.
|
The precision for MFLUX can be set in the [Config](src/mflux/models/common/config/config.py) class.
|
||||||
There is typically a noticeable but very small difference in the final image when switching between 16bit and 32bit precision.
|
There is typically a noticeable but very small difference in the final image when switching between 16bit and 32bit precision.
|
||||||
|
|
||||||
---
|
---
|
||||||
@ -837,6 +838,9 @@ In other words, you can reclaim the 34GB diskspace (per model) by deleting the f
|
|||||||
- [akx/FLUX.1-Kontext-dev-mflux-4bit](https://huggingface.co/akx/FLUX.1-Kontext-dev-mflux-4bit)
|
- [akx/FLUX.1-Kontext-dev-mflux-4bit](https://huggingface.co/akx/FLUX.1-Kontext-dev-mflux-4bit)
|
||||||
- [filipstrand/FLUX.1-Krea-dev-mflux-4bit](https://huggingface.co/filipstrand/FLUX.1-Krea-dev-mflux-4bit)
|
- [filipstrand/FLUX.1-Krea-dev-mflux-4bit](https://huggingface.co/filipstrand/FLUX.1-Krea-dev-mflux-4bit)
|
||||||
- [filipstrand/Qwen-Image-mflux-6bit](https://huggingface.co/filipstrand/Qwen-Image-mflux-6bit)
|
- [filipstrand/Qwen-Image-mflux-6bit](https://huggingface.co/filipstrand/Qwen-Image-mflux-6bit)
|
||||||
|
- [filipstrand/Z-Image-Turbo-mflux-4bit](https://huggingface.co/filipstrand/Z-Image-Turbo-mflux-4bit)
|
||||||
|
- [briaai/Fibo-mlx-4bit](https://huggingface.co/briaai/Fibo-mlx-4bit)
|
||||||
|
- [briaai/Fibo-mlx-8bit](https://huggingface.co/briaai/Fibo-mlx-8bit)
|
||||||
|
|
||||||
|
|
||||||
Using the [community model support](#-third-party-huggingface-model-support), the quantized weights can be also be automatically downloaded when running the generate command:
|
Using the [community model support](#-third-party-huggingface-model-support), the quantized weights can be also be automatically downloaded when running the generate command:
|
||||||
@ -850,23 +854,33 @@ mflux-generate \
|
|||||||
--seed 2674888
|
--seed 2674888
|
||||||
```
|
```
|
||||||
|
|
||||||
|
```sh
|
||||||
|
mflux-generate-fibo \
|
||||||
|
--model briaai/Fibo-mlx-4bit \
|
||||||
|
--prompt-file ~/Desktop/bird.json \
|
||||||
|
--width 1024 \
|
||||||
|
--height 1024 \
|
||||||
|
--steps 20 \
|
||||||
|
--seed 42 \
|
||||||
|
```
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
### 💽 Running a non-quantized model directly from disk
|
### 💽 Running a model directly from disk
|
||||||
|
|
||||||
MFLUX also supports running a non-quantized model directly from a custom location.
|
MFLUX supports running a model directly from a custom location using the `--model` flag with a local path.
|
||||||
In the example below, the model is placed in `/Users/filipstrand/Desktop/schnell`:
|
In the example below, the model is placed in `/Users/filipstrand/Desktop/schnell`:
|
||||||
|
|
||||||
```sh
|
```sh
|
||||||
mflux-generate \
|
mflux-generate \
|
||||||
--path "/Users/filipstrand/Desktop/schnell" \
|
--model "/Users/filipstrand/Desktop/schnell" \
|
||||||
--model schnell \
|
--base-model schnell \
|
||||||
--steps 2 \
|
--steps 2 \
|
||||||
--seed 2 \
|
--seed 2 \
|
||||||
--prompt "Luxury food photograph"
|
--prompt "Luxury food photograph"
|
||||||
```
|
```
|
||||||
|
|
||||||
Note that the `--model` flag must be set when loading a model from disk.
|
When loading from a local path, use `--base-model` to specify the architecture (e.g., `schnell`, `dev`).
|
||||||
|
|
||||||
Also note that unlike when using the typical `alias` way of initializing the model (which internally handles that the required resources are downloaded),
|
Also note that unlike when using the typical `alias` way of initializing the model (which internally handles that the required resources are downloaded),
|
||||||
when loading a model directly from disk, we require the downloaded models to look like the following:
|
when loading a model directly from disk, we require the downloaded models to look like the following:
|
||||||
@ -909,17 +923,30 @@ processed a bit differently, which is why we require this structure above.*
|
|||||||
|
|
||||||
### 🌐 Third-Party HuggingFace Model Support
|
### 🌐 Third-Party HuggingFace Model Support
|
||||||
|
|
||||||
MFLUX now supports compatible third-party models from HuggingFace that follow the FLUX architecture. This opens up the ecosystem to community-created models that may offer different capabilities, sizes, or specializations.
|
MFLUX supports compatible third-party models from HuggingFace that follow the FLUX architecture. The `--model` parameter accepts:
|
||||||
|
|
||||||
To use a third-party model, specify the HuggingFace repository ID with the `--model` parameter and indicate which base architecture (dev or schnell) it's derived from using the `--base-model` parameter:
|
- **Predefined names**: `dev`, `schnell`, `fibo`, `z-image-turbo`, etc.
|
||||||
|
- **HuggingFace repos**: `Freepik/flux.1-lite-8B`, `briaai/Fibo-mlx-4bit`
|
||||||
|
- **Local paths**: `/Users/me/models/my-model`, `~/my-model`
|
||||||
|
|
||||||
|
This unified interface mirrors how LoRA paths work, making it easy to switch between local and remote models.
|
||||||
|
|
||||||
```sh
|
```sh
|
||||||
|
# Using a HuggingFace repo
|
||||||
mflux-generate \
|
mflux-generate \
|
||||||
--model Freepik/flux.1-lite-8B \
|
--model Freepik/flux.1-lite-8B \
|
||||||
--base-model schnell \
|
--base-model schnell \
|
||||||
--steps 4 \
|
--steps 4 \
|
||||||
--seed 42 \
|
--seed 42 \
|
||||||
--prompt "A beautiful landscape with mountains and a lake"
|
--prompt "A beautiful landscape with mountains and a lake"
|
||||||
|
|
||||||
|
# Using a local path
|
||||||
|
mflux-generate \
|
||||||
|
--model /Users/me/models/flux-lite \
|
||||||
|
--base-model schnell \
|
||||||
|
--steps 4 \
|
||||||
|
--seed 42 \
|
||||||
|
--prompt "A beautiful landscape with mountains and a lake"
|
||||||
```
|
```
|
||||||
|
|
||||||
Some examples of compatible third-party models include:
|
Some examples of compatible third-party models include:
|
||||||
@ -1426,6 +1453,36 @@ Note: The optional `--prompt` parameter allows you to guide the VLM's interpreta
|
|||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
|
### ⚡ Z-Image
|
||||||
|
|
||||||
|
MFLUX supports [Z-Image-Turbo](https://huggingface.co/Tongyi-MAI/Z-Image-Turbo) from Tongyi Lab (Alibaba), released in November 2025. Z-Image is an efficient 6B-parameter image generation model with a single-stream DiT architecture. Z-Image-Turbo delivers high-quality images in just 9 steps, making it one of the fastest open-source models available. All the standard modes such as, img2img, LoRA and quantizations are supported for this model. See the [technical paper](https://arxiv.org/abs/2511.22699) for more details.
|
||||||
|
|
||||||
|

|
||||||
|
|
||||||
|
#### Example
|
||||||
|
|
||||||
|
The following uses the pre-quantized 4-bit model from [filipstrand/Z-Image-Turbo-mflux-4bit](https://huggingface.co/filipstrand/Z-Image-Turbo-mflux-4bit) to generate a vibrant 1960s style image with a LoRA adapter [Technically Color](https://huggingface.co/renderartist/Technically-Color-Z-Image-Turbo) for enhanced film color:
|
||||||
|
|
||||||
|
```sh
|
||||||
|
mflux-generate-z-image-turbo \
|
||||||
|
--model filipstrand/Z-Image-Turbo-mflux-4bit \
|
||||||
|
--prompt "t3chnic4lly vibrant 1960s close-up of a woman sitting under a tree in a blue skirt and white blouse, she has blonde wavy short hair and a smile with green eyes lake scene by a garden with flowers in the foreground 1960s style film She's holding her hand out there is a small smooth frog in her palm, she's making eye contact with the toad." \
|
||||||
|
--width 1280 \
|
||||||
|
--height 720 \
|
||||||
|
--seed 456 \
|
||||||
|
--steps 9 \
|
||||||
|
--lora-paths renderartist/Technically-Color-Z-Image-Turbo \
|
||||||
|
--lora-scales 0.5
|
||||||
|
```
|
||||||
|
|
||||||
|
⚠️ *Note: Z-Image-Turbo requires downloading the `Tongyi-MAI/Z-Image-Turbo` model weights (~31GB), or use quantization for smaller sizes.*
|
||||||
|
|
||||||
|
*Dreambooth fine-tuning for Z-Image is not yet supported in MFLUX but is planned. In the meantime, you can train Z-Image-Turbo LoRAs using [AI Toolkit](https://github.com/ostris/ai-toolkit) - see [How to Train a Z-Image-Turbo LoRA with AI Toolkit](https://www.youtube.com/watch?v=Kmve1_jiDpQ) by Ostris AI.*
|
||||||
|
|
||||||
|
*For a Swift MLX implementation of Z-Image, see [zimage.swift](https://github.com/mzbac/zimage.swift) by [@mzbac](https://github.com/mzbac).*
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
### 🔌 LoRA
|
### 🔌 LoRA
|
||||||
|
|
||||||
MFLUX support loading trained [LoRA](https://huggingface.co/docs/diffusers/en/training/lora) adapters (actual training support is coming).
|
MFLUX support loading trained [LoRA](https://huggingface.co/docs/diffusers/en/training/lora) adapters (actual training support is coming).
|
||||||
@ -1469,6 +1526,28 @@ mflux-generate \
|
|||||||
Just to see the difference, this image displays the four cases: One of having both adapters fully active, partially active and no LoRA at all.
|
Just to see the difference, this image displays the four cases: One of having both adapters fully active, partially active and no LoRA at all.
|
||||||
The example above also show the usage of `--lora-scales` flag.
|
The example above also show the usage of `--lora-scales` flag.
|
||||||
|
|
||||||
|
#### HuggingFace LoRA Downloads
|
||||||
|
|
||||||
|
MFLUX can automatically download LoRAs directly from HuggingFace. Simply pass the repository ID to `--lora-paths`:
|
||||||
|
|
||||||
|
```sh
|
||||||
|
# Download from a HuggingFace repo (auto-finds the .safetensors file)
|
||||||
|
mflux-generate \
|
||||||
|
--prompt "a portrait" \
|
||||||
|
--lora-paths "author/lora-model"
|
||||||
|
```
|
||||||
|
|
||||||
|
For repositories with multiple LoRA files (collections), use the `repo_id:filename` format to specify which file to download:
|
||||||
|
|
||||||
|
```sh
|
||||||
|
# Download a specific file from a collection
|
||||||
|
mflux-generate \
|
||||||
|
--prompt "film storyboard style, a cat" \
|
||||||
|
--lora-paths "ali-vilab/In-Context-LoRA:film-storyboard.safetensors"
|
||||||
|
```
|
||||||
|
|
||||||
|
Downloaded LoRAs are cached locally and reused on subsequent runs.
|
||||||
|
|
||||||
#### LoRA Library Path
|
#### LoRA Library Path
|
||||||
|
|
||||||
MFLUX supports a convenient LoRA library feature that allows you to reference LoRA files by their basename instead of full paths. This is particularly useful when you have a collection of LoRA files organized in one or more directories.
|
MFLUX supports a convenient LoRA library feature that allows you to reference LoRA files by their basename instead of full paths. This is particularly useful when you have a collection of LoRA files organized in one or more directories.
|
||||||
|
|||||||
@ -3,21 +3,18 @@ requires = ["uv_build>=0.7.19,<0.8.0"]
|
|||||||
build-backend = "uv_build"
|
build-backend = "uv_build"
|
||||||
|
|
||||||
[tool.uv.build-backend]
|
[tool.uv.build-backend]
|
||||||
default-excludes = true # __pycache__, *.pyc, and *.pyo
|
default-excludes = true
|
||||||
module-name = "mflux"
|
module-name = "mflux"
|
||||||
namespace = true
|
namespace = true
|
||||||
source-exclude = [
|
source-exclude = [
|
||||||
# documentation assets were 27MB on 2025-07-05
|
|
||||||
"**/assets/**",
|
"**/assets/**",
|
||||||
# dreambooth examples ~=5MB on 2025-07-05
|
|
||||||
"**/models/flux/variants/dreambooth/_example/images/**",
|
"**/models/flux/variants/dreambooth/_example/images/**",
|
||||||
# loss pdf/tex does not need to be distributed
|
|
||||||
"**/optimization/_loss_derivation/**",
|
"**/optimization/_loss_derivation/**",
|
||||||
]
|
]
|
||||||
|
|
||||||
[project]
|
[project]
|
||||||
name = "mflux"
|
name = "mflux"
|
||||||
version = "0.12.1"
|
version = "0.13.0.dev0"
|
||||||
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"]
|
||||||
@ -42,9 +39,8 @@ dependencies = [
|
|||||||
"requests>=2.32.4",
|
"requests>=2.32.4",
|
||||||
"safetensors>=0.4.4,<1.0",
|
"safetensors>=0.4.4,<1.0",
|
||||||
"sentencepiece>=0.2.0,<1.0; python_version<'3.13'",
|
"sentencepiece>=0.2.0,<1.0; python_version<'3.13'",
|
||||||
# sentencepiece 0.2.1 is first release with 3.13 and 3.14 wheels: https://pypi.org/project/sentencepiece/0.2.1/
|
|
||||||
"sentencepiece>=0.2.1,<1.0; python_version>='3.13'",
|
"sentencepiece>=0.2.1,<1.0; python_version>='3.13'",
|
||||||
"tokenizers>=0.20.3; python_version>='3.13'", # transformers -> tokenizers
|
"tokenizers>=0.20.3; python_version>='3.13'",
|
||||||
"toml>=0.10.2,<1.0",
|
"toml>=0.10.2,<1.0",
|
||||||
"torch>=2.7.1",
|
"torch>=2.7.1",
|
||||||
"torch>=2.3.1,<3.0; python_version<'3.13'",
|
"torch>=2.3.1,<3.0; python_version<'3.13'",
|
||||||
@ -68,36 +64,37 @@ dev = [
|
|||||||
"matplotlib>3.10,<4.0",
|
"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.29.2", # Used ONLY during test runs to ensure deterministic test results
|
"mlx==0.29.2",
|
||||||
]
|
]
|
||||||
|
|
||||||
[project.urls]
|
[project.urls]
|
||||||
homepage = "https://github.com/filipstrand/mflux"
|
homepage = "https://github.com/filipstrand/mflux"
|
||||||
|
|
||||||
[project.scripts]
|
[project.scripts]
|
||||||
mflux-generate = "mflux.generate:main"
|
mflux-generate = "mflux.models.flux.cli.flux_generate:main"
|
||||||
mflux-generate-controlnet = "mflux.generate_controlnet:main"
|
mflux-generate-controlnet = "mflux.models.flux.cli.flux_generate_controlnet:main"
|
||||||
mflux-generate-in-context = "mflux.generate_in_context_dev:main"
|
mflux-generate-in-context = "mflux.models.flux.cli.flux_generate_in_context_dev:main"
|
||||||
mflux-generate-in-context-edit = "mflux.generate_in_context_edit:main"
|
mflux-generate-in-context-edit = "mflux.models.flux.cli.flux_generate_in_context_edit:main"
|
||||||
mflux-generate-in-context-catvton = "mflux.generate_in_context_catvton:main"
|
mflux-generate-in-context-catvton = "mflux.models.flux.cli.flux_generate_in_context_catvton:main"
|
||||||
mflux-generate-fill = "mflux.generate_fill:main"
|
mflux-generate-fill = "mflux.models.flux.cli.flux_generate_fill:main"
|
||||||
mflux-generate-depth = "mflux.generate_depth:main"
|
mflux-generate-depth = "mflux.models.flux.cli.flux_generate_depth:main"
|
||||||
mflux-generate-redux = "mflux.generate_redux:main"
|
mflux-generate-redux = "mflux.models.flux.cli.flux_generate_redux:main"
|
||||||
mflux-generate-kontext = "mflux.generate_kontext:main"
|
mflux-generate-kontext = "mflux.models.flux.cli.flux_generate_kontext:main"
|
||||||
mflux-generate-qwen = "mflux.generate_qwen:main"
|
mflux-generate-qwen = "mflux.models.qwen.cli.qwen_image_generate:main"
|
||||||
mflux-generate-qwen-edit = "mflux.generate_qwen_edit:main"
|
mflux-generate-qwen-edit = "mflux.models.qwen.cli.qwen_image_edit_generate:main"
|
||||||
mflux-generate-fibo = "mflux.generate_fibo:main"
|
mflux-generate-fibo = "mflux.models.fibo.cli.fibo_generate:main"
|
||||||
mflux-refine-fibo = "mflux.refine_fibo:main"
|
mflux-generate-z-image-turbo = "mflux.models.z_image.cli.z_image_turbo_generate:main"
|
||||||
mflux-inspire-fibo = "mflux.inspire_fibo:main"
|
mflux-refine-fibo = "mflux.models.fibo_vlm.cli.fibo_refine:main"
|
||||||
mflux-concept = "mflux.concept:main"
|
mflux-inspire-fibo = "mflux.models.fibo_vlm.cli.fibo_inspire:main"
|
||||||
mflux-concept-from-image = "mflux.concept_from_image:main"
|
mflux-concept = "mflux.models.flux.cli.flux_concept:main"
|
||||||
mflux-save = "mflux.save:main"
|
mflux-concept-from-image = "mflux.models.flux.cli.flux_concept_from_image:main"
|
||||||
mflux-save-depth = "mflux.save_depth:main"
|
mflux-save = "mflux.models.common.cli.save:main"
|
||||||
mflux-train = "mflux.train:main"
|
mflux-save-depth = "mflux.models.depth_pro.cli.save_depth:main"
|
||||||
mflux-upscale = "mflux.upscale:main"
|
mflux-train = "mflux.models.common.cli.train:main"
|
||||||
mflux-lora-library = "mflux.lora_library:main"
|
mflux-upscale = "mflux.models.flux.cli.flux_upscale:main"
|
||||||
mflux-info = "mflux.info:main"
|
mflux-lora-library = "mflux.models.common.cli.lora_library:main"
|
||||||
mflux-completions = "mflux.ui.cli.completions.install:main"
|
mflux-info = "mflux.models.common.cli.info:main"
|
||||||
|
mflux-completions = "mflux.cli.completions.install:main"
|
||||||
|
|
||||||
|
|
||||||
[tool.ruff]
|
[tool.ruff]
|
||||||
@ -107,44 +104,18 @@ target-version = "py310"
|
|||||||
respect-gitignore = true
|
respect-gitignore = true
|
||||||
|
|
||||||
[tool.ruff.lint]
|
[tool.ruff.lint]
|
||||||
# Enable Pyflakes (`F`) and a subset of the pycodestyle (`E`) codes by default.
|
|
||||||
# Unlike Flake8, Ruff doesn't enable pycodestyle warnings (`W`) or
|
|
||||||
# McCabe complexity (`C901`) by default.
|
|
||||||
select = ["BLE", "E4", "E7", "E9", "F", "I", "ICN", "LOG", "PERF", "W"]
|
select = ["BLE", "E4", "E7", "E9", "F", "I", "ICN", "LOG", "PERF", "W"]
|
||||||
ignore = []
|
ignore = []
|
||||||
|
|
||||||
# Allow fix for all enabled rules (when `--fix`) is provided.
|
|
||||||
fixable = ["ALL"]
|
fixable = ["ALL"]
|
||||||
unfixable = []
|
unfixable = []
|
||||||
|
|
||||||
# Allow unused variables when underscore-prefixed.
|
|
||||||
dummy-variable-rgx = "^(_+|(_+[a-zA-Z0-9_]*[a-zA-Z0-9]+?))$"
|
dummy-variable-rgx = "^(_+|(_+[a-zA-Z0-9_]*[a-zA-Z0-9]+?))$"
|
||||||
|
|
||||||
[tool.ruff.format]
|
[tool.ruff.format]
|
||||||
# Like Black, use double quotes for strings.
|
|
||||||
quote-style = "double"
|
quote-style = "double"
|
||||||
|
|
||||||
# Like Black, indent with spaces, rather than tabs.
|
|
||||||
indent-style = "space"
|
indent-style = "space"
|
||||||
|
|
||||||
# Like Black, respect magic trailing commas.
|
|
||||||
skip-magic-trailing-comma = false
|
skip-magic-trailing-comma = false
|
||||||
|
|
||||||
# Like Black, automatically detect the appropriate line ending.
|
|
||||||
line-ending = "auto"
|
line-ending = "auto"
|
||||||
|
|
||||||
# Enable auto-formatting of code examples in docstrings. Markdown,
|
|
||||||
# reStructuredText code/literal blocks and doctests are all supported.
|
|
||||||
#
|
|
||||||
# This is currently disabled by default, but it is planned for this
|
|
||||||
# to be opt-out in the future.
|
|
||||||
docstring-code-format = false
|
docstring-code-format = false
|
||||||
|
|
||||||
# Set the line length limit used when formatting code snippets in
|
|
||||||
# docstrings.
|
|
||||||
#
|
|
||||||
# This only has an effect when the `docstring-code-format` setting is
|
|
||||||
# enabled.
|
|
||||||
docstring-code-line-length = "dynamic"
|
docstring-code-line-length = "dynamic"
|
||||||
|
|
||||||
[tool.pytest.ini_options]
|
[tool.pytest.ini_options]
|
||||||
@ -155,7 +126,6 @@ markers = [
|
|||||||
"high_memory_requirement: marks tests that require high memory (deselect with '-m \"not high_memory_requirement\"')",
|
"high_memory_requirement: marks tests that require high memory (deselect with '-m \"not high_memory_requirement\"')",
|
||||||
]
|
]
|
||||||
|
|
||||||
# https://docs.astral.sh/ruff/settings/#lintisort
|
|
||||||
[tool.ruff.lint.isort]
|
[tool.ruff.lint.isort]
|
||||||
case-sensitive = false
|
case-sensitive = false
|
||||||
combine-as-imports = true
|
combine-as-imports = true
|
||||||
@ -173,7 +143,6 @@ section-order = [
|
|||||||
|
|
||||||
[tool.mypy]
|
[tool.mypy]
|
||||||
disable_error_code = [
|
disable_error_code = [
|
||||||
# TODO: for each error code - clean them up in a single PR then remove entry
|
|
||||||
"annotation-unchecked",
|
"annotation-unchecked",
|
||||||
"arg-type",
|
"arg-type",
|
||||||
"assignment",
|
"assignment",
|
||||||
@ -188,5 +157,4 @@ disable_error_code = [
|
|||||||
error_summary = true
|
error_summary = true
|
||||||
ignore_missing_imports = true
|
ignore_missing_imports = true
|
||||||
implicit_optional = true
|
implicit_optional = true
|
||||||
# we should feel free to use new hinting features even if we back support to earlier versions
|
|
||||||
python_version = "3.12"
|
python_version = "3.12"
|
||||||
|
|||||||
@ -4,3 +4,8 @@ import os
|
|||||||
# This must be set before any tokenizers are imported/used
|
# This must be set before any tokenizers are imported/used
|
||||||
if "TOKENIZERS_PARALLELISM" not in os.environ:
|
if "TOKENIZERS_PARALLELISM" not in os.environ:
|
||||||
os.environ["TOKENIZERS_PARALLELISM"] = "false"
|
os.environ["TOKENIZERS_PARALLELISM"] = "false"
|
||||||
|
|
||||||
|
# Export main classes for convenient import
|
||||||
|
from mflux.models.flux.variants.txt2img.flux import Flux1
|
||||||
|
|
||||||
|
__all__ = ["Flux1"]
|
||||||
|
|||||||
BIN
src/mflux/assets/z_image_turbo_example.jpg
Normal file
BIN
src/mflux/assets/z_image_turbo_example.jpg
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 1.3 MiB |
@ -4,7 +4,7 @@ import mlx.core as mx
|
|||||||
import PIL.Image
|
import PIL.Image
|
||||||
import tqdm
|
import tqdm
|
||||||
|
|
||||||
from mflux.config.runtime_config import RuntimeConfig
|
from mflux.models.common.config.config import Config
|
||||||
|
|
||||||
|
|
||||||
class BeforeLoopCallback(Protocol):
|
class BeforeLoopCallback(Protocol):
|
||||||
@ -13,7 +13,7 @@ class BeforeLoopCallback(Protocol):
|
|||||||
seed: int,
|
seed: int,
|
||||||
prompt: str,
|
prompt: str,
|
||||||
latents: mx.array,
|
latents: mx.array,
|
||||||
config: RuntimeConfig,
|
config: Config,
|
||||||
canny_image: PIL.Image.Image | None = None,
|
canny_image: PIL.Image.Image | None = None,
|
||||||
depth_image: PIL.Image.Image | None = None,
|
depth_image: PIL.Image.Image | None = None,
|
||||||
) -> None: ...
|
) -> None: ...
|
||||||
@ -26,7 +26,7 @@ class InLoopCallback(Protocol):
|
|||||||
seed: int,
|
seed: int,
|
||||||
prompt: str,
|
prompt: str,
|
||||||
latents: mx.array,
|
latents: mx.array,
|
||||||
config: RuntimeConfig,
|
config: Config,
|
||||||
time_steps: tqdm,
|
time_steps: tqdm,
|
||||||
) -> None: ...
|
) -> None: ...
|
||||||
|
|
||||||
@ -37,7 +37,7 @@ class AfterLoopCallback(Protocol):
|
|||||||
seed: int,
|
seed: int,
|
||||||
prompt: str,
|
prompt: str,
|
||||||
latents: mx.array,
|
latents: mx.array,
|
||||||
config: RuntimeConfig,
|
config: Config,
|
||||||
) -> None: ...
|
) -> None: ...
|
||||||
|
|
||||||
|
|
||||||
@ -48,6 +48,6 @@ class InterruptCallback(Protocol):
|
|||||||
seed: int,
|
seed: int,
|
||||||
prompt: str,
|
prompt: str,
|
||||||
latents: mx.array,
|
latents: mx.array,
|
||||||
config: RuntimeConfig,
|
config: Config,
|
||||||
time_steps: tqdm,
|
time_steps: tqdm,
|
||||||
) -> None: ...
|
) -> None: ...
|
||||||
|
|||||||
@ -1,6 +1,5 @@
|
|||||||
from argparse import Namespace
|
from argparse import Namespace
|
||||||
|
|
||||||
from mflux.callbacks.callback_registry import CallbackRegistry
|
|
||||||
from mflux.callbacks.instances.battery_saver import BatterySaver
|
from mflux.callbacks.instances.battery_saver import BatterySaver
|
||||||
from mflux.callbacks.instances.canny_saver import CannyImageSaver
|
from mflux.callbacks.instances.canny_saver import CannyImageSaver
|
||||||
from mflux.callbacks.instances.depth_saver import DepthImageSaver
|
from mflux.callbacks.instances.depth_saver import DepthImageSaver
|
||||||
@ -13,32 +12,33 @@ class CallbackManager:
|
|||||||
def register_callbacks(
|
def register_callbacks(
|
||||||
args: Namespace,
|
args: Namespace,
|
||||||
model,
|
model,
|
||||||
|
latent_creator,
|
||||||
enable_canny_saver: bool = False,
|
enable_canny_saver: bool = False,
|
||||||
enable_depth_saver: bool = False,
|
enable_depth_saver: bool = False,
|
||||||
) -> MemorySaver | None:
|
) -> MemorySaver | None:
|
||||||
# Battery saver (always enabled)
|
# Battery saver (always enabled)
|
||||||
CallbackManager._register_battery_saver(args)
|
CallbackManager._register_battery_saver(args, model)
|
||||||
|
|
||||||
# VAE Tiling (if requested)
|
# VAE Tiling (if requested)
|
||||||
CallbackManager._register_vae_tiling(args, model)
|
CallbackManager._register_vae_tiling(args, model)
|
||||||
|
|
||||||
# Specialized savers (based on flags)
|
# Specialized savers (based on flags)
|
||||||
if enable_canny_saver:
|
if enable_canny_saver:
|
||||||
CallbackManager._register_canny_saver(args)
|
CallbackManager._register_canny_saver(args, model)
|
||||||
|
|
||||||
if enable_depth_saver:
|
if enable_depth_saver:
|
||||||
CallbackManager._register_depth_saver(args)
|
CallbackManager._register_depth_saver(args, model)
|
||||||
|
|
||||||
# Stepwise handler (if requested)
|
# Stepwise handler (if requested)
|
||||||
CallbackManager._register_stepwise_handler(args, model)
|
CallbackManager._register_stepwise_handler(args, model, latent_creator)
|
||||||
|
|
||||||
# Memory saver (if requested)
|
# Memory saver (if requested)
|
||||||
return CallbackManager._register_memory_saver(args, model)
|
return CallbackManager._register_memory_saver(args, model)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _register_battery_saver(args: Namespace) -> None:
|
def _register_battery_saver(args: Namespace, model) -> None:
|
||||||
battery_saver = BatterySaver(battery_percentage_stop_limit=args.battery_percentage_stop_limit)
|
battery_saver = BatterySaver(battery_percentage_stop_limit=args.battery_percentage_stop_limit)
|
||||||
CallbackRegistry.register_before_loop(battery_saver)
|
model.callbacks.register(battery_saver)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _register_vae_tiling(args: Namespace, model) -> None:
|
def _register_vae_tiling(args: Namespace, model) -> None:
|
||||||
@ -47,31 +47,31 @@ class CallbackManager:
|
|||||||
model.vae.decoder.split_direction = args.vae_tiling_split
|
model.vae.decoder.split_direction = args.vae_tiling_split
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _register_canny_saver(args: Namespace) -> None:
|
def _register_canny_saver(args: Namespace, model) -> None:
|
||||||
if hasattr(args, "controlnet_save_canny") and args.controlnet_save_canny:
|
if hasattr(args, "controlnet_save_canny") and args.controlnet_save_canny:
|
||||||
canny_image_saver = CannyImageSaver(path=args.output)
|
canny_image_saver = CannyImageSaver(path=args.output)
|
||||||
CallbackRegistry.register_before_loop(canny_image_saver)
|
model.callbacks.register(canny_image_saver)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _register_depth_saver(args: Namespace) -> None:
|
def _register_depth_saver(args: Namespace, model) -> None:
|
||||||
if hasattr(args, "save_depth_map") and args.save_depth_map:
|
if hasattr(args, "save_depth_map") and args.save_depth_map:
|
||||||
depth_image_saver = DepthImageSaver(path=args.output)
|
depth_image_saver = DepthImageSaver(path=args.output)
|
||||||
CallbackRegistry.register_before_loop(depth_image_saver)
|
model.callbacks.register(depth_image_saver)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _register_stepwise_handler(args: Namespace, model) -> None:
|
def _register_stepwise_handler(args: Namespace, model, latent_creator) -> None:
|
||||||
if args.stepwise_image_output_dir:
|
if args.stepwise_image_output_dir:
|
||||||
handler = StepwiseHandler(model=model, output_dir=args.stepwise_image_output_dir)
|
handler = StepwiseHandler(
|
||||||
CallbackRegistry.register_before_loop(handler)
|
model=model,
|
||||||
CallbackRegistry.register_in_loop(handler)
|
latent_creator=latent_creator,
|
||||||
CallbackRegistry.register_interrupt(handler)
|
output_dir=args.stepwise_image_output_dir,
|
||||||
|
)
|
||||||
|
model.callbacks.register(handler)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _register_memory_saver(args: Namespace, model) -> MemorySaver | None:
|
def _register_memory_saver(args: Namespace, model) -> MemorySaver | None:
|
||||||
memory_saver = None
|
memory_saver = None
|
||||||
if args.low_ram:
|
if args.low_ram:
|
||||||
memory_saver = MemorySaver(model=model, keep_transformer=len(args.seed) > 1)
|
memory_saver = MemorySaver(model=model, keep_transformer=len(args.seed) > 1)
|
||||||
CallbackRegistry.register_before_loop(memory_saver)
|
model.callbacks.register(memory_saver)
|
||||||
CallbackRegistry.register_in_loop(memory_saver)
|
|
||||||
CallbackRegistry.register_after_loop(memory_saver)
|
|
||||||
return memory_saver
|
return memory_saver
|
||||||
|
|||||||
@ -1,40 +1,44 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
from mflux.callbacks.callback import AfterLoopCallback, BeforeLoopCallback, InLoopCallback, InterruptCallback
|
from mflux.callbacks.callback import AfterLoopCallback, BeforeLoopCallback, InLoopCallback, InterruptCallback
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from mflux.callbacks.generation_context import GenerationContext
|
||||||
|
from mflux.models.common.config.config import Config
|
||||||
|
|
||||||
|
|
||||||
class CallbackRegistry:
|
class CallbackRegistry:
|
||||||
in_loop = []
|
def __init__(self):
|
||||||
before_loop = []
|
self.in_loop: list[InLoopCallback] = []
|
||||||
interrupt = []
|
self.before_loop: list[BeforeLoopCallback] = []
|
||||||
after_loop = []
|
self.interrupt: list[InterruptCallback] = []
|
||||||
|
self.after_loop: list[AfterLoopCallback] = []
|
||||||
|
|
||||||
@staticmethod
|
def register(self, callback) -> None:
|
||||||
def register_in_loop(callback: InLoopCallback) -> None:
|
if hasattr(callback, "call_before_loop"):
|
||||||
CallbackRegistry.in_loop.append(callback)
|
self.before_loop.append(callback)
|
||||||
|
if hasattr(callback, "call_in_loop"):
|
||||||
|
self.in_loop.append(callback)
|
||||||
|
if hasattr(callback, "call_after_loop"):
|
||||||
|
self.after_loop.append(callback)
|
||||||
|
if hasattr(callback, "call_interrupt"):
|
||||||
|
self.interrupt.append(callback)
|
||||||
|
|
||||||
@staticmethod
|
def start(self, seed: int, prompt: str, config: Config) -> GenerationContext:
|
||||||
def register_before_loop(callback: BeforeLoopCallback) -> None:
|
from mflux.callbacks.generation_context import GenerationContext
|
||||||
CallbackRegistry.before_loop.append(callback)
|
|
||||||
|
|
||||||
@staticmethod
|
return GenerationContext(self, seed, prompt, config)
|
||||||
def register_after_loop(callback: AfterLoopCallback) -> None:
|
|
||||||
CallbackRegistry.after_loop.append(callback)
|
|
||||||
|
|
||||||
@staticmethod
|
def before_loop_callbacks(self) -> list[BeforeLoopCallback]:
|
||||||
def register_interrupt(callback: InterruptCallback) -> None:
|
return self.before_loop
|
||||||
CallbackRegistry.interrupt.append(callback)
|
|
||||||
|
|
||||||
@staticmethod
|
def in_loop_callbacks(self) -> list[InLoopCallback]:
|
||||||
def before_loop_callbacks() -> list[BeforeLoopCallback]:
|
return self.in_loop
|
||||||
return CallbackRegistry.before_loop
|
|
||||||
|
|
||||||
@staticmethod
|
def after_loop_callbacks(self) -> list[AfterLoopCallback]:
|
||||||
def in_loop_callbacks() -> list[InLoopCallback]:
|
return self.after_loop
|
||||||
return CallbackRegistry.in_loop
|
|
||||||
|
|
||||||
@staticmethod
|
def interrupt_callbacks(self) -> list[InterruptCallback]:
|
||||||
def after_loop_callbacks() -> list[AfterLoopCallback]:
|
return self.interrupt
|
||||||
return CallbackRegistry.after_loop
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def interrupt_callbacks() -> list[InterruptCallback]:
|
|
||||||
return CallbackRegistry.interrupt
|
|
||||||
|
|||||||
@ -3,20 +3,21 @@ import PIL.Image
|
|||||||
import tqdm
|
import tqdm
|
||||||
|
|
||||||
from mflux.callbacks.callback_registry import CallbackRegistry
|
from mflux.callbacks.callback_registry import CallbackRegistry
|
||||||
from mflux.config.runtime_config import RuntimeConfig
|
from mflux.models.common.config.config import Config
|
||||||
|
|
||||||
|
|
||||||
class Callbacks:
|
class Callbacks:
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def before_loop(
|
def before_loop(
|
||||||
|
registry: CallbackRegistry,
|
||||||
seed: int,
|
seed: int,
|
||||||
prompt: str,
|
prompt: str,
|
||||||
latents: mx.array,
|
latents: mx.array,
|
||||||
config: RuntimeConfig,
|
config: Config,
|
||||||
canny_image: PIL.Image.Image | None = None,
|
canny_image: PIL.Image.Image | None = None,
|
||||||
depth_image: PIL.Image.Image | None = None,
|
depth_image: PIL.Image.Image | None = None,
|
||||||
):
|
):
|
||||||
for subscriber in CallbackRegistry.before_loop_callbacks():
|
for subscriber in registry.before_loop_callbacks():
|
||||||
subscriber.call_before_loop(
|
subscriber.call_before_loop(
|
||||||
seed=seed,
|
seed=seed,
|
||||||
prompt=prompt,
|
prompt=prompt,
|
||||||
@ -28,14 +29,15 @@ class Callbacks:
|
|||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def in_loop(
|
def in_loop(
|
||||||
|
registry: CallbackRegistry,
|
||||||
t: int,
|
t: int,
|
||||||
seed: int,
|
seed: int,
|
||||||
prompt: str,
|
prompt: str,
|
||||||
latents: mx.array,
|
latents: mx.array,
|
||||||
config: RuntimeConfig,
|
config: Config,
|
||||||
time_steps: tqdm,
|
time_steps: tqdm,
|
||||||
):
|
):
|
||||||
for subscriber in CallbackRegistry.in_loop_callbacks():
|
for subscriber in registry.in_loop_callbacks():
|
||||||
subscriber.call_in_loop(
|
subscriber.call_in_loop(
|
||||||
t=t,
|
t=t,
|
||||||
seed=seed,
|
seed=seed,
|
||||||
@ -47,24 +49,26 @@ class Callbacks:
|
|||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def after_loop(
|
def after_loop(
|
||||||
|
registry: CallbackRegistry,
|
||||||
seed: int,
|
seed: int,
|
||||||
prompt: str,
|
prompt: str,
|
||||||
latents: mx.array,
|
latents: mx.array,
|
||||||
config: RuntimeConfig,
|
config: Config,
|
||||||
):
|
):
|
||||||
for subscriber in CallbackRegistry.after_loop_callbacks():
|
for subscriber in registry.after_loop_callbacks():
|
||||||
subscriber.call_after_loop(seed=seed, prompt=prompt, latents=latents, config=config)
|
subscriber.call_after_loop(seed=seed, prompt=prompt, latents=latents, config=config)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def interruption(
|
def interruption(
|
||||||
|
registry: CallbackRegistry,
|
||||||
t: int,
|
t: int,
|
||||||
seed: int,
|
seed: int,
|
||||||
prompt: str,
|
prompt: str,
|
||||||
latents: mx.array,
|
latents: mx.array,
|
||||||
config: RuntimeConfig,
|
config: Config,
|
||||||
time_steps: tqdm,
|
time_steps: tqdm,
|
||||||
):
|
):
|
||||||
for subscriber in CallbackRegistry.interrupt_callbacks():
|
for subscriber in registry.interrupt_callbacks():
|
||||||
subscriber.call_interrupt(
|
subscriber.call_interrupt(
|
||||||
t=t,
|
t=t,
|
||||||
seed=seed,
|
seed=seed,
|
||||||
|
|||||||
75
src/mflux/callbacks/generation_context.py
Normal file
75
src/mflux/callbacks/generation_context.py
Normal file
@ -0,0 +1,75 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
import mlx.core as mx
|
||||||
|
import PIL.Image
|
||||||
|
import tqdm
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from mflux.callbacks.callback_registry import CallbackRegistry
|
||||||
|
from mflux.models.common.config.config import Config
|
||||||
|
|
||||||
|
|
||||||
|
class GenerationContext:
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
registry: CallbackRegistry,
|
||||||
|
seed: int,
|
||||||
|
prompt: str,
|
||||||
|
config: Config,
|
||||||
|
):
|
||||||
|
self._registry = registry
|
||||||
|
self._seed = seed
|
||||||
|
self._prompt = prompt
|
||||||
|
self._config = config
|
||||||
|
|
||||||
|
def before_loop(
|
||||||
|
self,
|
||||||
|
latents: mx.array,
|
||||||
|
*,
|
||||||
|
canny_image: PIL.Image.Image | None = None,
|
||||||
|
depth_image: PIL.Image.Image | None = None,
|
||||||
|
) -> None:
|
||||||
|
for subscriber in self._registry.before_loop_callbacks():
|
||||||
|
subscriber.call_before_loop(
|
||||||
|
seed=self._seed,
|
||||||
|
prompt=self._prompt,
|
||||||
|
latents=latents,
|
||||||
|
config=self._config,
|
||||||
|
canny_image=canny_image,
|
||||||
|
depth_image=depth_image,
|
||||||
|
)
|
||||||
|
|
||||||
|
def in_loop(self, t: int, latents: mx.array, time_steps: tqdm = None) -> None:
|
||||||
|
time_steps = time_steps or self._config.time_steps
|
||||||
|
for subscriber in self._registry.in_loop_callbacks():
|
||||||
|
subscriber.call_in_loop(
|
||||||
|
t=t,
|
||||||
|
seed=self._seed,
|
||||||
|
prompt=self._prompt,
|
||||||
|
latents=latents,
|
||||||
|
config=self._config,
|
||||||
|
time_steps=time_steps,
|
||||||
|
)
|
||||||
|
|
||||||
|
def after_loop(self, latents: mx.array) -> None:
|
||||||
|
for subscriber in self._registry.after_loop_callbacks():
|
||||||
|
subscriber.call_after_loop(
|
||||||
|
seed=self._seed,
|
||||||
|
prompt=self._prompt,
|
||||||
|
latents=latents,
|
||||||
|
config=self._config,
|
||||||
|
)
|
||||||
|
|
||||||
|
def interruption(self, t: int, latents: mx.array, time_steps: tqdm = None) -> None:
|
||||||
|
time_steps = time_steps or self._config.time_steps
|
||||||
|
for subscriber in self._registry.interrupt_callbacks():
|
||||||
|
subscriber.call_interrupt(
|
||||||
|
t=t,
|
||||||
|
seed=self._seed,
|
||||||
|
prompt=self._prompt,
|
||||||
|
latents=latents,
|
||||||
|
config=self._config,
|
||||||
|
time_steps=time_steps,
|
||||||
|
)
|
||||||
@ -6,57 +6,67 @@ import subprocess
|
|||||||
from mflux.callbacks.callback import BeforeLoopCallback
|
from mflux.callbacks.callback import BeforeLoopCallback
|
||||||
from mflux.utils.exceptions import StopImageGenerationException
|
from mflux.utils.exceptions import StopImageGenerationException
|
||||||
|
|
||||||
PMSET_AC_POWER_STATUS = "Now drawing from 'AC Power'"
|
|
||||||
PMSET_BATT_STATUS_PATTERN = r"InternalBattery-.+?(\d+)%"
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
def _get_machine_model() -> str:
|
|
||||||
"""Get the Mac machine model using system_profiler."""
|
|
||||||
try:
|
|
||||||
result = subprocess.run(
|
|
||||||
["system_profiler", "-json", "SPHardwareDataType"], capture_output=True, text=True, check=True
|
|
||||||
)
|
|
||||||
data = json.loads(result.stdout)
|
|
||||||
return data["SPHardwareDataType"][0]["machine_model"]
|
|
||||||
except (subprocess.CalledProcessError, json.JSONDecodeError, IndexError, KeyError) as e:
|
|
||||||
logger.warning(f"Cannot determine machine model via 'system_profiler -json SPHardwareDataType': {e}")
|
|
||||||
return "Unknown"
|
|
||||||
|
|
||||||
|
|
||||||
MACHINE_MODEL = _get_machine_model()
|
|
||||||
# assumption: all Apple Silicon models powered by battery are "MacBook"s
|
|
||||||
MACHINE_IS_BATTERY_POWERED = "MacBook" in MACHINE_MODEL
|
|
||||||
|
|
||||||
|
|
||||||
def get_battery_percentage() -> int | None:
|
|
||||||
"""Get the current battery percentage of a battery-powered Mac.
|
|
||||||
Returns None if Mac is not a battery-powered machine."""
|
|
||||||
if not MACHINE_IS_BATTERY_POWERED:
|
|
||||||
return None
|
|
||||||
percentage = None
|
|
||||||
try:
|
|
||||||
# running the subprocess would be expensive in a tight loop
|
|
||||||
# but in mflux use case, we would call this only once every
|
|
||||||
# few minutes due to N-minutes-long generation times
|
|
||||||
result = subprocess.run(["pmset", "-g", "batt"], capture_output=True, text=True, check=True)
|
|
||||||
if PMSET_AC_POWER_STATUS not in result.stdout:
|
|
||||||
if match := re.search(PMSET_BATT_STATUS_PATTERN, result.stdout):
|
|
||||||
percentage = int(match.group(1))
|
|
||||||
except (subprocess.CalledProcessError, TypeError) as e:
|
|
||||||
logger.warning(
|
|
||||||
f"Cannot read battery percentage via 'pmset -g batt': {e}. Battery saver functionality is disabled and the program will continue running."
|
|
||||||
)
|
|
||||||
|
|
||||||
return percentage
|
|
||||||
|
|
||||||
|
|
||||||
class BatterySaver(BeforeLoopCallback):
|
class BatterySaver(BeforeLoopCallback):
|
||||||
def __init__(self, battery_percentage_stop_limit=10):
|
PMSET_AC_POWER_STATUS = "Now drawing from 'AC Power'"
|
||||||
|
PMSET_BATT_STATUS_PATTERN = r"InternalBattery-.+?(\d+)%"
|
||||||
|
|
||||||
|
_machine_model: str | None = None
|
||||||
|
_is_battery_powered: bool | None = None
|
||||||
|
|
||||||
|
def __init__(self, battery_percentage_stop_limit: int = 10):
|
||||||
self.limit = battery_percentage_stop_limit
|
self.limit = battery_percentage_stop_limit
|
||||||
|
|
||||||
def call_before_loop(self, **kwargs) -> None: # type: ignore
|
def call_before_loop(self, **kwargs) -> None: # type: ignore
|
||||||
current_pct: int | None = get_battery_percentage()
|
current_pct = self._get_battery_percentage()
|
||||||
if current_pct is not None and current_pct <= self.limit:
|
if current_pct is not None and current_pct <= self.limit:
|
||||||
raise StopImageGenerationException(f"Battery below {self.limit}% threshold: {current_pct}%")
|
raise StopImageGenerationException(f"Battery below {self.limit}% threshold: {current_pct}%")
|
||||||
|
|
||||||
|
def _get_battery_percentage(self) -> int | None:
|
||||||
|
if not self._is_machine_battery_powered():
|
||||||
|
return None
|
||||||
|
|
||||||
|
percentage = None
|
||||||
|
try:
|
||||||
|
result = subprocess.run(
|
||||||
|
["pmset", "-g", "batt"],
|
||||||
|
capture_output=True,
|
||||||
|
text=True,
|
||||||
|
check=True,
|
||||||
|
)
|
||||||
|
if self.PMSET_AC_POWER_STATUS not in result.stdout:
|
||||||
|
if match := re.search(self.PMSET_BATT_STATUS_PATTERN, result.stdout):
|
||||||
|
percentage = int(match.group(1))
|
||||||
|
except (subprocess.CalledProcessError, TypeError) as e:
|
||||||
|
logger.warning(
|
||||||
|
f"Cannot read battery percentage via 'pmset -g batt': {e}. "
|
||||||
|
f"Battery saver functionality is disabled and the program will continue running."
|
||||||
|
)
|
||||||
|
|
||||||
|
return percentage
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _is_machine_battery_powered(cls) -> bool:
|
||||||
|
if cls._is_battery_powered is None:
|
||||||
|
machine_model = cls._get_machine_model()
|
||||||
|
cls._is_battery_powered = "MacBook" in machine_model
|
||||||
|
return cls._is_battery_powered
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _get_machine_model(cls) -> str:
|
||||||
|
if cls._machine_model is None:
|
||||||
|
try:
|
||||||
|
result = subprocess.run(
|
||||||
|
["system_profiler", "-json", "SPHardwareDataType"],
|
||||||
|
capture_output=True,
|
||||||
|
text=True,
|
||||||
|
check=True,
|
||||||
|
)
|
||||||
|
data = json.loads(result.stdout)
|
||||||
|
cls._machine_model = data["SPHardwareDataType"][0]["machine_model"]
|
||||||
|
except (subprocess.CalledProcessError, json.JSONDecodeError, IndexError, KeyError) as e:
|
||||||
|
logger.warning(f"Cannot determine machine model via 'system_profiler -json SPHardwareDataType': {e}")
|
||||||
|
cls._machine_model = "Unknown"
|
||||||
|
return cls._machine_model
|
||||||
|
|||||||
@ -5,7 +5,7 @@ import mlx.core as mx
|
|||||||
import PIL.Image
|
import PIL.Image
|
||||||
|
|
||||||
from mflux.callbacks.callback import BeforeLoopCallback
|
from mflux.callbacks.callback import BeforeLoopCallback
|
||||||
from mflux.config.runtime_config import RuntimeConfig
|
from mflux.models.common.config.config import Config
|
||||||
from mflux.utils.image_util import ImageUtil
|
from mflux.utils.image_util import ImageUtil
|
||||||
|
|
||||||
|
|
||||||
@ -18,7 +18,7 @@ class CannyImageSaver(BeforeLoopCallback):
|
|||||||
seed: int,
|
seed: int,
|
||||||
prompt: str,
|
prompt: str,
|
||||||
latents: mx.array,
|
latents: mx.array,
|
||||||
config: RuntimeConfig,
|
config: Config,
|
||||||
canny_image: PIL.Image.Image | None = None,
|
canny_image: PIL.Image.Image | None = None,
|
||||||
depth_image: PIL.Image.Image | None = None,
|
depth_image: PIL.Image.Image | None = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
|
|||||||
@ -5,7 +5,7 @@ import mlx.core as mx
|
|||||||
import PIL.Image
|
import PIL.Image
|
||||||
|
|
||||||
from mflux.callbacks.callback import BeforeLoopCallback
|
from mflux.callbacks.callback import BeforeLoopCallback
|
||||||
from mflux.config.runtime_config import RuntimeConfig
|
from mflux.models.common.config.config import Config
|
||||||
from mflux.utils.image_util import ImageUtil
|
from mflux.utils.image_util import ImageUtil
|
||||||
|
|
||||||
|
|
||||||
@ -18,7 +18,7 @@ class DepthImageSaver(BeforeLoopCallback):
|
|||||||
seed: int,
|
seed: int,
|
||||||
prompt: str,
|
prompt: str,
|
||||||
latents: mx.array,
|
latents: mx.array,
|
||||||
config: RuntimeConfig,
|
config: Config,
|
||||||
canny_image: PIL.Image.Image | None = None,
|
canny_image: PIL.Image.Image | None = None,
|
||||||
depth_image: PIL.Image.Image | None = None,
|
depth_image: PIL.Image.Image | None = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
@ -28,5 +28,5 @@ class DepthImageSaver(BeforeLoopCallback):
|
|||||||
base, ext = os.path.splitext(self.path)
|
base, ext = os.path.splitext(self.path)
|
||||||
ImageUtil.save_image(
|
ImageUtil.save_image(
|
||||||
image=depth_image,
|
image=depth_image,
|
||||||
path=f"{base}_depth_map{ext}"
|
path=f"{base}_depth_map{ext}",
|
||||||
) # fmt: off
|
)
|
||||||
|
|||||||
@ -5,15 +5,10 @@ import PIL.Image
|
|||||||
from tqdm import tqdm
|
from tqdm import tqdm
|
||||||
|
|
||||||
from mflux.callbacks.callback import AfterLoopCallback, BeforeLoopCallback, InLoopCallback
|
from mflux.callbacks.callback import AfterLoopCallback, BeforeLoopCallback, InLoopCallback
|
||||||
from mflux.config.runtime_config import RuntimeConfig
|
from mflux.models.common.config.config import Config
|
||||||
|
|
||||||
|
|
||||||
class MemorySaver(BeforeLoopCallback, InLoopCallback, AfterLoopCallback):
|
class MemorySaver(BeforeLoopCallback, InLoopCallback, AfterLoopCallback):
|
||||||
"""
|
|
||||||
Optimizes memory usage by clearing caches and removing unused model
|
|
||||||
components at strategic points in the execution cycle.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(self, model, keep_transformer: bool = True, cache_limit_bytes: int = 1000**3):
|
def __init__(self, model, keep_transformer: bool = True, cache_limit_bytes: int = 1000**3):
|
||||||
self.model = model
|
self.model = model
|
||||||
self.keep_transformer = keep_transformer
|
self.keep_transformer = keep_transformer
|
||||||
@ -27,7 +22,7 @@ class MemorySaver(BeforeLoopCallback, InLoopCallback, AfterLoopCallback):
|
|||||||
seed: int,
|
seed: int,
|
||||||
prompt: str,
|
prompt: str,
|
||||||
latents: mx.array,
|
latents: mx.array,
|
||||||
config: RuntimeConfig,
|
config: Config,
|
||||||
canny_image: PIL.Image.Image | None = None,
|
canny_image: PIL.Image.Image | None = None,
|
||||||
depth_image: PIL.Image.Image | None = None,
|
depth_image: PIL.Image.Image | None = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
@ -40,7 +35,7 @@ class MemorySaver(BeforeLoopCallback, InLoopCallback, AfterLoopCallback):
|
|||||||
seed: int,
|
seed: int,
|
||||||
prompt: str,
|
prompt: str,
|
||||||
latents: mx.array,
|
latents: mx.array,
|
||||||
config: RuntimeConfig,
|
config: Config,
|
||||||
time_steps: tqdm,
|
time_steps: tqdm,
|
||||||
) -> None:
|
) -> None:
|
||||||
self.peak_memory = mx.get_peak_memory()
|
self.peak_memory = mx.get_peak_memory()
|
||||||
@ -50,7 +45,7 @@ class MemorySaver(BeforeLoopCallback, InLoopCallback, AfterLoopCallback):
|
|||||||
seed: int,
|
seed: int,
|
||||||
prompt: str,
|
prompt: str,
|
||||||
latents: mx.array,
|
latents: mx.array,
|
||||||
config: RuntimeConfig,
|
config: Config,
|
||||||
) -> None:
|
) -> None:
|
||||||
self.peak_memory = mx.get_peak_memory()
|
self.peak_memory = mx.get_peak_memory()
|
||||||
if not self.keep_transformer:
|
if not self.keep_transformer:
|
||||||
@ -66,8 +61,9 @@ class MemorySaver(BeforeLoopCallback, InLoopCallback, AfterLoopCallback):
|
|||||||
self.model.text_encoder = None
|
self.model.text_encoder = None
|
||||||
if hasattr(self.model, "qwen_vl_encoder") and self.model.qwen_vl_encoder is not None:
|
if hasattr(self.model, "qwen_vl_encoder") and self.model.qwen_vl_encoder is not None:
|
||||||
self.model.qwen_vl_encoder = None
|
self.model.qwen_vl_encoder = None
|
||||||
if hasattr(self.model, "qwen_vl_tokenizer") and self.model.qwen_vl_tokenizer is not None:
|
# Clear VLM tokenizers from the tokenizers dict if present
|
||||||
self.model.qwen_vl_tokenizer = None
|
if hasattr(self.model, "tokenizers") and "qwen_vl" in self.model.tokenizers:
|
||||||
|
self.model.tokenizers["qwen_vl"] = None
|
||||||
gc.collect()
|
gc.collect()
|
||||||
mx.clear_cache()
|
mx.clear_cache()
|
||||||
|
|
||||||
|
|||||||
@ -5,8 +5,7 @@ import PIL.Image
|
|||||||
import tqdm
|
import tqdm
|
||||||
|
|
||||||
from mflux.callbacks.callback import BeforeLoopCallback, InLoopCallback, InterruptCallback
|
from mflux.callbacks.callback import BeforeLoopCallback, InLoopCallback, InterruptCallback
|
||||||
from mflux.config.runtime_config import RuntimeConfig
|
from mflux.models.common.config.config import Config
|
||||||
from mflux.utils.array_util import ArrayUtil
|
|
||||||
from mflux.utils.image_util import ImageUtil
|
from mflux.utils.image_util import ImageUtil
|
||||||
|
|
||||||
|
|
||||||
@ -15,9 +14,11 @@ class StepwiseHandler(BeforeLoopCallback, InLoopCallback, InterruptCallback):
|
|||||||
self,
|
self,
|
||||||
model,
|
model,
|
||||||
output_dir: str,
|
output_dir: str,
|
||||||
|
latent_creator,
|
||||||
):
|
):
|
||||||
self.model = model
|
self.model = model
|
||||||
self.output_dir = Path(output_dir)
|
self.output_dir = Path(output_dir)
|
||||||
|
self.latent_creator = latent_creator
|
||||||
self.step_wise_images = []
|
self.step_wise_images = []
|
||||||
|
|
||||||
if self.output_dir:
|
if self.output_dir:
|
||||||
@ -28,7 +29,7 @@ class StepwiseHandler(BeforeLoopCallback, InLoopCallback, InterruptCallback):
|
|||||||
seed: int,
|
seed: int,
|
||||||
prompt: str,
|
prompt: str,
|
||||||
latents: mx.array,
|
latents: mx.array,
|
||||||
config: RuntimeConfig,
|
config: Config,
|
||||||
canny_image: PIL.Image.Image | None = None,
|
canny_image: PIL.Image.Image | None = None,
|
||||||
depth_image: PIL.Image.Image | None = None,
|
depth_image: PIL.Image.Image | None = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
@ -47,7 +48,7 @@ class StepwiseHandler(BeforeLoopCallback, InLoopCallback, InterruptCallback):
|
|||||||
seed: int,
|
seed: int,
|
||||||
prompt: str,
|
prompt: str,
|
||||||
latents: mx.array,
|
latents: mx.array,
|
||||||
config: RuntimeConfig,
|
config: Config,
|
||||||
time_steps: tqdm,
|
time_steps: tqdm,
|
||||||
) -> None:
|
) -> None:
|
||||||
self._save_image(
|
self._save_image(
|
||||||
@ -65,7 +66,7 @@ class StepwiseHandler(BeforeLoopCallback, InLoopCallback, InterruptCallback):
|
|||||||
seed: int,
|
seed: int,
|
||||||
prompt: str,
|
prompt: str,
|
||||||
latents: mx.array,
|
latents: mx.array,
|
||||||
config: RuntimeConfig,
|
config: Config,
|
||||||
time_steps: tqdm,
|
time_steps: tqdm,
|
||||||
) -> None:
|
) -> None:
|
||||||
self._save_composite(seed=seed)
|
self._save_composite(seed=seed)
|
||||||
@ -76,10 +77,10 @@ class StepwiseHandler(BeforeLoopCallback, InLoopCallback, InterruptCallback):
|
|||||||
seed: int,
|
seed: int,
|
||||||
prompt: str,
|
prompt: str,
|
||||||
latents: mx.array,
|
latents: mx.array,
|
||||||
config: RuntimeConfig,
|
config: Config,
|
||||||
time_steps: tqdm,
|
time_steps: tqdm,
|
||||||
) -> None:
|
) -> None:
|
||||||
unpack_latents = ArrayUtil.unpack_latents(latents=latents, height=config.height, width=config.width)
|
unpack_latents = self.latent_creator.unpack_latents(latents=latents, height=config.height, width=config.width)
|
||||||
stepwise_decoded = self.model.vae.decode(unpack_latents)
|
stepwise_decoded = self.model.vae.decode(unpack_latents)
|
||||||
generation_time = time_steps.format_dict["elapsed"] if time_steps is not None else 0
|
generation_time = time_steps.format_dict["elapsed"] if time_steps is not None else 0
|
||||||
stepwise_img = ImageUtil.to_image(
|
stepwise_img = ImageUtil.to_image(
|
||||||
|
|||||||
@ -1,16 +1,12 @@
|
|||||||
"""Generate ZSH completion scripts for mflux commands."""
|
|
||||||
|
|
||||||
import argparse
|
import argparse
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
|
from mflux.cli.defaults import defaults as ui_defaults
|
||||||
|
from mflux.cli.parser.parsers import CommandLineParser
|
||||||
from mflux.models.flux.variants.in_context.utils.in_context_loras import LORA_NAME_MAP
|
from mflux.models.flux.variants.in_context.utils.in_context_loras import LORA_NAME_MAP
|
||||||
from mflux.ui import defaults as ui_defaults
|
|
||||||
from mflux.ui.cli.parsers import CommandLineParser
|
|
||||||
|
|
||||||
|
|
||||||
class CompletionGenerator:
|
class CompletionGenerator:
|
||||||
"""Generate ZSH completion scripts by introspecting argparse parsers."""
|
|
||||||
|
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
self.commands = [
|
self.commands = [
|
||||||
"mflux-generate",
|
"mflux-generate",
|
||||||
@ -22,6 +18,12 @@ class CompletionGenerator:
|
|||||||
"mflux-generate-fill",
|
"mflux-generate-fill",
|
||||||
"mflux-generate-depth",
|
"mflux-generate-depth",
|
||||||
"mflux-generate-redux",
|
"mflux-generate-redux",
|
||||||
|
"mflux-generate-qwen",
|
||||||
|
"mflux-generate-qwen-edit",
|
||||||
|
"mflux-generate-fibo",
|
||||||
|
"mflux-generate-z-image-turbo",
|
||||||
|
"mflux-refine-fibo",
|
||||||
|
"mflux-inspire-fibo",
|
||||||
"mflux-concept",
|
"mflux-concept",
|
||||||
"mflux-concept-from-image",
|
"mflux-concept-from-image",
|
||||||
"mflux-save",
|
"mflux-save",
|
||||||
@ -29,10 +31,10 @@ class CompletionGenerator:
|
|||||||
"mflux-train",
|
"mflux-train",
|
||||||
"mflux-upscale",
|
"mflux-upscale",
|
||||||
"mflux-lora-library",
|
"mflux-lora-library",
|
||||||
|
"mflux-info",
|
||||||
]
|
]
|
||||||
|
|
||||||
def create_parser_for_command(self, command: str) -> CommandLineParser:
|
def create_parser_for_command(self, command: str) -> CommandLineParser:
|
||||||
"""Create the appropriate parser for a given command."""
|
|
||||||
parser = CommandLineParser(prog=command, add_help=False)
|
parser = CommandLineParser(prog=command, add_help=False)
|
||||||
|
|
||||||
if command == "mflux-generate":
|
if command == "mflux-generate":
|
||||||
@ -111,6 +113,58 @@ class CompletionGenerator:
|
|||||||
parser.add_redux_arguments()
|
parser.add_redux_arguments()
|
||||||
parser.add_output_arguments()
|
parser.add_output_arguments()
|
||||||
|
|
||||||
|
elif command == "mflux-generate-qwen":
|
||||||
|
parser.add_general_arguments()
|
||||||
|
parser.add_model_arguments(require_model_arg=False)
|
||||||
|
parser.add_lora_arguments()
|
||||||
|
parser.add_image_generator_arguments(supports_metadata_config=True)
|
||||||
|
parser.add_image_to_image_arguments()
|
||||||
|
parser.add_output_arguments()
|
||||||
|
|
||||||
|
elif command == "mflux-generate-qwen-edit":
|
||||||
|
parser.add_general_arguments()
|
||||||
|
parser.add_model_arguments(require_model_arg=False)
|
||||||
|
parser.add_lora_arguments()
|
||||||
|
parser.add_image_generator_arguments(supports_metadata_config=True)
|
||||||
|
parser.add_argument("--image-paths", type=Path, nargs="+", required=True, help="Local paths to init images")
|
||||||
|
parser.add_output_arguments()
|
||||||
|
|
||||||
|
elif command == "mflux-generate-fibo":
|
||||||
|
parser.add_general_arguments()
|
||||||
|
parser.add_model_arguments(require_model_arg=False)
|
||||||
|
parser.add_lora_arguments()
|
||||||
|
parser.add_image_generator_arguments(supports_metadata_config=True)
|
||||||
|
parser.add_image_to_image_arguments()
|
||||||
|
parser.add_output_arguments()
|
||||||
|
|
||||||
|
elif command == "mflux-generate-z-image-turbo":
|
||||||
|
parser.add_general_arguments()
|
||||||
|
parser.add_model_arguments(require_model_arg=False)
|
||||||
|
parser.add_lora_arguments()
|
||||||
|
parser.add_image_generator_arguments(supports_metadata_config=True)
|
||||||
|
parser.add_image_to_image_arguments()
|
||||||
|
parser.add_output_arguments()
|
||||||
|
|
||||||
|
elif command == "mflux-refine-fibo":
|
||||||
|
parser.add_argument("--prompt-file", type=Path, required=True, help="Path to JSON prompt file to refine")
|
||||||
|
parser.add_argument("--instructions", type=str, required=True, help="Text instructions for refinement")
|
||||||
|
parser.add_argument("--output", type=Path, help="Output path for refined JSON prompt")
|
||||||
|
parser.add_argument("--path", type=str, help="Local path for VLM model")
|
||||||
|
parser.add_argument("--top-p", type=float, help="Top-p sampling for VLM")
|
||||||
|
parser.add_argument("--temperature", type=float, help="Temperature for VLM")
|
||||||
|
parser.add_argument("--max-tokens", type=int, help="Max tokens for VLM generation")
|
||||||
|
parser.add_argument("--seed", type=int, help="Seed for VLM generation")
|
||||||
|
|
||||||
|
elif command == "mflux-inspire-fibo":
|
||||||
|
parser.add_argument("--image-path", type=Path, required=True, help="Path to image file to inspire from")
|
||||||
|
parser.add_argument("--prompt", type=str, help="Optional text prompt to blend with the image")
|
||||||
|
parser.add_argument("--output", type=Path, help="Output path for generated JSON prompt")
|
||||||
|
parser.add_argument("--path", type=str, help="Local path for VLM model")
|
||||||
|
parser.add_argument("--top-p", type=float, help="Top-p sampling for VLM")
|
||||||
|
parser.add_argument("--temperature", type=float, help="Temperature for VLM")
|
||||||
|
parser.add_argument("--max-tokens", type=int, help="Max tokens for VLM generation")
|
||||||
|
parser.add_argument("--seed", type=int, help="Seed for VLM generation")
|
||||||
|
|
||||||
elif command == "mflux-concept":
|
elif command == "mflux-concept":
|
||||||
parser.add_general_arguments()
|
parser.add_general_arguments()
|
||||||
parser.add_model_arguments()
|
parser.add_model_arguments()
|
||||||
@ -149,10 +203,12 @@ class CompletionGenerator:
|
|||||||
list_parser = subparsers.add_parser("list")
|
list_parser = subparsers.add_parser("list")
|
||||||
list_parser.add_argument("--paths", action="store_true")
|
list_parser.add_argument("--paths", action="store_true")
|
||||||
|
|
||||||
|
elif command == "mflux-info":
|
||||||
|
parser.add_info_arguments()
|
||||||
|
|
||||||
return parser
|
return parser
|
||||||
|
|
||||||
def escape_description(self, desc: str) -> str:
|
def escape_description(self, desc: str) -> str:
|
||||||
"""Escape special characters in descriptions for ZSH."""
|
|
||||||
if not desc:
|
if not desc:
|
||||||
return ""
|
return ""
|
||||||
# Escape brackets, quotes, and other special characters
|
# Escape brackets, quotes, and other special characters
|
||||||
@ -165,7 +221,6 @@ class CompletionGenerator:
|
|||||||
return desc
|
return desc
|
||||||
|
|
||||||
def format_argument_spec(self, action: argparse.Action) -> list[str]:
|
def format_argument_spec(self, action: argparse.Action) -> list[str]:
|
||||||
"""Format an argparse action into ZSH completion syntax."""
|
|
||||||
specs = []
|
specs = []
|
||||||
|
|
||||||
# Handle options with both short and long forms
|
# Handle options with both short and long forms
|
||||||
@ -195,7 +250,6 @@ class CompletionGenerator:
|
|||||||
return specs
|
return specs
|
||||||
|
|
||||||
def get_value_spec(self, action: argparse.Action) -> str:
|
def get_value_spec(self, action: argparse.Action) -> str:
|
||||||
"""Get the value specification for an argument."""
|
|
||||||
# Check for special cases first
|
# Check for special cases first
|
||||||
if action.dest == "model":
|
if action.dest == "model":
|
||||||
return ":model:_mflux_models"
|
return ":model:_mflux_models"
|
||||||
@ -241,7 +295,6 @@ class CompletionGenerator:
|
|||||||
return ""
|
return ""
|
||||||
|
|
||||||
def generate_command_function(self, command: str, parser: CommandLineParser) -> str:
|
def generate_command_function(self, command: str, parser: CommandLineParser) -> str:
|
||||||
"""Generate the ZSH completion function for a specific command."""
|
|
||||||
func_name = command.replace("-", "_")
|
func_name = command.replace("-", "_")
|
||||||
lines = [f"_{func_name}() {{"]
|
lines = [f"_{func_name}() {{"]
|
||||||
lines.append(" local -a args")
|
lines.append(" local -a args")
|
||||||
@ -262,7 +315,6 @@ class CompletionGenerator:
|
|||||||
return "\n".join(lines)
|
return "\n".join(lines)
|
||||||
|
|
||||||
def generate_header(self) -> str:
|
def generate_header(self) -> str:
|
||||||
"""Generate the completion script header."""
|
|
||||||
commands = " ".join(self.commands)
|
commands = " ".join(self.commands)
|
||||||
return f"""#compdef {commands}
|
return f"""#compdef {commands}
|
||||||
# ZSH completion for mflux commands
|
# ZSH completion for mflux commands
|
||||||
@ -271,7 +323,6 @@ class CompletionGenerator:
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
def generate_helper_functions(self) -> str:
|
def generate_helper_functions(self) -> str:
|
||||||
"""Generate helper functions for common completions."""
|
|
||||||
helpers = []
|
helpers = []
|
||||||
|
|
||||||
# Model completion helper
|
# Model completion helper
|
||||||
@ -303,7 +354,6 @@ class CompletionGenerator:
|
|||||||
return "\n".join(helpers)
|
return "\n".join(helpers)
|
||||||
|
|
||||||
def generate_main_function(self) -> str:
|
def generate_main_function(self) -> str:
|
||||||
"""Generate the main completion dispatcher."""
|
|
||||||
lines = ["# Main completion dispatcher", "_mflux() {", " local cmd=$words[1]", " case $cmd in"]
|
lines = ["# Main completion dispatcher", "_mflux() {", " local cmd=$words[1]", " case $cmd in"]
|
||||||
|
|
||||||
for command in self.commands:
|
for command in self.commands:
|
||||||
@ -317,7 +367,6 @@ class CompletionGenerator:
|
|||||||
return "\n".join(lines)
|
return "\n".join(lines)
|
||||||
|
|
||||||
def generate(self) -> str:
|
def generate(self) -> str:
|
||||||
"""Generate the complete ZSH completion script."""
|
|
||||||
script = [self.generate_header()]
|
script = [self.generate_header()]
|
||||||
script.append(self.generate_helper_functions())
|
script.append(self.generate_helper_functions())
|
||||||
|
|
||||||
@ -1,17 +1,14 @@
|
|||||||
#!/usr/bin/env python3
|
#!/usr/bin/env python3
|
||||||
"""Install ZSH completions for mflux commands."""
|
|
||||||
|
|
||||||
import argparse
|
import argparse
|
||||||
import os
|
import os
|
||||||
import subprocess
|
import subprocess
|
||||||
import sys
|
import sys
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
from mflux.ui.cli.completions.generator import CompletionGenerator
|
from mflux.cli.completions.generator import CompletionGenerator
|
||||||
|
|
||||||
|
|
||||||
def get_zsh_fpath():
|
def get_zsh_fpath():
|
||||||
"""Get the ZSH fpath directories."""
|
|
||||||
try:
|
try:
|
||||||
result = subprocess.run(
|
result = subprocess.run(
|
||||||
["zsh", "-c", "echo $fpath"],
|
["zsh", "-c", "echo $fpath"],
|
||||||
@ -25,7 +22,6 @@ def get_zsh_fpath():
|
|||||||
|
|
||||||
|
|
||||||
def find_completion_dir():
|
def find_completion_dir():
|
||||||
"""Find appropriate directory for completion files."""
|
|
||||||
# Common completion directories in order of preference
|
# Common completion directories in order of preference
|
||||||
candidates = [
|
candidates = [
|
||||||
Path.home() / ".zsh" / "completions",
|
Path.home() / ".zsh" / "completions",
|
||||||
@ -53,7 +49,6 @@ def find_completion_dir():
|
|||||||
|
|
||||||
|
|
||||||
def check_installation():
|
def check_installation():
|
||||||
"""Check if completions are properly installed and accessible."""
|
|
||||||
print("Checking mflux completions installation...\n")
|
print("Checking mflux completions installation...\n")
|
||||||
|
|
||||||
# Check if completion file exists
|
# Check if completion file exists
|
||||||
@ -107,7 +102,6 @@ def check_installation():
|
|||||||
|
|
||||||
|
|
||||||
def main():
|
def main():
|
||||||
"""Main entry point for completion installation."""
|
|
||||||
parser = argparse.ArgumentParser(
|
parser = argparse.ArgumentParser(
|
||||||
description="Install or generate ZSH completions for mflux commands",
|
description="Install or generate ZSH completions for mflux commands",
|
||||||
formatter_class=argparse.RawDescriptionHelpFormatter,
|
formatter_class=argparse.RawDescriptionHelpFormatter,
|
||||||
31
src/mflux/cli/defaults/defaults.py
Normal file
31
src/mflux/cli/defaults/defaults.py
Normal file
@ -0,0 +1,31 @@
|
|||||||
|
import os
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import platformdirs
|
||||||
|
|
||||||
|
BATTERY_PERCENTAGE_STOP_LIMIT = 5
|
||||||
|
CONTROLNET_STRENGTH = 0.4
|
||||||
|
DEFAULT_DEV_FILL_GUIDANCE = 30
|
||||||
|
DEFAULT_DEPTH_GUIDANCE = 10
|
||||||
|
DIMENSION_STEP_PIXELS = 16
|
||||||
|
GUIDANCE_SCALE = 3.5
|
||||||
|
GUIDANCE_SCALE_KONTEXT = 2.5
|
||||||
|
HEIGHT, WIDTH = 1024, 1024
|
||||||
|
IMAGE_STRENGTH = 0.4
|
||||||
|
MODEL_CHOICES = ["dev", "schnell", "krea-dev", "dev-krea", "qwen", "fibo", "z-image-turbo"]
|
||||||
|
MODEL_INFERENCE_STEPS = {
|
||||||
|
"dev": 25,
|
||||||
|
"schnell": 4,
|
||||||
|
"krea-dev": 25,
|
||||||
|
"qwen": 20,
|
||||||
|
"fibo": 20,
|
||||||
|
"z-image-turbo": 9,
|
||||||
|
}
|
||||||
|
QUANTIZE_CHOICES = [3, 5, 4, 6, 8]
|
||||||
|
|
||||||
|
if os.environ.get("MFLUX_CACHE_DIR"):
|
||||||
|
MFLUX_CACHE_DIR = Path(os.environ["MFLUX_CACHE_DIR"]).resolve()
|
||||||
|
else:
|
||||||
|
MFLUX_CACHE_DIR = Path(platformdirs.user_cache_dir(appname="mflux"))
|
||||||
|
|
||||||
|
MFLUX_LORA_CACHE_DIR = MFLUX_CACHE_DIR / "loras"
|
||||||
@ -5,28 +5,19 @@ import time
|
|||||||
import typing as t
|
import typing as t
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
from mflux.models.common.lora.download.lora_library import get_lora_path
|
from mflux.cli.defaults import defaults as ui_defaults
|
||||||
from mflux.models.flux.variants.in_context.utils.in_context_loras import LORA_NAME_MAP, LORA_REPO_ID
|
from mflux.models.common.resolution.lora_resolution import LoraResolution
|
||||||
from mflux.ui import (
|
from mflux.models.flux.variants.in_context.utils.in_context_loras import LORA_NAME_MAP
|
||||||
box_values,
|
from mflux.utils import box_values, scale_factor
|
||||||
defaults as ui_defaults,
|
|
||||||
scale_factor,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
class ModelSpecAction(argparse.Action):
|
class ModelSpecAction(argparse.Action):
|
||||||
def __call__(self, parser, namespace, values, option_string=None):
|
def __call__(self, parser, namespace, values, option_string=None):
|
||||||
if values in ui_defaults.MODEL_CHOICES:
|
# Accept:
|
||||||
setattr(namespace, self.dest, values)
|
# 1. Predefined model names (dev, schnell, fibo, etc.)
|
||||||
return
|
# 2. HuggingFace repos (org/model format)
|
||||||
|
# 3. Local paths (/path/to/model, ./model, ~/model)
|
||||||
if values.count("/") != 1:
|
# The WeightLoader will determine if it's a local path or HuggingFace repo
|
||||||
raise argparse.ArgumentError(
|
|
||||||
self,
|
|
||||||
(f'Value must be either {" ".join(ui_defaults.MODEL_CHOICES)} or in format "org/model". Got: {values}'),
|
|
||||||
)
|
|
||||||
|
|
||||||
# If we got here, values contains exactly one slash
|
|
||||||
setattr(namespace, self.dest, values)
|
setattr(namespace, self.dest, values)
|
||||||
|
|
||||||
|
|
||||||
@ -42,7 +33,7 @@ def int_or_special_value(value) -> int | scale_factor.ScaleFactor:
|
|||||||
|
|
||||||
# If not an integer, try to parse as scale factor
|
# If not an integer, try to parse as scale factor
|
||||||
try:
|
try:
|
||||||
return scale_factor.parse_scale_factor(value)
|
return scale_factor.ScaleFactor.parse(value)
|
||||||
except ValueError:
|
except ValueError:
|
||||||
raise argparse.ArgumentTypeError(
|
raise argparse.ArgumentTypeError(
|
||||||
f"'{value}' is not a valid integer or 'auto' or a scale factor like '2x' or '3.5x'"
|
f"'{value}' is not a valid integer or 'auto' or a scale factor like '2x' or '3.5x'"
|
||||||
@ -71,10 +62,8 @@ 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.require_model_arg = require_model_arg
|
self.require_model_arg = require_model_arg
|
||||||
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).")
|
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)}, a HuggingFace repo org/model, or a local path).")
|
||||||
if path_type == "load":
|
if path_type == "save":
|
||||||
self.add_argument("--path", type=str, default=None, help="Local path for loading a model from disk")
|
|
||||||
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("--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)")
|
||||||
@ -83,10 +72,8 @@ class CommandLineParser(argparse.ArgumentParser):
|
|||||||
self.supports_lora = True
|
self.supports_lora = True
|
||||||
lora_group = self.add_argument_group("LoRA configuration")
|
lora_group = self.add_argument_group("LoRA configuration")
|
||||||
lora_group.add_argument("--lora-style", type=str, choices=sorted(LORA_NAME_MAP.keys()), help="Style of the LoRA to use (e.g., 'storyboard' for film storyboard style)")
|
lora_group.add_argument("--lora-style", type=str, choices=sorted(LORA_NAME_MAP.keys()), help="Style of the LoRA to use (e.g., 'storyboard' for film storyboard style)")
|
||||||
self.add_argument("--lora-paths", type=str, nargs="*", default=None, help="Local safetensors for applying LORA from disk")
|
self.add_argument("--lora-paths", type=str, nargs="*", default=None, help="LoRA paths: local files, HuggingFace repos (org/model), or collection format (repo:filename.safetensors)")
|
||||||
self.add_argument("--lora-scales", type=float, nargs="*", default=None, help="Scaling factor to adjust the impact of LoRA weights on the model. A value of 1.0 applies the LoRA weights as they are.")
|
self.add_argument("--lora-scales", type=float, nargs="*", default=None, help="Scaling factor to adjust the impact of LoRA weights on the model. A value of 1.0 applies the LoRA weights as they are.")
|
||||||
lora_group.add_argument("--lora-name", type=str, help="Name of the LoRA to download from Hugging Face")
|
|
||||||
lora_group.add_argument("--lora-repo-id", type=str, default=LORA_REPO_ID, help=f"Hugging Face repository ID for LoRAs (default: {LORA_REPO_ID})")
|
|
||||||
|
|
||||||
def _add_image_generator_common_arguments(self, supports_dimension_scale_factor=False) -> None:
|
def _add_image_generator_common_arguments(self, supports_dimension_scale_factor=False) -> None:
|
||||||
self.supports_image_generation = True
|
self.supports_image_generation = True
|
||||||
@ -312,7 +299,7 @@ class CommandLineParser(argparse.ArgumentParser):
|
|||||||
|
|
||||||
if self.supports_image_outpaint and namespace.image_outpaint_padding is not None:
|
if self.supports_image_outpaint and namespace.image_outpaint_padding is not None:
|
||||||
# parse and normalize any acceptable 1,2,3,4-tuple box value to 4-tuple
|
# parse and normalize any acceptable 1,2,3,4-tuple box value to 4-tuple
|
||||||
namespace.image_outpaint_padding = box_values.parse_box_value(namespace.image_outpaint_padding)
|
namespace.image_outpaint_padding = box_values.BoxValues.parse(namespace.image_outpaint_padding)
|
||||||
print(f"{namespace.image_outpaint_padding=}")
|
print(f"{namespace.image_outpaint_padding=}")
|
||||||
|
|
||||||
# Resolve lora paths from library if needed
|
# Resolve lora paths from library if needed
|
||||||
@ -320,10 +307,17 @@ class CommandLineParser(argparse.ArgumentParser):
|
|||||||
resolved_paths = []
|
resolved_paths = []
|
||||||
for lora_path in namespace.lora_paths:
|
for lora_path in namespace.lora_paths:
|
||||||
try:
|
try:
|
||||||
resolved_path = get_lora_path(lora_path)
|
resolved_path = LoraResolution.resolve(lora_path)
|
||||||
resolved_paths.append(resolved_path)
|
resolved_paths.append(resolved_path)
|
||||||
except FileNotFoundError as e: # noqa: PERF203
|
except FileNotFoundError as e: # noqa: PERF203
|
||||||
self.error(str(e))
|
self.error(str(e))
|
||||||
namespace.lora_paths = resolved_paths
|
namespace.lora_paths = resolved_paths
|
||||||
|
|
||||||
|
# Compute model_path: None for predefined names, otherwise use the model value
|
||||||
|
# Predefined names like "schnell", "dev" are handled by ModelConfig, not PathResolution
|
||||||
|
if hasattr(namespace, "model") and namespace.model is not None:
|
||||||
|
namespace.model_path = None if namespace.model in ui_defaults.MODEL_CHOICES else namespace.model
|
||||||
|
else:
|
||||||
|
namespace.model_path = None
|
||||||
|
|
||||||
return namespace
|
return namespace
|
||||||
@ -1,40 +0,0 @@
|
|||||||
import logging
|
|
||||||
from pathlib import Path
|
|
||||||
|
|
||||||
import mlx.core as mx
|
|
||||||
|
|
||||||
log = logging.getLogger(__name__)
|
|
||||||
|
|
||||||
|
|
||||||
class Config:
|
|
||||||
precision: mx.Dtype = mx.bfloat16
|
|
||||||
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
num_inference_steps: int = 4,
|
|
||||||
width: int = 1024,
|
|
||||||
height: int = 1024,
|
|
||||||
guidance: float = 4.0,
|
|
||||||
image_path: Path | None = None,
|
|
||||||
image_strength: float | None = None,
|
|
||||||
depth_image_path: Path | None = None,
|
|
||||||
redux_image_paths: list[Path] | None = None,
|
|
||||||
redux_image_strengths: list[float] | None = None,
|
|
||||||
masked_image_path: Path | None = None,
|
|
||||||
controlnet_strength: float | None = None,
|
|
||||||
scheduler: str = "linear",
|
|
||||||
):
|
|
||||||
if width % 16 != 0 or height % 16 != 0:
|
|
||||||
log.warning("Width and height should be multiples of 16. Rounding down.")
|
|
||||||
self.width = 16 * (width // 16)
|
|
||||||
self.height = 16 * (height // 16)
|
|
||||||
self.num_inference_steps = num_inference_steps
|
|
||||||
self.guidance = guidance
|
|
||||||
self.image_path = image_path
|
|
||||||
self.image_strength = image_strength
|
|
||||||
self.depth_image_path = depth_image_path
|
|
||||||
self.redux_image_paths = redux_image_paths
|
|
||||||
self.redux_image_strengths = redux_image_strengths
|
|
||||||
self.masked_image_path = masked_image_path
|
|
||||||
self.controlnet_strength = controlnet_strength
|
|
||||||
self.scheduler_str = scheduler
|
|
||||||
@ -1,116 +0,0 @@
|
|||||||
import logging
|
|
||||||
from pathlib import Path
|
|
||||||
|
|
||||||
import mlx.core as mx
|
|
||||||
|
|
||||||
from mflux.config.config import Config
|
|
||||||
from mflux.config.model_config import ModelConfig
|
|
||||||
from mflux.models.common.schedulers import SCHEDULER_REGISTRY, try_import_external_scheduler
|
|
||||||
from mflux.models.common.schedulers.linear_scheduler import LinearScheduler
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
|
||||||
|
|
||||||
|
|
||||||
class RuntimeConfig:
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
config: Config,
|
|
||||||
model_config: ModelConfig,
|
|
||||||
):
|
|
||||||
self.config = config
|
|
||||||
self.model_config = model_config
|
|
||||||
self._scheduler = None
|
|
||||||
|
|
||||||
@property
|
|
||||||
def height(self) -> int:
|
|
||||||
return self.config.height
|
|
||||||
|
|
||||||
@property
|
|
||||||
def width(self) -> int:
|
|
||||||
return self.config.width
|
|
||||||
|
|
||||||
@width.setter
|
|
||||||
def width(self, value):
|
|
||||||
self.config.width = value
|
|
||||||
|
|
||||||
@property
|
|
||||||
def guidance(self) -> float:
|
|
||||||
return self.config.guidance
|
|
||||||
|
|
||||||
@property
|
|
||||||
def num_inference_steps(self) -> int:
|
|
||||||
return self.config.num_inference_steps
|
|
||||||
|
|
||||||
@property
|
|
||||||
def precision(self) -> mx.Dtype:
|
|
||||||
return self.config.precision
|
|
||||||
|
|
||||||
@property
|
|
||||||
def num_train_steps(self) -> int:
|
|
||||||
return self.model_config.num_train_steps
|
|
||||||
|
|
||||||
@property
|
|
||||||
def image_path(self) -> Path | None:
|
|
||||||
return self.config.image_path
|
|
||||||
|
|
||||||
@property
|
|
||||||
def image_strength(self) -> float | None:
|
|
||||||
return self.config.image_strength
|
|
||||||
|
|
||||||
@property
|
|
||||||
def depth_image_path(self) -> Path | None:
|
|
||||||
return self.config.depth_image_path
|
|
||||||
|
|
||||||
@property
|
|
||||||
def redux_image_paths(self) -> list[Path] | None:
|
|
||||||
return self.config.redux_image_paths
|
|
||||||
|
|
||||||
@property
|
|
||||||
def redux_image_strengths(self) -> list[float] | None:
|
|
||||||
return self.config.redux_image_strengths
|
|
||||||
|
|
||||||
@property
|
|
||||||
def masked_image_path(self) -> Path | None:
|
|
||||||
return self.config.masked_image_path
|
|
||||||
|
|
||||||
@property
|
|
||||||
def init_time_step(self) -> int:
|
|
||||||
is_img2img = (
|
|
||||||
self.config.image_path is not None and
|
|
||||||
self.image_strength is not None and
|
|
||||||
self.image_strength > 0.0
|
|
||||||
) # fmt: off
|
|
||||||
|
|
||||||
if is_img2img:
|
|
||||||
# 1. Clamp strength to [0, 1]
|
|
||||||
strength = max(0.0, min(1.0, self.config.image_strength)) # type: ignore
|
|
||||||
|
|
||||||
# 2. Return start time in [1, floor(num_steps * strength)]
|
|
||||||
return max(1, int(self.num_inference_steps * strength)) # type: ignore
|
|
||||||
else:
|
|
||||||
return 0
|
|
||||||
|
|
||||||
@property
|
|
||||||
def controlnet_strength(self) -> float | None:
|
|
||||||
if self.config.controlnet_strength is not None:
|
|
||||||
return self.config.controlnet_strength
|
|
||||||
|
|
||||||
return None
|
|
||||||
|
|
||||||
@property
|
|
||||||
def scheduler(self):
|
|
||||||
if self._scheduler is not None:
|
|
||||||
return self._scheduler
|
|
||||||
|
|
||||||
if self.config.scheduler_str == "linear":
|
|
||||||
self._scheduler = LinearScheduler(self)
|
|
||||||
elif (registered_scheduler := SCHEDULER_REGISTRY.get(self.config.scheduler_str, None)) is not None:
|
|
||||||
self._scheduler = registered_scheduler(self)
|
|
||||||
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.")
|
|
||||||
|
|
||||||
return self._scheduler
|
|
||||||
@ -1,123 +0,0 @@
|
|||||||
"""Display metadata information from MFLUX generated images."""
|
|
||||||
|
|
||||||
import sys
|
|
||||||
from datetime import datetime
|
|
||||||
from pathlib import Path
|
|
||||||
|
|
||||||
from mflux.ui.cli.parsers import CommandLineParser
|
|
||||||
from mflux.utils.metadata_reader import MetadataReader
|
|
||||||
|
|
||||||
|
|
||||||
def format_metadata(metadata: dict) -> str:
|
|
||||||
"""Format metadata in a clean, readable format."""
|
|
||||||
exif = metadata.get("exif", {})
|
|
||||||
if not exif:
|
|
||||||
return "No metadata found"
|
|
||||||
|
|
||||||
lines = []
|
|
||||||
lines.append("=" * 60)
|
|
||||||
lines.append("MFLUX Image Information")
|
|
||||||
lines.append("=" * 60)
|
|
||||||
|
|
||||||
# Prompt
|
|
||||||
if prompt := exif.get("prompt"):
|
|
||||||
lines.append(f"\nPrompt: {prompt}")
|
|
||||||
|
|
||||||
if negative_prompt := exif.get("negative_prompt"):
|
|
||||||
lines.append(f"Negative Prompt: {negative_prompt}")
|
|
||||||
|
|
||||||
# Model information
|
|
||||||
lines.append("")
|
|
||||||
if model := exif.get("model"):
|
|
||||||
lines.append(f"Model: {model}")
|
|
||||||
|
|
||||||
# Image dimensions
|
|
||||||
if width := exif.get("width"):
|
|
||||||
lines.append(f"Width: {width}")
|
|
||||||
if height := exif.get("height"):
|
|
||||||
lines.append(f"Height: {height}")
|
|
||||||
|
|
||||||
# Generation parameters
|
|
||||||
lines.append("")
|
|
||||||
if seed := exif.get("seed"):
|
|
||||||
lines.append(f"Seed: {seed}")
|
|
||||||
if steps := exif.get("steps"):
|
|
||||||
lines.append(f"Steps: {steps}")
|
|
||||||
if guidance := exif.get("guidance"):
|
|
||||||
lines.append(f"Guidance: {guidance}")
|
|
||||||
|
|
||||||
# Technical settings
|
|
||||||
if quantize := exif.get("quantize"):
|
|
||||||
lines.append(f"Quantization: {quantize}-bit")
|
|
||||||
if precision := exif.get("precision"):
|
|
||||||
lines.append(f"Precision: {precision}")
|
|
||||||
|
|
||||||
# LoRA information
|
|
||||||
if lora_paths := exif.get("lora_paths"):
|
|
||||||
lines.append("")
|
|
||||||
lines.append(f"LoRAs ({len(lora_paths)}):")
|
|
||||||
lora_scales = exif.get("lora_scales") or []
|
|
||||||
for i, lora in enumerate(lora_paths):
|
|
||||||
scale = lora_scales[i] if i < len(lora_scales) else 1.0
|
|
||||||
lora_name = Path(lora).name
|
|
||||||
lines.append(f" - {lora_name} (scale: {scale})")
|
|
||||||
|
|
||||||
# Image-to-image parameters
|
|
||||||
if image_path := exif.get("image_path"):
|
|
||||||
lines.append("")
|
|
||||||
lines.append(f"Source Image: {Path(image_path).name}")
|
|
||||||
if image_strength := exif.get("image_strength"):
|
|
||||||
lines.append(f"Image Strength: {image_strength}")
|
|
||||||
|
|
||||||
# ControlNet parameters
|
|
||||||
if controlnet_path := exif.get("controlnet_image_path"):
|
|
||||||
lines.append("")
|
|
||||||
lines.append(f"ControlNet Image: {Path(controlnet_path).name}")
|
|
||||||
if controlnet_strength := exif.get("controlnet_strength"):
|
|
||||||
lines.append(f"ControlNet Strength: {controlnet_strength}")
|
|
||||||
|
|
||||||
# Generation metadata
|
|
||||||
lines.append("")
|
|
||||||
if gen_time := exif.get("generation_time_seconds"):
|
|
||||||
lines.append(f"Generation Time: {gen_time:.2f}s")
|
|
||||||
|
|
||||||
if created_at := exif.get("created_at"):
|
|
||||||
try:
|
|
||||||
dt = datetime.fromisoformat(created_at)
|
|
||||||
lines.append(f"Created: {dt.strftime('%Y-%m-%d %H:%M:%S')}")
|
|
||||||
except (ValueError, AttributeError):
|
|
||||||
lines.append(f"Created: {created_at}")
|
|
||||||
|
|
||||||
if version := exif.get("mflux_version"):
|
|
||||||
lines.append(f"MFLUX Version: {version}")
|
|
||||||
|
|
||||||
lines.append("=" * 60)
|
|
||||||
return "\n".join(lines)
|
|
||||||
|
|
||||||
|
|
||||||
def main():
|
|
||||||
# Parse command line arguments
|
|
||||||
parser = CommandLineParser(description="Display metadata from MFLUX generated images")
|
|
||||||
parser.add_info_arguments()
|
|
||||||
args = parser.parse_args()
|
|
||||||
|
|
||||||
# Check if file exists
|
|
||||||
image_path = Path(args.image_path)
|
|
||||||
if not image_path.exists():
|
|
||||||
print(f"Error: Image file not found: {image_path}")
|
|
||||||
sys.exit(1)
|
|
||||||
|
|
||||||
# Read metadata
|
|
||||||
metadata = MetadataReader.read_all_metadata(image_path)
|
|
||||||
|
|
||||||
# Check if metadata was found
|
|
||||||
if not metadata or (not metadata.get("exif") and not metadata.get("xmp")):
|
|
||||||
print("No metadata found")
|
|
||||||
sys.exit(1)
|
|
||||||
|
|
||||||
# Format and display
|
|
||||||
print(format_metadata(metadata))
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
|
||||||
main()
|
|
||||||
@ -1,133 +0,0 @@
|
|||||||
#!/usr/bin/env python3
|
|
||||||
"""CLI tool for managing the MFLUX LoRA library."""
|
|
||||||
|
|
||||||
import argparse
|
|
||||||
import os
|
|
||||||
import sys
|
|
||||||
from collections import defaultdict
|
|
||||||
from pathlib import Path
|
|
||||||
|
|
||||||
from mflux.models.common.lora.download.lora_library import _discover_lora_files
|
|
||||||
|
|
||||||
|
|
||||||
def list_loras(paths: list[str] | None = None) -> int:
|
|
||||||
"""List all discovered LoRA files from specified paths or LORA_LIBRARY_PATH.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
paths: Optional list of paths to use instead of LORA_LIBRARY_PATH
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Exit code (0 for success, 1 for error)
|
|
||||||
"""
|
|
||||||
if paths:
|
|
||||||
# Use provided paths
|
|
||||||
library_paths = [Path(p.strip()) for p in paths]
|
|
||||||
else:
|
|
||||||
# Use environment variable
|
|
||||||
library_path_env = os.environ.get("LORA_LIBRARY_PATH")
|
|
||||||
|
|
||||||
if not library_path_env:
|
|
||||||
print("LORA_LIBRARY_PATH environment variable is not set.", file=sys.stderr)
|
|
||||||
print("Set it to one or more colon-separated directories containing .safetensors files.", file=sys.stderr)
|
|
||||||
print("Alternatively, use --paths to specify directories directly.", file=sys.stderr)
|
|
||||||
return 1
|
|
||||||
|
|
||||||
# Parse library paths from environment
|
|
||||||
library_paths = [Path(p.strip()) for p in library_path_env.split(":") if p.strip()]
|
|
||||||
|
|
||||||
# Check which paths exist
|
|
||||||
valid_paths = []
|
|
||||||
for path in library_paths:
|
|
||||||
if path.exists() and path.is_dir():
|
|
||||||
valid_paths.append(path)
|
|
||||||
else:
|
|
||||||
print(f"Warning: Path does not exist or is not a directory: {path}", file=sys.stderr)
|
|
||||||
|
|
||||||
if not valid_paths:
|
|
||||||
print("No valid directories found in LORA_LIBRARY_PATH.", file=sys.stderr)
|
|
||||||
return 1
|
|
||||||
|
|
||||||
# Discover all LoRA files
|
|
||||||
lora_registry = _discover_lora_files(valid_paths)
|
|
||||||
|
|
||||||
if not lora_registry:
|
|
||||||
print("No .safetensors files found in the specified directories.")
|
|
||||||
return 0
|
|
||||||
|
|
||||||
# Sort by basename for consistent output
|
|
||||||
sorted_items = sorted(lora_registry.items())
|
|
||||||
|
|
||||||
# Print all discovered LoRAs
|
|
||||||
print("Discovered LoRA files:")
|
|
||||||
print("-" * 80)
|
|
||||||
for basename, full_path in sorted_items:
|
|
||||||
print(f"{basename} -> {full_path}")
|
|
||||||
|
|
||||||
# Calculate statistics per top-level directory
|
|
||||||
stats = defaultdict(int)
|
|
||||||
for full_path in lora_registry.values():
|
|
||||||
# Find which library path this file belongs to
|
|
||||||
for lib_path in valid_paths:
|
|
||||||
try:
|
|
||||||
# Check if full_path is relative to lib_path
|
|
||||||
full_path.relative_to(lib_path.resolve())
|
|
||||||
stats[str(lib_path)] += 1
|
|
||||||
break
|
|
||||||
except ValueError:
|
|
||||||
# Not relative to this lib_path, continue
|
|
||||||
continue
|
|
||||||
|
|
||||||
# Print summary
|
|
||||||
print("-" * 80)
|
|
||||||
print(f"\nTotal LoRA files found: {len(lora_registry)}")
|
|
||||||
|
|
||||||
if len(valid_paths) > 1:
|
|
||||||
print("\nBreakdown by library path:")
|
|
||||||
for lib_path in valid_paths:
|
|
||||||
count = stats.get(str(lib_path), 0)
|
|
||||||
print(f" {lib_path}: {count} files")
|
|
||||||
|
|
||||||
return 0
|
|
||||||
|
|
||||||
|
|
||||||
def main():
|
|
||||||
"""Main entry point for the CLI."""
|
|
||||||
parser = argparse.ArgumentParser(
|
|
||||||
description="MFLUX LoRA Library management tool",
|
|
||||||
formatter_class=argparse.RawDescriptionHelpFormatter,
|
|
||||||
epilog="""
|
|
||||||
Environment Variables:
|
|
||||||
LORA_LIBRARY_PATH Colon-separated list of directories containing .safetensors files
|
|
||||||
Example: /path/to/loras:/another/path/to/loras
|
|
||||||
|
|
||||||
Examples:
|
|
||||||
# List all discovered LoRA files using LORA_LIBRARY_PATH
|
|
||||||
mflux-lora-library list
|
|
||||||
|
|
||||||
# With environment variable set
|
|
||||||
LORA_LIBRARY_PATH=/home/user/loras:/opt/shared/loras mflux-lora-library list
|
|
||||||
|
|
||||||
# Override with specific paths
|
|
||||||
mflux-lora-library list --paths /path/to/loras /another/path/to/loras
|
|
||||||
""",
|
|
||||||
)
|
|
||||||
|
|
||||||
subparsers = parser.add_subparsers(dest="command", help="Available commands")
|
|
||||||
|
|
||||||
# Add 'list' command
|
|
||||||
list_parser = subparsers.add_parser("list", help="List all discovered LoRA files")
|
|
||||||
list_parser.add_argument(
|
|
||||||
"--paths", nargs="+", help="Override LORA_LIBRARY_PATH with these directories (space-separated)"
|
|
||||||
)
|
|
||||||
|
|
||||||
args = parser.parse_args()
|
|
||||||
|
|
||||||
if args.command == "list":
|
|
||||||
return list_loras(paths=args.paths)
|
|
||||||
else:
|
|
||||||
parser.print_help()
|
|
||||||
return 1
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
|
||||||
sys.exit(main())
|
|
||||||
34
src/mflux/models/common/cli/info.py
Normal file
34
src/mflux/models/common/cli/info.py
Normal file
@ -0,0 +1,34 @@
|
|||||||
|
import sys
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
from mflux.cli.parser.parsers import CommandLineParser
|
||||||
|
from mflux.utils.info_util import InfoUtil
|
||||||
|
from mflux.utils.metadata_reader import MetadataReader
|
||||||
|
|
||||||
|
|
||||||
|
def main():
|
||||||
|
# Parse command line arguments
|
||||||
|
parser = CommandLineParser(description="Display metadata from MFLUX generated images")
|
||||||
|
parser.add_info_arguments()
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
# Check if file exists
|
||||||
|
image_path = Path(args.image_path)
|
||||||
|
if not image_path.exists():
|
||||||
|
print(f"Error: Image file not found: {image_path}")
|
||||||
|
sys.exit(1)
|
||||||
|
|
||||||
|
# Read metadata
|
||||||
|
metadata = MetadataReader.read_all_metadata(image_path)
|
||||||
|
|
||||||
|
# Check if metadata was found
|
||||||
|
if not metadata or (not metadata.get("exif") and not metadata.get("xmp")):
|
||||||
|
print("No metadata found")
|
||||||
|
sys.exit(1)
|
||||||
|
|
||||||
|
# Format and display
|
||||||
|
print(InfoUtil.format_metadata(metadata))
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
26
src/mflux/models/common/cli/lora_library.py
Normal file
26
src/mflux/models/common/cli/lora_library.py
Normal file
@ -0,0 +1,26 @@
|
|||||||
|
import argparse
|
||||||
|
import sys
|
||||||
|
|
||||||
|
from mflux.utils.lora_library_util import LoraLibraryUtil
|
||||||
|
|
||||||
|
|
||||||
|
def main():
|
||||||
|
parser = argparse.ArgumentParser(
|
||||||
|
description="MFLUX LoRA Library management tool",
|
||||||
|
formatter_class=argparse.RawDescriptionHelpFormatter,
|
||||||
|
epilog=LoraLibraryUtil.epilog(),
|
||||||
|
)
|
||||||
|
subparsers = parser.add_subparsers(dest="command", help="Available commands")
|
||||||
|
list_parser = subparsers.add_parser("list", help="List all discovered LoRA files")
|
||||||
|
list_parser.add_argument("--paths", nargs="+", help="Override LORA_LIBRARY_PATH with these directories (space-separated)") # fmt: off
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
if args.command == "list":
|
||||||
|
return LoraLibraryUtil.list_loras(paths=args.paths)
|
||||||
|
else:
|
||||||
|
parser.print_help()
|
||||||
|
return 1
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
sys.exit(main())
|
||||||
@ -1,8 +1,9 @@
|
|||||||
from mflux.config.model_config import ModelConfig
|
from mflux.cli.parser.parsers import CommandLineParser
|
||||||
|
from mflux.models.common.config import ModelConfig
|
||||||
from mflux.models.fibo.variants.txt2img.fibo import FIBO
|
from mflux.models.fibo.variants.txt2img.fibo import FIBO
|
||||||
from mflux.models.flux.variants.txt2img.flux import Flux1
|
from mflux.models.flux.variants.txt2img.flux import Flux1
|
||||||
from mflux.models.qwen.variants.txt2img.qwen_image import QwenImage
|
from mflux.models.qwen.variants.txt2img.qwen_image import QwenImage
|
||||||
from mflux.ui.cli.parsers import CommandLineParser
|
from mflux.models.z_image.variants.turbo.z_image_turbo import ZImageTurbo
|
||||||
|
|
||||||
|
|
||||||
def main():
|
def main():
|
||||||
@ -18,19 +19,18 @@ def main():
|
|||||||
model_class = QwenImage
|
model_class = QwenImage
|
||||||
elif "fibo" in model_name_lower:
|
elif "fibo" in model_name_lower:
|
||||||
model_class = FIBO
|
model_class = FIBO
|
||||||
|
elif "z-image" in model_name_lower or "zimage" in model_name_lower:
|
||||||
|
model_class = ZImageTurbo
|
||||||
else:
|
else:
|
||||||
model_class = Flux1
|
model_class = Flux1
|
||||||
|
|
||||||
# 2. Load, quantize and save the model
|
# 2. Load, quantize and save the model
|
||||||
model_kwargs = {
|
model = model_class(
|
||||||
"model_config": ModelConfig.from_name(args.model, base_model=args.base_model),
|
quantize=args.quantize,
|
||||||
"quantize": args.quantize,
|
lora_paths=args.lora_paths,
|
||||||
}
|
lora_scales=args.lora_scales,
|
||||||
if args.lora_paths and model_class != FIBO:
|
model_config=ModelConfig.from_name(args.model, base_model=args.base_model),
|
||||||
model_kwargs["lora_paths"] = args.lora_paths
|
)
|
||||||
model_kwargs["lora_scales"] = args.lora_scales
|
|
||||||
|
|
||||||
model = model_class(**model_kwargs)
|
|
||||||
model.save_model(args.path)
|
model.save_model(args.path)
|
||||||
|
|
||||||
|
|
||||||
@ -1,6 +1,6 @@
|
|||||||
|
from mflux.cli.parser.parsers import CommandLineParser
|
||||||
from mflux.models.flux.variants.dreambooth.dreambooth import DreamBooth
|
from mflux.models.flux.variants.dreambooth.dreambooth import DreamBooth
|
||||||
from mflux.models.flux.variants.dreambooth.dreambooth_initializer import DreamBoothInitializer
|
from mflux.models.flux.variants.dreambooth.dreambooth_initializer import DreamBoothInitializer
|
||||||
from mflux.ui.cli.parsers import CommandLineParser
|
|
||||||
from mflux.utils.exceptions import StopTrainingException
|
from mflux.utils.exceptions import StopTrainingException
|
||||||
|
|
||||||
|
|
||||||
@ -13,7 +13,7 @@ def main():
|
|||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
|
|
||||||
# 1. Initialize the required resources
|
# 1. Initialize the required resources
|
||||||
flux, runtime_config, training_spec, training_state = DreamBoothInitializer.initialize(
|
flux, config, training_spec, training_state = DreamBoothInitializer.initialize(
|
||||||
config_path=args.train_config,
|
config_path=args.train_config,
|
||||||
checkpoint_path=args.train_checkpoint,
|
checkpoint_path=args.train_checkpoint,
|
||||||
)
|
)
|
||||||
@ -22,12 +22,12 @@ def main():
|
|||||||
try:
|
try:
|
||||||
DreamBooth.train(
|
DreamBooth.train(
|
||||||
flux=flux,
|
flux=flux,
|
||||||
runtime_config=runtime_config,
|
config=config,
|
||||||
training_spec=training_spec,
|
training_spec=training_spec,
|
||||||
training_state=training_state,
|
training_state=training_state,
|
||||||
)
|
)
|
||||||
except StopTrainingException as stop_exc:
|
except StopTrainingException as stop_exc:
|
||||||
training_state.save(training_spec)
|
training_state.save(flux, training_spec)
|
||||||
print(stop_exc)
|
print(stop_exc)
|
||||||
|
|
||||||
|
|
||||||
4
src/mflux/models/common/config/__init__.py
Normal file
4
src/mflux/models/common/config/__init__.py
Normal file
@ -0,0 +1,4 @@
|
|||||||
|
from mflux.models.common.config.config import Config
|
||||||
|
from mflux.models.common.config.model_config import ModelConfig
|
||||||
|
|
||||||
|
__all__ = ["Config", "ModelConfig"]
|
||||||
148
src/mflux/models/common/config/config.py
Normal file
148
src/mflux/models/common/config/config.py
Normal file
@ -0,0 +1,148 @@
|
|||||||
|
import logging
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import mlx.core as mx
|
||||||
|
from tqdm import tqdm
|
||||||
|
|
||||||
|
from mflux.models.common.config.model_config import ModelConfig
|
||||||
|
from mflux.models.common.schedulers import SCHEDULER_REGISTRY, try_import_external_scheduler
|
||||||
|
from mflux.models.common.schedulers.linear_scheduler import LinearScheduler
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class Config:
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
model_config: ModelConfig,
|
||||||
|
num_inference_steps: int = 4,
|
||||||
|
height: int = 1024,
|
||||||
|
width: int = 1024,
|
||||||
|
guidance: float = 4.0,
|
||||||
|
image_path: Path | str | None = None,
|
||||||
|
image_strength: float | None = None,
|
||||||
|
depth_image_path: Path | str | None = None,
|
||||||
|
redux_image_paths: list[Path | str] | None = None,
|
||||||
|
redux_image_strengths: list[float] | None = None,
|
||||||
|
masked_image_path: Path | str | None = None,
|
||||||
|
controlnet_strength: float | None = None,
|
||||||
|
scheduler: str = "linear",
|
||||||
|
):
|
||||||
|
# Ensure dimensions are multiples of 16
|
||||||
|
if width % 16 != 0 or height % 16 != 0:
|
||||||
|
logger.warning("Width and height should be multiples of 16. Rounding down.")
|
||||||
|
|
||||||
|
self.model_config = model_config
|
||||||
|
self._num_inference_steps = num_inference_steps
|
||||||
|
self._height = 16 * (height // 16)
|
||||||
|
self._width = 16 * (width // 16)
|
||||||
|
self._guidance = guidance
|
||||||
|
self._image_path = Path(image_path) if isinstance(image_path, str) else image_path
|
||||||
|
self._image_strength = image_strength
|
||||||
|
self._depth_image_path = Path(depth_image_path) if isinstance(depth_image_path, str) else depth_image_path
|
||||||
|
self._redux_image_paths = (
|
||||||
|
[Path(p) if isinstance(p, str) else p for p in redux_image_paths] if redux_image_paths else None
|
||||||
|
)
|
||||||
|
self._redux_image_strengths = redux_image_strengths
|
||||||
|
self._masked_image_path = Path(masked_image_path) if isinstance(masked_image_path, str) else masked_image_path
|
||||||
|
self._controlnet_strength = controlnet_strength
|
||||||
|
self._scheduler_str = scheduler
|
||||||
|
self._scheduler = None
|
||||||
|
self._time_steps = None
|
||||||
|
|
||||||
|
@property
|
||||||
|
def height(self) -> int:
|
||||||
|
return self._height
|
||||||
|
|
||||||
|
@property
|
||||||
|
def width(self) -> int:
|
||||||
|
return self._width
|
||||||
|
|
||||||
|
@width.setter
|
||||||
|
def width(self, value):
|
||||||
|
self._width = value
|
||||||
|
|
||||||
|
@property
|
||||||
|
def guidance(self) -> float:
|
||||||
|
return self._guidance
|
||||||
|
|
||||||
|
@property
|
||||||
|
def num_inference_steps(self) -> int:
|
||||||
|
return self._num_inference_steps
|
||||||
|
|
||||||
|
@property
|
||||||
|
def precision(self) -> mx.Dtype:
|
||||||
|
return ModelConfig.precision
|
||||||
|
|
||||||
|
@property
|
||||||
|
def num_train_steps(self) -> int:
|
||||||
|
return self.model_config.num_train_steps
|
||||||
|
|
||||||
|
@property
|
||||||
|
def image_path(self) -> Path | None:
|
||||||
|
return self._image_path
|
||||||
|
|
||||||
|
@property
|
||||||
|
def image_strength(self) -> float | None:
|
||||||
|
return self._image_strength
|
||||||
|
|
||||||
|
@property
|
||||||
|
def depth_image_path(self) -> Path | None:
|
||||||
|
return self._depth_image_path
|
||||||
|
|
||||||
|
@property
|
||||||
|
def redux_image_paths(self) -> list[Path] | None:
|
||||||
|
return self._redux_image_paths
|
||||||
|
|
||||||
|
@property
|
||||||
|
def redux_image_strengths(self) -> list[float] | None:
|
||||||
|
return self._redux_image_strengths
|
||||||
|
|
||||||
|
@property
|
||||||
|
def masked_image_path(self) -> Path | None:
|
||||||
|
return self._masked_image_path
|
||||||
|
|
||||||
|
@property
|
||||||
|
def init_time_step(self) -> int:
|
||||||
|
is_img2img = (
|
||||||
|
self._image_path is not None and
|
||||||
|
self._image_strength is not None and
|
||||||
|
self._image_strength > 0.0
|
||||||
|
) # fmt: off
|
||||||
|
|
||||||
|
if is_img2img:
|
||||||
|
# 1. Clamp strength to [0, 1]
|
||||||
|
strength = max(0.0, min(1.0, self._image_strength)) # type: ignore
|
||||||
|
|
||||||
|
# 2. Return start time in [1, floor(num_steps * strength)]
|
||||||
|
return max(1, int(self._num_inference_steps * strength)) # type: ignore
|
||||||
|
else:
|
||||||
|
return 0
|
||||||
|
|
||||||
|
@property
|
||||||
|
def time_steps(self) -> tqdm:
|
||||||
|
if self._time_steps is None:
|
||||||
|
self._time_steps = tqdm(range(self.init_time_step, self.num_inference_steps))
|
||||||
|
return self._time_steps
|
||||||
|
|
||||||
|
@property
|
||||||
|
def controlnet_strength(self) -> float | None:
|
||||||
|
return self._controlnet_strength
|
||||||
|
|
||||||
|
@property
|
||||||
|
def scheduler(self):
|
||||||
|
if self._scheduler is not None:
|
||||||
|
return self._scheduler
|
||||||
|
|
||||||
|
if self._scheduler_str == "linear":
|
||||||
|
self._scheduler = LinearScheduler(self)
|
||||||
|
elif (registered_scheduler := SCHEDULER_REGISTRY.get(self._scheduler_str, None)) is not None:
|
||||||
|
self._scheduler = registered_scheduler(self)
|
||||||
|
elif "." in self._scheduler_str:
|
||||||
|
# this raises ValueError if scheduler is not importable
|
||||||
|
scheduler_cls = try_import_external_scheduler(self._scheduler_str)
|
||||||
|
self._scheduler = scheduler_cls(self)
|
||||||
|
else:
|
||||||
|
raise NotImplementedError(f"The scheduler {self._scheduler_str!r} is not implemented by mflux.")
|
||||||
|
|
||||||
|
return self._scheduler
|
||||||
@ -1,10 +1,14 @@
|
|||||||
from functools import lru_cache
|
from functools import lru_cache
|
||||||
from typing import Literal
|
from typing import Literal
|
||||||
|
|
||||||
from mflux.utils.exceptions import InvalidBaseModel, ModelConfigError
|
import mlx.core as mx
|
||||||
|
|
||||||
|
from mflux.models.common.resolution.config_resolution import ConfigResolution
|
||||||
|
|
||||||
|
|
||||||
class ModelConfig:
|
class ModelConfig:
|
||||||
|
precision: mx.Dtype = mx.bfloat16
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
aliases: list[str],
|
aliases: list[str],
|
||||||
@ -99,6 +103,11 @@ class ModelConfig:
|
|||||||
def fibo() -> "ModelConfig":
|
def fibo() -> "ModelConfig":
|
||||||
return AVAILABLE_MODELS["fibo"]
|
return AVAILABLE_MODELS["fibo"]
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
@lru_cache
|
||||||
|
def z_image_turbo() -> "ModelConfig":
|
||||||
|
return AVAILABLE_MODELS["z-image-turbo"]
|
||||||
|
|
||||||
def x_embedder_input_dim(self) -> int:
|
def x_embedder_input_dim(self) -> int:
|
||||||
if "Fill" in self.model_name:
|
if "Fill" in self.model_name:
|
||||||
return 384
|
return 384
|
||||||
@ -115,56 +124,7 @@ class ModelConfig:
|
|||||||
model_name: str,
|
model_name: str,
|
||||||
base_model: Literal["dev", "schnell", "krea-dev"] | None = None,
|
base_model: Literal["dev", "schnell", "krea-dev"] | None = None,
|
||||||
) -> "ModelConfig":
|
) -> "ModelConfig":
|
||||||
# 0. Get all base models (where base_model is None) sorted by priority
|
return ConfigResolution.resolve(model_name=model_name, base_model=base_model)
|
||||||
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 aliases or full name
|
|
||||||
for base in base_models:
|
|
||||||
if model_name == base.model_name or model_name in base.aliases:
|
|
||||||
return base
|
|
||||||
|
|
||||||
# 2. Validate explicit base_model
|
|
||||||
allowed_names = []
|
|
||||||
for base in base_models:
|
|
||||||
allowed_names.extend(base.aliases + [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 (check all aliases)
|
|
||||||
default_base = next((b for b in base_models if base_model == b.model_name or base_model in b.aliases), None)
|
|
||||||
else:
|
|
||||||
# Infer from model_name substring - prefer longer matches (more specific)
|
|
||||||
# Use case-insensitive matching for better compatibility
|
|
||||||
model_name_lower = model_name.lower()
|
|
||||||
matching_bases = [
|
|
||||||
(b, alias) for b in base_models for alias in b.aliases if alias and alias.lower() in model_name_lower
|
|
||||||
]
|
|
||||||
|
|
||||||
if matching_bases:
|
|
||||||
# Sort by alias length descending, then by priority ascending
|
|
||||||
default_base = sorted(matching_bases, key=lambda x: (-len(x[1]), x[0].priority))[0][0]
|
|
||||||
else:
|
|
||||||
default_base = None
|
|
||||||
if not default_base:
|
|
||||||
raise ModelConfigError(f"Cannot infer base_model from {model_name}")
|
|
||||||
|
|
||||||
# 4. Construct the config
|
|
||||||
return ModelConfig(
|
|
||||||
aliases=default_base.aliases,
|
|
||||||
model_name=model_name,
|
|
||||||
base_model=default_base.model_name,
|
|
||||||
controlnet_model=default_base.controlnet_model,
|
|
||||||
custom_transformer_model=default_base.custom_transformer_model,
|
|
||||||
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 = {
|
AVAILABLE_MODELS = {
|
||||||
@ -336,4 +296,16 @@ AVAILABLE_MODELS = {
|
|||||||
requires_sigma_shift=False,
|
requires_sigma_shift=False,
|
||||||
priority=13,
|
priority=13,
|
||||||
),
|
),
|
||||||
|
"z-image-turbo": ModelConfig(
|
||||||
|
aliases=["z-image-turbo", "z-image", "zimage-turbo", "zimage"],
|
||||||
|
model_name="Tongyi-MAI/Z-Image-Turbo",
|
||||||
|
base_model=None,
|
||||||
|
controlnet_model=None,
|
||||||
|
custom_transformer_model=None,
|
||||||
|
num_train_steps=1000,
|
||||||
|
max_sequence_length=512,
|
||||||
|
supports_guidance=False, # Turbo model uses guidance_scale=0
|
||||||
|
requires_sigma_shift=True,
|
||||||
|
priority=14,
|
||||||
|
),
|
||||||
}
|
}
|
||||||
@ -1,5 +1,5 @@
|
|||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING, TypeAlias
|
||||||
|
|
||||||
import mlx.core as mx
|
import mlx.core as mx
|
||||||
from mlx import nn
|
from mlx import nn
|
||||||
@ -10,13 +10,16 @@ if TYPE_CHECKING:
|
|||||||
from mflux.models.fibo.latent_creator.fibo_latent_creator import FiboLatentCreator
|
from mflux.models.fibo.latent_creator.fibo_latent_creator import FiboLatentCreator
|
||||||
from mflux.models.flux.latent_creator.flux_latent_creator import FluxLatentCreator
|
from mflux.models.flux.latent_creator.flux_latent_creator import FluxLatentCreator
|
||||||
from mflux.models.qwen.latent_creator.qwen_latent_creator import QwenLatentCreator
|
from mflux.models.qwen.latent_creator.qwen_latent_creator import QwenLatentCreator
|
||||||
|
from mflux.models.z_image.latent_creator.z_image_latent_creator import ZImageLatentCreator
|
||||||
|
|
||||||
|
LatentCreatorType: TypeAlias = type[FiboLatentCreator | FluxLatentCreator | QwenLatentCreator | ZImageLatentCreator]
|
||||||
|
|
||||||
|
|
||||||
class Img2Img:
|
class Img2Img:
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
vae: nn.Module,
|
vae: nn.Module,
|
||||||
latent_creator: type["FiboLatentCreator"] | type["FluxLatentCreator"] | type["QwenLatentCreator"],
|
latent_creator: "LatentCreatorType",
|
||||||
sigmas: mx.array,
|
sigmas: mx.array,
|
||||||
init_time_step: int,
|
init_time_step: int,
|
||||||
image_path: str | Path | None,
|
image_path: str | Path | None,
|
||||||
|
|||||||
@ -1,101 +0,0 @@
|
|||||||
import shutil
|
|
||||||
from pathlib import Path
|
|
||||||
|
|
||||||
from mflux.ui.defaults import MFLUX_LORA_CACHE_DIR
|
|
||||||
from mflux.utils.download import snapshot_download
|
|
||||||
|
|
||||||
|
|
||||||
class LoRAHuggingFaceDownloader:
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def download_loras(
|
|
||||||
lora_names: list[str] | None = None,
|
|
||||||
repo_id: str | None = None,
|
|
||||||
cache_dir: Path | str | None = None,
|
|
||||||
model_name: str = "LoRA",
|
|
||||||
) -> list[str]:
|
|
||||||
if not lora_names or not repo_id:
|
|
||||||
return []
|
|
||||||
|
|
||||||
lora_paths = []
|
|
||||||
for lora_name in lora_names:
|
|
||||||
lora_path = LoRAHuggingFaceDownloader.download_lora(
|
|
||||||
repo_id=repo_id,
|
|
||||||
lora_name=lora_name,
|
|
||||||
cache_dir=cache_dir,
|
|
||||||
model_name=model_name,
|
|
||||||
)
|
|
||||||
lora_paths.append(lora_path)
|
|
||||||
|
|
||||||
return lora_paths
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def download_lora(
|
|
||||||
repo_id: str,
|
|
||||||
lora_name: str,
|
|
||||||
cache_dir: Path | str | None = None,
|
|
||||||
model_name: str = "LoRA", # For logging purposes
|
|
||||||
) -> str:
|
|
||||||
# Ensure cache_dir is a Path object
|
|
||||||
if cache_dir is None:
|
|
||||||
cache_path = MFLUX_LORA_CACHE_DIR
|
|
||||||
else:
|
|
||||||
cache_path = Path(cache_dir)
|
|
||||||
|
|
||||||
cache_path.mkdir(parents=True, exist_ok=True)
|
|
||||||
|
|
||||||
# Check if already cached
|
|
||||||
cached_file_path = cache_path / lora_name
|
|
||||||
if cached_file_path.exists() and cached_file_path.is_file():
|
|
||||||
try:
|
|
||||||
# Verify the file is actually readable (catches broken symlinks)
|
|
||||||
with open(cached_file_path, "rb") as f:
|
|
||||||
f.read(1) # Try to read just 1 byte to verify it works
|
|
||||||
print(f"Using cached {model_name} LoRA: {cached_file_path}")
|
|
||||||
return str(cached_file_path)
|
|
||||||
except (OSError, IOError):
|
|
||||||
# File exists but is not readable (broken symlink, permissions, etc.)
|
|
||||||
print(f"Cached {model_name} LoRA file is corrupted or inaccessible, re-downloading: {cached_file_path}")
|
|
||||||
try:
|
|
||||||
cached_file_path.unlink() # Remove the broken file/symlink
|
|
||||||
except OSError:
|
|
||||||
pass # Ignore if we can't remove it
|
|
||||||
|
|
||||||
# Download the LoRA from Hugging Face
|
|
||||||
print(f"Downloading {model_name} LoRA '{lora_name}' from {repo_id}...")
|
|
||||||
download_path = Path(
|
|
||||||
snapshot_download(
|
|
||||||
repo_id=repo_id,
|
|
||||||
allow_patterns=[f"*{lora_name}*"],
|
|
||||||
cache_dir=str(cache_path),
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
# Find the downloaded file
|
|
||||||
print(f"🔍 Searching for downloaded files in: {download_path}")
|
|
||||||
found_files = list(download_path.glob(f"**/*{lora_name}*"))
|
|
||||||
print(f"📁 Found files matching pattern: {found_files}")
|
|
||||||
|
|
||||||
for file in found_files:
|
|
||||||
print(f"📄 Checking file: {file} (suffix: {file.suffix}, size: {file.stat().st_size} bytes)")
|
|
||||||
if file.is_file() and file.suffix in [".safetensors", ".bin"]:
|
|
||||||
# Ensure the target path has the correct extension
|
|
||||||
if not lora_name.endswith(file.suffix):
|
|
||||||
target_name = f"{lora_name}{file.suffix}"
|
|
||||||
else:
|
|
||||||
target_name = lora_name
|
|
||||||
|
|
||||||
target_path = cache_path / target_name
|
|
||||||
if not target_path.exists():
|
|
||||||
# Create a symlink or copy the file
|
|
||||||
try:
|
|
||||||
target_path.symlink_to(file)
|
|
||||||
print(f"🔗 Created symlink: {target_path} -> {file}")
|
|
||||||
except (OSError, AttributeError):
|
|
||||||
shutil.copy2(file, target_path)
|
|
||||||
print(f"📋 Copied file: {file} -> {target_path}")
|
|
||||||
|
|
||||||
print(f"{model_name} LoRA downloaded to: {target_path}")
|
|
||||||
return str(target_path)
|
|
||||||
|
|
||||||
raise FileNotFoundError(f"Could not find {model_name} LoRA file '{lora_name}' in the downloaded repository.")
|
|
||||||
@ -1,87 +0,0 @@
|
|||||||
import os
|
|
||||||
from pathlib import Path
|
|
||||||
|
|
||||||
|
|
||||||
def _discover_lora_files(library_paths: list[Path]) -> dict[str, Path]:
|
|
||||||
"""
|
|
||||||
Discover all .safetensors files in the library paths and their subdirectories.
|
|
||||||
Earlier paths in the list have higher precedence for duplicate basenames.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
library_paths: List of paths to LORA library directories (in precedence order)
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Dictionary mapping basename (without extension) to full path
|
|
||||||
"""
|
|
||||||
lora_files = {}
|
|
||||||
|
|
||||||
# Process paths in reverse order so earlier paths overwrite later ones
|
|
||||||
for library_path in reversed(library_paths):
|
|
||||||
if not library_path.exists() or not library_path.is_dir():
|
|
||||||
continue
|
|
||||||
|
|
||||||
# Find all .safetensors files recursively
|
|
||||||
for safetensor_path in library_path.rglob("*.safetensors"):
|
|
||||||
# Use the basename without extension as the key
|
|
||||||
basename = safetensor_path.stem
|
|
||||||
|
|
||||||
# Skip files with digit-only names (0-9) in transformer directories
|
|
||||||
if basename.isdigit() and safetensor_path.parent.name == "transformer":
|
|
||||||
continue
|
|
||||||
|
|
||||||
# Earlier paths in the list take precedence (overwrite)
|
|
||||||
lora_files[basename] = safetensor_path.resolve()
|
|
||||||
|
|
||||||
return lora_files
|
|
||||||
|
|
||||||
|
|
||||||
# Global registry that will be populated on module import
|
|
||||||
_LORA_REGISTRY: dict[str, Path] = {}
|
|
||||||
|
|
||||||
|
|
||||||
def _initialize_registry() -> None:
|
|
||||||
"""Initialize the global LORA registry from LORA_LIBRARY_PATH environment variable."""
|
|
||||||
global _LORA_REGISTRY
|
|
||||||
|
|
||||||
library_path_env = os.environ.get("LORA_LIBRARY_PATH")
|
|
||||||
if library_path_env:
|
|
||||||
# Split by colon to support multiple paths
|
|
||||||
library_paths = [Path(p.strip()) for p in library_path_env.split(":") if p.strip()]
|
|
||||||
_LORA_REGISTRY = _discover_lora_files(library_paths)
|
|
||||||
|
|
||||||
|
|
||||||
def get_lora_path(path_or_name: str) -> str:
|
|
||||||
"""
|
|
||||||
Get the full path for a LORA file, resolving from library if needed.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
path_or_name: Either a full path or a basename that exists in the library
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
The resolved path as a string
|
|
||||||
|
|
||||||
Raises:
|
|
||||||
FileNotFoundError: If the file cannot be found either as a path or in the registry
|
|
||||||
"""
|
|
||||||
# If it's already a path that exists, return it as-is
|
|
||||||
path = Path(path_or_name)
|
|
||||||
if path.exists():
|
|
||||||
return str(path)
|
|
||||||
|
|
||||||
# Otherwise, check if it's in the registry
|
|
||||||
if path_or_name in _LORA_REGISTRY:
|
|
||||||
return str(_LORA_REGISTRY[path_or_name])
|
|
||||||
|
|
||||||
# If not found, raise FileNotFoundError
|
|
||||||
raise FileNotFoundError(
|
|
||||||
f"LoRA file not found: '{path_or_name}'. File does not exist and is not in the LoRA library."
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def get_registry() -> dict[str, Path]:
|
|
||||||
"""Get a copy of the current LORA registry."""
|
|
||||||
return _LORA_REGISTRY.copy()
|
|
||||||
|
|
||||||
|
|
||||||
# Initialize the registry when the module is imported
|
|
||||||
_initialize_registry()
|
|
||||||
@ -1,6 +1,7 @@
|
|||||||
import re
|
import re
|
||||||
|
from collections.abc import Callable
|
||||||
|
from dataclasses import dataclass
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Dict, Tuple
|
|
||||||
|
|
||||||
import mlx.core as mx
|
import mlx.core as mx
|
||||||
import mlx.nn as nn
|
import mlx.nn as nn
|
||||||
@ -8,6 +9,16 @@ import mlx.nn as nn
|
|||||||
from mflux.models.common.lora.layer.fused_linear_lora_layer import FusedLoRALinear
|
from mflux.models.common.lora.layer.fused_linear_lora_layer import FusedLoRALinear
|
||||||
from mflux.models.common.lora.layer.linear_lora_layer import LoRALinear
|
from mflux.models.common.lora.layer.linear_lora_layer import LoRALinear
|
||||||
from mflux.models.common.lora.mapping.lora_mapping import LoRATarget
|
from mflux.models.common.lora.mapping.lora_mapping import LoRATarget
|
||||||
|
from mflux.models.common.resolution.lora_resolution import LoraResolution
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class PatternMatch:
|
||||||
|
source_pattern: str
|
||||||
|
target_path: str
|
||||||
|
matrix_name: str # "lora_A", "lora_B", or "alpha"
|
||||||
|
transpose: bool
|
||||||
|
transform: Callable[[mx.array], mx.array] | None = None
|
||||||
|
|
||||||
|
|
||||||
class LoRALoader:
|
class LoRALoader:
|
||||||
@ -15,30 +26,34 @@ class LoRALoader:
|
|||||||
def load_and_apply_lora(
|
def load_and_apply_lora(
|
||||||
lora_mapping: list[LoRATarget],
|
lora_mapping: list[LoRATarget],
|
||||||
transformer: nn.Module,
|
transformer: nn.Module,
|
||||||
lora_files: list[str],
|
lora_paths: list[str] | None = None,
|
||||||
lora_scales: list[float] | None = None,
|
lora_scales: list[float] | None = None,
|
||||||
) -> None:
|
) -> tuple[list[str], list[float]]:
|
||||||
if not lora_files:
|
resolved_paths = LoraResolution.resolve_paths(lora_paths)
|
||||||
return
|
if not resolved_paths:
|
||||||
|
return resolved_paths, []
|
||||||
|
|
||||||
# Validate scales - handle both None and empty list cases
|
resolved_scales = LoraResolution.resolve_scales(lora_scales, len(resolved_paths))
|
||||||
if lora_scales is None or len(lora_scales) == 0:
|
if len(resolved_scales) != len(resolved_paths):
|
||||||
lora_scales = [1.0] * len(lora_files)
|
|
||||||
elif len(lora_scales) != len(lora_files):
|
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Number of LoRA scales ({len(lora_scales)}) must match number of LoRA files ({len(lora_files)})"
|
f"Number of LoRA scales ({len(resolved_scales)}) must match number of LoRA files ({len(resolved_paths)})"
|
||||||
)
|
)
|
||||||
|
|
||||||
print(f"📦 Loading {len(lora_files)} LoRA file(s)...")
|
print(f"📦 Loading {len(resolved_paths)} LoRA file(s)...")
|
||||||
|
|
||||||
for lora_file, scale in zip(lora_files, lora_scales):
|
for lora_file, scale in zip(resolved_paths, resolved_scales):
|
||||||
LoRALoader._apply_single_lora(transformer, lora_file, scale, lora_mapping)
|
LoRALoader._apply_single_lora(transformer, lora_file, scale, lora_mapping)
|
||||||
|
|
||||||
print("✅ All LoRA weights applied successfully")
|
print("✅ All LoRA weights applied successfully")
|
||||||
|
|
||||||
|
return resolved_paths, resolved_scales
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _apply_single_lora(
|
def _apply_single_lora(
|
||||||
transformer: nn.Module, lora_file: str, scale: float, lora_mapping: list[LoRATarget]
|
transformer: nn.Module,
|
||||||
|
lora_file: str,
|
||||||
|
scale: float,
|
||||||
|
lora_mapping: list[LoRATarget],
|
||||||
) -> None:
|
) -> None:
|
||||||
# Load the LoRA weights
|
# Load the LoRA weights
|
||||||
if not Path(lora_file).exists():
|
if not Path(lora_file).exists():
|
||||||
@ -49,73 +64,136 @@ class LoRALoader:
|
|||||||
|
|
||||||
try:
|
try:
|
||||||
weights = dict(mx.load(lora_file, return_metadata=True)[0].items())
|
weights = dict(mx.load(lora_file, return_metadata=True)[0].items())
|
||||||
mx.eval(weights)
|
|
||||||
except (FileNotFoundError, ValueError, RuntimeError) as e:
|
except (FileNotFoundError, ValueError, RuntimeError) as e:
|
||||||
print(f"❌ Failed to load LoRA file: {e}")
|
print(f"❌ Failed to load LoRA file: {e}")
|
||||||
return
|
return
|
||||||
|
|
||||||
# Apply LoRA using the provided mapping
|
# Build pattern mappings from LoRATargets
|
||||||
flat_mapping = LoRALoader._get_flat_mapping(lora_mapping)
|
pattern_mappings = LoRALoader._build_pattern_mappings(lora_mapping)
|
||||||
applied_count = LoRALoader._apply_lora_with_mapping(transformer, weights, scale, flat_mapping)
|
|
||||||
|
|
||||||
print(f" ✅ Applied to {applied_count} layers")
|
# Apply LoRA using the mappings (allows multiple targets per source)
|
||||||
|
applied_count, matched_keys = LoRALoader._apply_lora_with_mapping(transformer, weights, scale, pattern_mappings)
|
||||||
|
|
||||||
|
# Report results
|
||||||
|
total_keys = len(weights)
|
||||||
|
unmatched_keys = set(weights.keys()) - matched_keys
|
||||||
|
|
||||||
|
print(f" ✅ Applied to {applied_count} layers ({len(matched_keys)}/{total_keys} keys matched)")
|
||||||
|
|
||||||
|
if unmatched_keys:
|
||||||
|
print(f" ⚠️ {len(unmatched_keys)} unmatched keys in LoRA file:")
|
||||||
|
for key in sorted(unmatched_keys)[:5]:
|
||||||
|
print(f" - {key}")
|
||||||
|
if len(unmatched_keys) > 5:
|
||||||
|
print(f" ... and {len(unmatched_keys) - 5} more")
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _build_pattern_mappings(targets: list[LoRATarget]) -> list[PatternMatch]:
|
||||||
|
mappings = []
|
||||||
|
|
||||||
|
for target in targets:
|
||||||
|
# Add up weight patterns (lora_B)
|
||||||
|
mappings.extend(
|
||||||
|
PatternMatch(
|
||||||
|
source_pattern=pattern,
|
||||||
|
target_path=target.model_path,
|
||||||
|
matrix_name="lora_B",
|
||||||
|
transpose=True,
|
||||||
|
transform=target.up_transform,
|
||||||
|
)
|
||||||
|
for pattern in target.possible_up_patterns
|
||||||
|
)
|
||||||
|
|
||||||
|
# Add down weight patterns (lora_A)
|
||||||
|
mappings.extend(
|
||||||
|
PatternMatch(
|
||||||
|
source_pattern=pattern,
|
||||||
|
target_path=target.model_path,
|
||||||
|
matrix_name="lora_A",
|
||||||
|
transpose=True,
|
||||||
|
transform=target.down_transform,
|
||||||
|
)
|
||||||
|
for pattern in target.possible_down_patterns
|
||||||
|
)
|
||||||
|
|
||||||
|
# Add alpha patterns (no transpose, no transform)
|
||||||
|
mappings.extend(
|
||||||
|
PatternMatch(
|
||||||
|
source_pattern=pattern,
|
||||||
|
target_path=target.model_path,
|
||||||
|
matrix_name="alpha",
|
||||||
|
transpose=False,
|
||||||
|
transform=None,
|
||||||
|
)
|
||||||
|
for pattern in target.possible_alpha_patterns
|
||||||
|
)
|
||||||
|
|
||||||
|
return mappings
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _apply_lora_with_mapping(
|
def _apply_lora_with_mapping(
|
||||||
transformer: nn.Module, weights: dict, scale: float, lora_mappings: Dict[str, Tuple[str, str, bool]]
|
transformer: nn.Module,
|
||||||
) -> int:
|
weights: dict,
|
||||||
|
scale: float,
|
||||||
|
pattern_mappings: list[PatternMatch],
|
||||||
|
) -> tuple[int, set]:
|
||||||
applied_count = 0
|
applied_count = 0
|
||||||
lora_data_by_target = {}
|
lora_data_by_target: dict[str, dict] = {}
|
||||||
|
matched_keys: set[str] = set()
|
||||||
|
|
||||||
# Group LoRA weights by their target layers
|
# For each weight key, find ALL matching patterns (not just first)
|
||||||
|
# This allows multiple targets to use the same source (e.g., QKV split)
|
||||||
for weight_key, weight_value in weights.items():
|
for weight_key, weight_value in weights.items():
|
||||||
found_mapping = None
|
for mapping in pattern_mappings:
|
||||||
block_idx = None
|
match_result = LoRALoader._match_pattern(weight_key, mapping.source_pattern)
|
||||||
|
if match_result is None:
|
||||||
|
continue
|
||||||
|
|
||||||
# Pattern matching logic
|
matched_keys.add(weight_key)
|
||||||
for pattern, mapping_info in lora_mappings.items():
|
block_idx = match_result
|
||||||
if "{block}" in pattern:
|
|
||||||
# Extract block number from the weight key - try both . and _ separators
|
|
||||||
# This handles both standard LoRA formats (dot-separated) and other formats (underscore-separated)
|
|
||||||
# Find all numbers in the weight key
|
|
||||||
numbers_in_key = re.findall(r"\d+", weight_key)
|
|
||||||
for num_str in numbers_in_key:
|
|
||||||
try:
|
|
||||||
test_block_idx = int(num_str)
|
|
||||||
concrete_pattern = pattern.format(block=test_block_idx)
|
|
||||||
if weight_key == concrete_pattern:
|
|
||||||
found_mapping = mapping_info
|
|
||||||
block_idx = test_block_idx
|
|
||||||
break
|
|
||||||
except (ValueError, KeyError):
|
|
||||||
continue
|
|
||||||
if found_mapping:
|
|
||||||
break
|
|
||||||
else:
|
|
||||||
if weight_key == pattern:
|
|
||||||
found_mapping = mapping_info
|
|
||||||
break
|
|
||||||
|
|
||||||
if found_mapping is None:
|
# Resolve target path with block index if needed
|
||||||
continue
|
target_path = mapping.target_path
|
||||||
|
if block_idx is not None and "{block}" in target_path:
|
||||||
|
target_path = target_path.format(block=block_idx)
|
||||||
|
|
||||||
target_path, matrix_name, transpose = found_mapping
|
# Apply transform if specified
|
||||||
|
transformed_value = weight_value
|
||||||
|
if mapping.transform is not None:
|
||||||
|
transformed_value = mapping.transform(weight_value)
|
||||||
|
|
||||||
# Handle block substitution in target path
|
# Apply transpose if needed
|
||||||
if block_idx is not None and "{block}" in target_path:
|
if mapping.transpose:
|
||||||
target_path = target_path.format(block=block_idx)
|
transformed_value = transformed_value.T
|
||||||
|
|
||||||
if target_path not in lora_data_by_target:
|
# Store for this target
|
||||||
lora_data_by_target[target_path] = {}
|
if target_path not in lora_data_by_target:
|
||||||
|
lora_data_by_target[target_path] = {}
|
||||||
|
|
||||||
lora_data_by_target[target_path][matrix_name] = (weight_value, transpose)
|
lora_data_by_target[target_path][mapping.matrix_name] = transformed_value
|
||||||
|
|
||||||
# Apply LoRA to each target
|
# Apply LoRA to each target
|
||||||
for target_path, lora_data in lora_data_by_target.items():
|
for target_path, lora_data in lora_data_by_target.items():
|
||||||
if LoRALoader._apply_lora_matrices_to_target(transformer, target_path, lora_data, scale):
|
if LoRALoader._apply_lora_matrices_to_target(transformer, target_path, lora_data, scale):
|
||||||
applied_count += 1
|
applied_count += 1
|
||||||
|
|
||||||
return applied_count
|
return applied_count, matched_keys
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _match_pattern(weight_key: str, pattern: str) -> int | None:
|
||||||
|
if "{block}" in pattern:
|
||||||
|
# Find all numbers in the weight key
|
||||||
|
numbers_in_key = re.findall(r"\d+", weight_key)
|
||||||
|
for num_str in numbers_in_key:
|
||||||
|
test_block_idx = int(num_str)
|
||||||
|
concrete_pattern = pattern.replace("{block}", str(test_block_idx))
|
||||||
|
if weight_key == concrete_pattern:
|
||||||
|
return test_block_idx
|
||||||
|
return None
|
||||||
|
else:
|
||||||
|
if weight_key == pattern:
|
||||||
|
return 0 # Return 0 to indicate match (no block)
|
||||||
|
return None
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _apply_lora_matrices_to_target(transformer: nn.Module, target_path: str, lora_data: dict, scale: float) -> bool:
|
def _apply_lora_matrices_to_target(transformer: nn.Module, target_path: str, lora_data: dict, scale: float) -> bool:
|
||||||
@ -138,19 +216,14 @@ class LoRALoader:
|
|||||||
print(f"❌ Missing required LoRA matrices for {target_path}")
|
print(f"❌ Missing required LoRA matrices for {target_path}")
|
||||||
return False
|
return False
|
||||||
|
|
||||||
lora_A, transpose_A = lora_data["lora_A"]
|
# Values are already transformed and transposed
|
||||||
lora_B, transpose_B = lora_data["lora_B"]
|
lora_A = lora_data["lora_A"]
|
||||||
|
lora_B = lora_data["lora_B"]
|
||||||
# Handle transposition
|
|
||||||
if transpose_A:
|
|
||||||
lora_A = lora_A.T
|
|
||||||
if transpose_B:
|
|
||||||
lora_B = lora_B.T
|
|
||||||
|
|
||||||
# Handle alpha scaling
|
# Handle alpha scaling
|
||||||
alpha_scale = 1.0
|
alpha_scale = 1.0
|
||||||
if "alpha" in lora_data:
|
if "alpha" in lora_data:
|
||||||
alpha_value, _ = lora_data["alpha"]
|
alpha_value = lora_data["alpha"]
|
||||||
rank = lora_A.shape[1]
|
rank = lora_A.shape[1]
|
||||||
alpha_scale = float(alpha_value) / rank
|
alpha_scale = float(alpha_value) / rank
|
||||||
|
|
||||||
@ -167,47 +240,33 @@ class LoRALoader:
|
|||||||
# Handle fusion: if the current module is already a LoRA layer, fuse them
|
# Handle fusion: if the current module is already a LoRA layer, fuse them
|
||||||
if is_lora_linear:
|
if is_lora_linear:
|
||||||
print(f" 🔀 Fusing with existing LoRA at {target_path}")
|
print(f" 🔀 Fusing with existing LoRA at {target_path}")
|
||||||
# Create a temporary LoRA layer from the base linear of the existing LoRA
|
|
||||||
lora_layer = LoRALinear.from_linear(current_module.linear, r=lora_A.shape[1], scale=effective_scale)
|
lora_layer = LoRALinear.from_linear(current_module.linear, r=lora_A.shape[1], scale=effective_scale)
|
||||||
# Set the LoRA matrices
|
|
||||||
lora_layer.lora_A = lora_A
|
lora_layer.lora_A = lora_A
|
||||||
lora_layer.lora_B = lora_B
|
lora_layer.lora_B = lora_B
|
||||||
# Apply alpha scaling to the matrices if present
|
|
||||||
if "alpha" in lora_data:
|
if "alpha" in lora_data:
|
||||||
lora_layer.lora_B = lora_layer.lora_B * alpha_scale
|
lora_layer.lora_B = lora_layer.lora_B * alpha_scale
|
||||||
|
|
||||||
# Create fused layer with the existing LoRA and the new one
|
|
||||||
fused_layer = FusedLoRALinear(base_linear=current_module.linear, loras=[current_module, lora_layer])
|
fused_layer = FusedLoRALinear(base_linear=current_module.linear, loras=[current_module, lora_layer])
|
||||||
replacement_layer = fused_layer
|
replacement_layer = fused_layer
|
||||||
elif is_fused_linear:
|
elif is_fused_linear:
|
||||||
print(f" 🔀 Adding to existing fusion at {target_path}")
|
print(f" 🔀 Adding to existing fusion at {target_path}")
|
||||||
# Create a temporary LoRA layer from the base linear
|
|
||||||
lora_layer = LoRALinear.from_linear(
|
lora_layer = LoRALinear.from_linear(
|
||||||
current_module.base_linear, r=lora_A.shape[1], scale=effective_scale
|
current_module.base_linear, r=lora_A.shape[1], scale=effective_scale
|
||||||
)
|
)
|
||||||
# Set the LoRA matrices
|
|
||||||
lora_layer.lora_A = lora_A
|
lora_layer.lora_A = lora_A
|
||||||
lora_layer.lora_B = lora_B
|
lora_layer.lora_B = lora_B
|
||||||
# Apply alpha scaling to the matrices if present
|
|
||||||
if "alpha" in lora_data:
|
if "alpha" in lora_data:
|
||||||
lora_layer.lora_B = lora_layer.lora_B * alpha_scale
|
lora_layer.lora_B = lora_layer.lora_B * alpha_scale
|
||||||
|
|
||||||
# Add to existing fusion
|
|
||||||
fused_layer = FusedLoRALinear(
|
fused_layer = FusedLoRALinear(
|
||||||
base_linear=current_module.base_linear, loras=current_module.loras + [lora_layer]
|
base_linear=current_module.base_linear, loras=current_module.loras + [lora_layer]
|
||||||
)
|
)
|
||||||
replacement_layer = fused_layer
|
replacement_layer = fused_layer
|
||||||
else:
|
else:
|
||||||
# First LoRA on this layer
|
# First LoRA on this layer
|
||||||
# Create LoRA layer
|
|
||||||
lora_layer = LoRALinear.from_linear(current_module, r=lora_A.shape[1], scale=effective_scale)
|
lora_layer = LoRALinear.from_linear(current_module, r=lora_A.shape[1], scale=effective_scale)
|
||||||
# Set the LoRA matrices - use the correct dimensions from the LoRA file
|
|
||||||
lora_layer.lora_A = lora_A
|
lora_layer.lora_A = lora_A
|
||||||
lora_layer.lora_B = lora_B
|
lora_layer.lora_B = lora_B
|
||||||
# Apply alpha scaling to the matrices if present
|
|
||||||
if "alpha" in lora_data:
|
if "alpha" in lora_data:
|
||||||
lora_layer.lora_B = lora_layer.lora_B * alpha_scale
|
lora_layer.lora_B = lora_layer.lora_B * alpha_scale
|
||||||
|
|
||||||
replacement_layer = lora_layer
|
replacement_layer = lora_layer
|
||||||
|
|
||||||
# Replace the layer in the parent module
|
# Replace the layer in the parent module
|
||||||
@ -228,22 +287,3 @@ class LoRALoader:
|
|||||||
else:
|
else:
|
||||||
print(f"❌ Target layer {target_path} is not a linear layer")
|
print(f"❌ Target layer {target_path} is not a linear layer")
|
||||||
return False
|
return False
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _get_flat_mapping(targets: list[LoRATarget]) -> Dict[str, Tuple[str, str, bool]]:
|
|
||||||
flat_mapping = {}
|
|
||||||
|
|
||||||
for target in targets:
|
|
||||||
# Add up weight patterns (lora_B, transposed)
|
|
||||||
for pattern in target.possible_up_patterns:
|
|
||||||
flat_mapping[pattern] = (target.model_path, "lora_B", True)
|
|
||||||
|
|
||||||
# Add down weight patterns (lora_A, transposed)
|
|
||||||
for pattern in target.possible_down_patterns:
|
|
||||||
flat_mapping[pattern] = (target.model_path, "lora_A", True)
|
|
||||||
|
|
||||||
# Add alpha patterns (no transpose)
|
|
||||||
for pattern in target.possible_alpha_patterns:
|
|
||||||
flat_mapping[pattern] = (target.model_path, "alpha", False)
|
|
||||||
|
|
||||||
return flat_mapping
|
|
||||||
|
|||||||
@ -1,17 +1,21 @@
|
|||||||
from dataclasses import dataclass
|
from collections.abc import Callable
|
||||||
|
from dataclasses import dataclass, field
|
||||||
from typing import List, Protocol
|
from typing import List, Protocol
|
||||||
|
|
||||||
|
import mlx.core as mx
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class LoRATarget:
|
class LoRATarget:
|
||||||
model_path: str
|
model_path: str
|
||||||
possible_up_patterns: List[str]
|
possible_up_patterns: List[str]
|
||||||
possible_down_patterns: List[str]
|
possible_down_patterns: List[str]
|
||||||
possible_alpha_patterns: List[str]
|
possible_alpha_patterns: List[str] = field(default_factory=list)
|
||||||
|
up_transform: Callable[[mx.array], mx.array] | None = None
|
||||||
|
down_transform: Callable[[mx.array], mx.array] | None = None
|
||||||
|
|
||||||
|
|
||||||
class LoRAMapping(Protocol):
|
class LoRAMapping(Protocol):
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def get_mapping() -> List[LoRATarget]:
|
def get_mapping() -> List[LoRATarget]:
|
||||||
return
|
return
|
||||||
|
|||||||
101
src/mflux/models/common/lora/mapping/lora_transforms.py
Normal file
101
src/mflux/models/common/lora/mapping/lora_transforms.py
Normal file
@ -0,0 +1,101 @@
|
|||||||
|
import mlx.core as mx
|
||||||
|
|
||||||
|
|
||||||
|
class LoraTransforms:
|
||||||
|
@staticmethod
|
||||||
|
def split_q_up(tensor: mx.array) -> mx.array:
|
||||||
|
return LoraTransforms._split_qkv_up(tensor, 0)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def split_k_up(tensor: mx.array) -> mx.array:
|
||||||
|
return LoraTransforms._split_qkv_up(tensor, 1)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def split_v_up(tensor: mx.array) -> mx.array:
|
||||||
|
return LoraTransforms._split_qkv_up(tensor, 2)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def split_q_down(tensor: mx.array) -> mx.array:
|
||||||
|
return LoraTransforms._split_qkv_down(tensor, 0)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def split_k_down(tensor: mx.array) -> mx.array:
|
||||||
|
return LoraTransforms._split_qkv_down(tensor, 1)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def split_v_down(tensor: mx.array) -> mx.array:
|
||||||
|
return LoraTransforms._split_qkv_down(tensor, 2)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def split_single_q_up(tensor: mx.array) -> mx.array:
|
||||||
|
return LoraTransforms._split_qkv_mlp_up(tensor, 0)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def split_single_k_up(tensor: mx.array) -> mx.array:
|
||||||
|
return LoraTransforms._split_qkv_mlp_up(tensor, 1)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def split_single_v_up(tensor: mx.array) -> mx.array:
|
||||||
|
return LoraTransforms._split_qkv_mlp_up(tensor, 2)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def split_single_mlp_up(tensor: mx.array) -> mx.array:
|
||||||
|
return LoraTransforms._split_qkv_mlp_up(tensor, 3)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def split_single_q_down(tensor: mx.array) -> mx.array:
|
||||||
|
return LoraTransforms._split_qkv_mlp_down(tensor, 0)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def split_single_k_down(tensor: mx.array) -> mx.array:
|
||||||
|
return LoraTransforms._split_qkv_mlp_down(tensor, 1)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def split_single_v_down(tensor: mx.array) -> mx.array:
|
||||||
|
return LoraTransforms._split_qkv_mlp_down(tensor, 2)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def split_single_mlp_down(tensor: mx.array) -> mx.array:
|
||||||
|
return LoraTransforms._split_qkv_mlp_down(tensor, 3)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _transpose(tensor: mx.array) -> mx.array:
|
||||||
|
return tensor.T
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _split_qkv_up(tensor: mx.array, index: int, num_splits: int = 3) -> mx.array:
|
||||||
|
split_size = tensor.shape[0] // num_splits
|
||||||
|
start = index * split_size
|
||||||
|
end = start + split_size
|
||||||
|
return tensor[start:end, :]
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _split_qkv_down(tensor: mx.array, index: int, num_splits: int = 3) -> mx.array:
|
||||||
|
rank = tensor.shape[0]
|
||||||
|
if rank % num_splits == 0:
|
||||||
|
chunk_size = rank // num_splits
|
||||||
|
start = index * chunk_size
|
||||||
|
end = start + chunk_size
|
||||||
|
return tensor[start:end, :]
|
||||||
|
else:
|
||||||
|
return tensor
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _split_qkv_mlp_up(tensor: mx.array, index: int, dims: list[int] | None = None) -> mx.array:
|
||||||
|
if dims is None:
|
||||||
|
dims = [3072, 3072, 3072, 12288]
|
||||||
|
|
||||||
|
start = sum(dims[:index])
|
||||||
|
end = start + dims[index]
|
||||||
|
return tensor[start:end, :]
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _split_qkv_mlp_down(tensor: mx.array, index: int, num_splits: int = 4) -> mx.array:
|
||||||
|
rank = tensor.shape[0]
|
||||||
|
if rank % num_splits == 0:
|
||||||
|
chunk_size = rank // num_splits
|
||||||
|
start = index * chunk_size
|
||||||
|
end = start + chunk_size
|
||||||
|
return tensor[start:end, :]
|
||||||
|
else:
|
||||||
|
return tensor
|
||||||
@ -1,112 +0,0 @@
|
|||||||
from typing import TYPE_CHECKING
|
|
||||||
|
|
||||||
import mlx.nn as nn
|
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
|
||||||
from mflux.models.fibo.weights.fibo_weight_handler import FIBOWeightHandler
|
|
||||||
from mflux.models.flux.variants.controlnet.weight_handler_controlnet import WeightHandlerControlnet
|
|
||||||
from mflux.models.flux.weights.weight_handler import WeightHandler
|
|
||||||
from mflux.models.qwen.weights.qwen_weight_handler import QwenWeightHandler
|
|
||||||
|
|
||||||
|
|
||||||
class QuantizationUtil:
|
|
||||||
@staticmethod
|
|
||||||
def quantize_model(
|
|
||||||
vae: nn.Module,
|
|
||||||
transformer: nn.Module,
|
|
||||||
t5_text_encoder: nn.Module,
|
|
||||||
clip_text_encoder: nn.Module,
|
|
||||||
quantize: int,
|
|
||||||
weights: "WeightHandler",
|
|
||||||
) -> None:
|
|
||||||
q_level = weights.meta_data.quantization_level
|
|
||||||
if q_level == "None":
|
|
||||||
q_level = None
|
|
||||||
if quantize == "None":
|
|
||||||
quantize = None
|
|
||||||
|
|
||||||
if quantize is not None or q_level is not None:
|
|
||||||
bits = int(q_level) if q_level is not None else quantize
|
|
||||||
nn.quantize(vae, bits=bits)
|
|
||||||
nn.quantize(transformer, bits=bits)
|
|
||||||
nn.quantize(t5_text_encoder, bits=bits)
|
|
||||||
nn.quantize(clip_text_encoder, bits=bits)
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def quantize_controlnet(
|
|
||||||
quantize: int,
|
|
||||||
weights: "WeightHandlerControlnet",
|
|
||||||
transformer_controlnet: nn.Module,
|
|
||||||
) -> None:
|
|
||||||
q_level = weights.meta_data.quantization_level
|
|
||||||
|
|
||||||
if quantize is not None or q_level is not None:
|
|
||||||
bits = int(q_level) if q_level is not None else quantize
|
|
||||||
nn.quantize(transformer_controlnet, bits=bits)
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def quantize_redux_models(
|
|
||||||
quantize: int,
|
|
||||||
weights: "WeightHandler",
|
|
||||||
redux_encoder: nn.Module,
|
|
||||||
siglip_vision_transformer: nn.Module,
|
|
||||||
) -> None:
|
|
||||||
q_level = weights.meta_data.quantization_level
|
|
||||||
|
|
||||||
if quantize is not None or q_level is not None:
|
|
||||||
bits = int(q_level) if q_level is not None else quantize
|
|
||||||
nn.quantize(redux_encoder, class_predicate=QuantizationUtil.quantization_predicate, bits=bits)
|
|
||||||
nn.quantize(siglip_vision_transformer, class_predicate=QuantizationUtil.quantization_predicate, bits=bits)
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def quantization_predicate(path, m):
|
|
||||||
# 1. Skip Conv2d layers
|
|
||||||
if isinstance(m, nn.Conv2d):
|
|
||||||
return False
|
|
||||||
|
|
||||||
# 2. Skip any layer with incompatible dimensions
|
|
||||||
if hasattr(m, "weight") and hasattr(m.weight, "shape"):
|
|
||||||
if m.weight.shape == (1152, 4304):
|
|
||||||
return False
|
|
||||||
|
|
||||||
if m.weight.shape[-1] % 64 != 0:
|
|
||||||
return False
|
|
||||||
|
|
||||||
# Only quantize layers that have to_quantized method
|
|
||||||
return hasattr(m, "to_quantized")
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def quantize_qwen_models(
|
|
||||||
text_encoder: nn.Module,
|
|
||||||
vae: nn.Module,
|
|
||||||
transformer: nn.Module,
|
|
||||||
quantize: int,
|
|
||||||
weights: "QwenWeightHandler",
|
|
||||||
) -> None:
|
|
||||||
q_level = weights.meta_data.quantization_level
|
|
||||||
|
|
||||||
if quantize is not None or q_level is not None:
|
|
||||||
bits = int(q_level) if q_level is not None else quantize
|
|
||||||
nn.quantize(vae, bits=bits)
|
|
||||||
nn.quantize(transformer, bits=bits)
|
|
||||||
# nn.quantize(text_encoder, bits=bits) # Quantization of text encoder causes significant semantic degradation
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def quantize_fibo_models(
|
|
||||||
vae: nn.Module,
|
|
||||||
transformer: nn.Module,
|
|
||||||
text_encoder: nn.Module | None,
|
|
||||||
quantize: int,
|
|
||||||
weights: "FIBOWeightHandler",
|
|
||||||
) -> None:
|
|
||||||
q_level = weights.meta_data.quantization_level
|
|
||||||
if q_level == "None":
|
|
||||||
q_level = None
|
|
||||||
if quantize == "None":
|
|
||||||
quantize = None
|
|
||||||
|
|
||||||
if quantize is not None or q_level is not None:
|
|
||||||
bits = int(q_level) if q_level is not None else quantize
|
|
||||||
nn.quantize(vae, class_predicate=QuantizationUtil.quantization_predicate, bits=bits)
|
|
||||||
nn.quantize(transformer, class_predicate=QuantizationUtil.quantization_predicate, bits=bits)
|
|
||||||
nn.quantize(text_encoder, class_predicate=QuantizationUtil.quantization_predicate, bits=bits)
|
|
||||||
6
src/mflux/models/common/resolution/__init__.py
Normal file
6
src/mflux/models/common/resolution/__init__.py
Normal file
@ -0,0 +1,6 @@
|
|||||||
|
from mflux.models.common.resolution.config_resolution import ConfigResolution
|
||||||
|
from mflux.models.common.resolution.lora_resolution import LoraResolution
|
||||||
|
from mflux.models.common.resolution.path_resolution import PathResolution
|
||||||
|
from mflux.models.common.resolution.quantization_resolution import QuantizationResolution
|
||||||
|
|
||||||
|
__all__ = ["ConfigResolution", "LoraResolution", "PathResolution", "QuantizationResolution"]
|
||||||
39
src/mflux/models/common/resolution/actions.py
Normal file
39
src/mflux/models/common/resolution/actions.py
Normal file
@ -0,0 +1,39 @@
|
|||||||
|
from enum import Enum
|
||||||
|
from typing import NamedTuple
|
||||||
|
|
||||||
|
|
||||||
|
class QuantizationAction(Enum):
|
||||||
|
NONE = "No quantization - use full precision"
|
||||||
|
STORED = "Use quantization level stored in weights"
|
||||||
|
REQUESTED = "Apply requested quantization on-the-fly"
|
||||||
|
|
||||||
|
|
||||||
|
class PathAction(Enum):
|
||||||
|
LOCAL = "Use local filesystem path"
|
||||||
|
HUGGINGFACE_CACHED = "Use cached HuggingFace files (no network)"
|
||||||
|
HUGGINGFACE = "Download from HuggingFace"
|
||||||
|
ERROR = "Path not found"
|
||||||
|
|
||||||
|
|
||||||
|
class LoraAction(Enum):
|
||||||
|
LOCAL = "Use local filesystem path"
|
||||||
|
REGISTRY = "Lookup in LORA_LIBRARY_PATH registry"
|
||||||
|
HUGGINGFACE_COLLECTION_CACHED = "Use cached file from HuggingFace collection (no network)"
|
||||||
|
HUGGINGFACE_COLLECTION = "Download specific file from HuggingFace collection"
|
||||||
|
HUGGINGFACE_REPO_CACHED = "Use cached HuggingFace repository (no network)"
|
||||||
|
HUGGINGFACE_REPO = "Download from HuggingFace repository"
|
||||||
|
ERROR = "LoRA not found"
|
||||||
|
|
||||||
|
|
||||||
|
class ConfigAction(Enum):
|
||||||
|
EXACT_MATCH = "Model name exactly matches a known alias"
|
||||||
|
EXPLICIT_BASE = "Use explicitly provided --base-model"
|
||||||
|
INFER_SUBSTRING = "Infer base model from substring match"
|
||||||
|
ERROR = "Cannot determine model configuration"
|
||||||
|
|
||||||
|
|
||||||
|
class Rule(NamedTuple):
|
||||||
|
priority: int
|
||||||
|
name: str
|
||||||
|
check: str
|
||||||
|
action: QuantizationAction | PathAction | LoraAction | ConfigAction
|
||||||
126
src/mflux/models/common/resolution/config_resolution.py
Normal file
126
src/mflux/models/common/resolution/config_resolution.py
Normal file
@ -0,0 +1,126 @@
|
|||||||
|
import logging
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
from mflux.models.common.resolution.actions import ConfigAction, Rule
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from mflux.models.common.config.model_config import ModelConfig
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class ConfigResolution:
|
||||||
|
RULES = frozenset(
|
||||||
|
{
|
||||||
|
Rule(priority=0, name="exact_match", check="is_exact_match", action=ConfigAction.EXACT_MATCH),
|
||||||
|
Rule(priority=1, name="explicit_base", check="has_explicit_base", action=ConfigAction.EXPLICIT_BASE),
|
||||||
|
Rule(priority=2, name="infer_substring", check="can_infer_substring", action=ConfigAction.INFER_SUBSTRING),
|
||||||
|
Rule(priority=3, name="error", check="always", action=ConfigAction.ERROR),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def resolve(model_name: str, base_model: str | None = None) -> "ModelConfig":
|
||||||
|
from mflux.models.common.config.model_config import AVAILABLE_MODELS, ModelConfig
|
||||||
|
from mflux.utils.exceptions import InvalidBaseModel, ModelConfigError
|
||||||
|
|
||||||
|
base_models = sorted(
|
||||||
|
[m for m in AVAILABLE_MODELS.values() if m.base_model is None],
|
||||||
|
key=lambda x: x.priority,
|
||||||
|
)
|
||||||
|
|
||||||
|
ctx = {
|
||||||
|
"model_name": model_name,
|
||||||
|
"base_model": base_model,
|
||||||
|
"base_models": base_models,
|
||||||
|
"ModelConfig": ModelConfig,
|
||||||
|
"InvalidBaseModel": InvalidBaseModel,
|
||||||
|
"ModelConfigError": ModelConfigError,
|
||||||
|
}
|
||||||
|
|
||||||
|
for rule in sorted(ConfigResolution.RULES, key=lambda r: r.priority):
|
||||||
|
if ConfigResolution._check(rule.check, ctx):
|
||||||
|
logger.debug(f"Config resolution: '{model_name}' → rule '{rule.name}' ({rule.action.value})")
|
||||||
|
return ConfigResolution._execute(rule.action, ctx)
|
||||||
|
|
||||||
|
raise ValueError(f"No rule matched for model_name: {model_name}")
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _check(check: str, ctx: dict) -> bool:
|
||||||
|
if check == "is_exact_match":
|
||||||
|
model_name = ctx["model_name"]
|
||||||
|
for base in ctx["base_models"]:
|
||||||
|
if model_name == base.model_name or model_name in base.aliases:
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
if check == "has_explicit_base":
|
||||||
|
return ctx["base_model"] is not None
|
||||||
|
if check == "can_infer_substring":
|
||||||
|
model_name_lower = ctx["model_name"].lower()
|
||||||
|
for base in ctx["base_models"]:
|
||||||
|
for alias in base.aliases:
|
||||||
|
if alias and alias.lower() in model_name_lower:
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
if check == "always":
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _execute(action: ConfigAction, ctx: dict) -> "ModelConfig":
|
||||||
|
model_name = ctx["model_name"]
|
||||||
|
base_model = ctx["base_model"]
|
||||||
|
base_models = ctx["base_models"]
|
||||||
|
ModelConfig = ctx["ModelConfig"]
|
||||||
|
InvalidBaseModel = ctx["InvalidBaseModel"]
|
||||||
|
ModelConfigError = ctx["ModelConfigError"]
|
||||||
|
|
||||||
|
if action == ConfigAction.EXACT_MATCH:
|
||||||
|
for base in base_models:
|
||||||
|
if model_name == base.model_name or model_name in base.aliases:
|
||||||
|
return base
|
||||||
|
raise ValueError("Exact match check passed but no match found")
|
||||||
|
|
||||||
|
if action == ConfigAction.EXPLICIT_BASE:
|
||||||
|
allowed_names = []
|
||||||
|
for base in base_models:
|
||||||
|
allowed_names.extend(base.aliases + [base.model_name])
|
||||||
|
if base_model not in allowed_names:
|
||||||
|
raise InvalidBaseModel(f"Invalid base_model. Choose one of {allowed_names}")
|
||||||
|
|
||||||
|
default_base = next(
|
||||||
|
(b for b in base_models if base_model == b.model_name or base_model in b.aliases),
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
return ConfigResolution._create_config(model_name, default_base, ModelConfig)
|
||||||
|
|
||||||
|
if action == ConfigAction.INFER_SUBSTRING:
|
||||||
|
model_name_lower = model_name.lower()
|
||||||
|
matching_bases = [
|
||||||
|
(b, alias) for b in base_models for alias in b.aliases if alias and alias.lower() in model_name_lower
|
||||||
|
]
|
||||||
|
if not matching_bases:
|
||||||
|
raise ModelConfigError(f"Cannot infer base_model from {model_name}")
|
||||||
|
|
||||||
|
default_base = sorted(matching_bases, key=lambda x: (-len(x[1]), x[0].priority))[0][0]
|
||||||
|
return ConfigResolution._create_config(model_name, default_base, ModelConfig)
|
||||||
|
|
||||||
|
if action == ConfigAction.ERROR:
|
||||||
|
raise ModelConfigError(f"Cannot infer base_model from {model_name}")
|
||||||
|
|
||||||
|
raise ValueError(f"Unknown action: {action}")
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _create_config(model_name: str, base: "ModelConfig", ModelConfig: type) -> "ModelConfig":
|
||||||
|
return ModelConfig(
|
||||||
|
aliases=base.aliases,
|
||||||
|
model_name=model_name,
|
||||||
|
base_model=base.model_name,
|
||||||
|
controlnet_model=base.controlnet_model,
|
||||||
|
custom_transformer_model=base.custom_transformer_model,
|
||||||
|
num_train_steps=base.num_train_steps,
|
||||||
|
max_sequence_length=base.max_sequence_length,
|
||||||
|
supports_guidance=base.supports_guidance,
|
||||||
|
requires_sigma_shift=base.requires_sigma_shift,
|
||||||
|
priority=base.priority,
|
||||||
|
)
|
||||||
298
src/mflux/models/common/resolution/lora_resolution.py
Normal file
298
src/mflux/models/common/resolution/lora_resolution.py
Normal file
@ -0,0 +1,298 @@
|
|||||||
|
import logging
|
||||||
|
import os
|
||||||
|
import shutil
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
from huggingface_hub import snapshot_download
|
||||||
|
from huggingface_hub.utils import LocalEntryNotFoundError
|
||||||
|
|
||||||
|
from mflux.cli.defaults.defaults import MFLUX_LORA_CACHE_DIR
|
||||||
|
from mflux.models.common.resolution.actions import LoraAction, Rule
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class LoraResolution:
|
||||||
|
RULES = frozenset(
|
||||||
|
{
|
||||||
|
Rule(priority=0, name="local", check="exists_locally", action=LoraAction.LOCAL),
|
||||||
|
Rule(priority=1, name="registry", check="in_registry", action=LoraAction.REGISTRY),
|
||||||
|
Rule(priority=2, name="collection_cached", check="is_collection_cached", action=LoraAction.HUGGINGFACE_COLLECTION_CACHED),
|
||||||
|
Rule(priority=3, name="collection_download", check="is_collection_format", action=LoraAction.HUGGINGFACE_COLLECTION),
|
||||||
|
Rule(priority=4, name="repo_cached", check="is_repo_cached", action=LoraAction.HUGGINGFACE_REPO_CACHED),
|
||||||
|
Rule(priority=5, name="repo_download", check="is_hf_format", action=LoraAction.HUGGINGFACE_REPO),
|
||||||
|
Rule(priority=6, name="error", check="always", action=LoraAction.ERROR),
|
||||||
|
}
|
||||||
|
) # fmt: off
|
||||||
|
|
||||||
|
_registry: dict[str, Path] = {}
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def resolve(path: str) -> str:
|
||||||
|
for rule in sorted(LoraResolution.RULES, key=lambda r: r.priority):
|
||||||
|
if LoraResolution._check(rule.check, path):
|
||||||
|
logger.debug(f"LoRA resolution: '{path}' → rule '{rule.name}' ({rule.action.value})")
|
||||||
|
return LoraResolution._execute(rule.action, path)
|
||||||
|
|
||||||
|
raise ValueError(f"No rule matched for LoRA path: {path}")
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def resolve_paths(paths: list[str] | None) -> list[str]:
|
||||||
|
if not paths:
|
||||||
|
return []
|
||||||
|
return [r for path in paths if (r := LoraResolution._try_resolve(path)) is not None]
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _try_resolve(path: str) -> str | None:
|
||||||
|
try:
|
||||||
|
return LoraResolution.resolve(path)
|
||||||
|
except FileNotFoundError as e:
|
||||||
|
print(f"⚠️ {e}")
|
||||||
|
return None
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def resolve_scales(scales: list[float] | None, num_paths: int) -> list[float]:
|
||||||
|
if not scales:
|
||||||
|
return [1.0] * num_paths
|
||||||
|
if len(scales) != num_paths:
|
||||||
|
print(
|
||||||
|
f"⚠️ Number of LoRA scales ({len(scales)}) doesn't match number of LoRA paths ({num_paths}). "
|
||||||
|
f"Using provided scales and defaulting remaining to 1.0."
|
||||||
|
)
|
||||||
|
# Pad with 1.0 if too few scales, truncate if too many
|
||||||
|
if len(scales) < num_paths:
|
||||||
|
return list(scales) + [1.0] * (num_paths - len(scales))
|
||||||
|
return list(scales[:num_paths])
|
||||||
|
return scales
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _is_collection_format(path: str) -> bool:
|
||||||
|
return ":" in path and "/" in path
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _is_hf_format(path: str) -> bool:
|
||||||
|
return "/" in path and path.count("/") == 1 and not path.startswith(("./", "../", "~/"))
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _check(check: str, path: str) -> bool:
|
||||||
|
if check == "exists_locally":
|
||||||
|
return Path(path).expanduser().exists()
|
||||||
|
if check == "in_registry":
|
||||||
|
return path in LoraResolution._registry
|
||||||
|
if check == "is_collection_cached":
|
||||||
|
if not LoraResolution._is_collection_format(path):
|
||||||
|
return False
|
||||||
|
repo_id, filename = path.split(":", 1)
|
||||||
|
return LoraResolution._is_collection_in_cache(repo_id, filename)
|
||||||
|
if check == "is_collection_format":
|
||||||
|
return LoraResolution._is_collection_format(path)
|
||||||
|
if check == "is_repo_cached":
|
||||||
|
if not LoraResolution._is_hf_format(path):
|
||||||
|
return False
|
||||||
|
return LoraResolution._is_repo_in_cache(path)
|
||||||
|
if check == "is_hf_format":
|
||||||
|
return LoraResolution._is_hf_format(path)
|
||||||
|
if check == "always":
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _is_collection_in_cache(repo_id: str, filename: str) -> bool:
|
||||||
|
cache_path = MFLUX_LORA_CACHE_DIR
|
||||||
|
# Check mflux cache
|
||||||
|
cached_file_path = cache_path / filename
|
||||||
|
if cached_file_path.exists() and cached_file_path.is_file():
|
||||||
|
return True
|
||||||
|
# Check HF cache
|
||||||
|
try:
|
||||||
|
snapshot_download(
|
||||||
|
repo_id=repo_id,
|
||||||
|
allow_patterns=[f"*{filename}*"],
|
||||||
|
cache_dir=str(cache_path),
|
||||||
|
local_files_only=True,
|
||||||
|
)
|
||||||
|
return True
|
||||||
|
except LocalEntryNotFoundError:
|
||||||
|
return False
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _is_repo_in_cache(repo_id: str) -> bool:
|
||||||
|
cache_path = MFLUX_LORA_CACHE_DIR
|
||||||
|
try:
|
||||||
|
snapshot_download(
|
||||||
|
repo_id=repo_id,
|
||||||
|
allow_patterns=["*.safetensors"],
|
||||||
|
cache_dir=str(cache_path),
|
||||||
|
local_files_only=True,
|
||||||
|
)
|
||||||
|
return True
|
||||||
|
except LocalEntryNotFoundError:
|
||||||
|
return False
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _execute(action: LoraAction, path: str) -> str:
|
||||||
|
if action == LoraAction.LOCAL:
|
||||||
|
return str(Path(path).expanduser())
|
||||||
|
if action == LoraAction.REGISTRY:
|
||||||
|
return str(LoraResolution._registry[path])
|
||||||
|
if action == LoraAction.HUGGINGFACE_COLLECTION_CACHED:
|
||||||
|
repo_id, filename = path.split(":", 1)
|
||||||
|
return LoraResolution._load_collection_from_cache(repo_id, filename)
|
||||||
|
if action == LoraAction.HUGGINGFACE_COLLECTION:
|
||||||
|
repo_id, filename = path.split(":", 1)
|
||||||
|
return LoraResolution._download_collection(repo_id, filename)
|
||||||
|
if action == LoraAction.HUGGINGFACE_REPO_CACHED:
|
||||||
|
return LoraResolution._load_repo_from_cache(path)
|
||||||
|
if action == LoraAction.HUGGINGFACE_REPO:
|
||||||
|
return LoraResolution._download_repo(path)
|
||||||
|
if action == LoraAction.ERROR:
|
||||||
|
raise FileNotFoundError(
|
||||||
|
f"LoRA file not found: '{path}'. File does not exist and is not in the LoRA library."
|
||||||
|
)
|
||||||
|
raise ValueError(f"Unknown action: {action}")
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _load_repo_from_cache(repo_id: str) -> str:
|
||||||
|
cache_path = MFLUX_LORA_CACHE_DIR
|
||||||
|
cache_path.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
|
download_path = Path(
|
||||||
|
snapshot_download(
|
||||||
|
repo_id=repo_id,
|
||||||
|
allow_patterns=["*.safetensors"],
|
||||||
|
cache_dir=str(cache_path),
|
||||||
|
local_files_only=True,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
safetensor_files = list(download_path.glob("*.safetensors"))
|
||||||
|
if not safetensor_files:
|
||||||
|
raise FileNotFoundError(f"No .safetensors file found in cached repo: {repo_id}")
|
||||||
|
|
||||||
|
if len(safetensor_files) > 1:
|
||||||
|
file_names = [f.name for f in safetensor_files]
|
||||||
|
raise ValueError(
|
||||||
|
f"Multiple .safetensors files found in '{repo_id}': {file_names}. "
|
||||||
|
f"Please specify which file to use with the collection format: '{repo_id}:<filename>.safetensors'"
|
||||||
|
)
|
||||||
|
|
||||||
|
return str(safetensor_files[0])
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _download_repo(repo_id: str) -> str:
|
||||||
|
cache_path = MFLUX_LORA_CACHE_DIR
|
||||||
|
cache_path.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
|
print(f"Downloading LoRA from HuggingFace: {repo_id}...")
|
||||||
|
download_path = Path(
|
||||||
|
snapshot_download(
|
||||||
|
repo_id=repo_id,
|
||||||
|
allow_patterns=["*.safetensors"],
|
||||||
|
cache_dir=str(cache_path),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
safetensor_files = list(download_path.glob("*.safetensors"))
|
||||||
|
if not safetensor_files:
|
||||||
|
raise FileNotFoundError(f"No .safetensors file found in HuggingFace repo: {repo_id}")
|
||||||
|
|
||||||
|
if len(safetensor_files) > 1:
|
||||||
|
file_names = [f.name for f in safetensor_files]
|
||||||
|
raise ValueError(
|
||||||
|
f"Multiple .safetensors files found in '{repo_id}': {file_names}. "
|
||||||
|
f"Please specify which file to use with the collection format: '{repo_id}:<filename>.safetensors'"
|
||||||
|
)
|
||||||
|
|
||||||
|
return str(safetensor_files[0])
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _load_collection_from_cache(repo_id: str, filename: str) -> str:
|
||||||
|
cache_path = MFLUX_LORA_CACHE_DIR
|
||||||
|
cache_path.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
|
# Check mflux cache first
|
||||||
|
cached_file_path = cache_path / filename
|
||||||
|
if cached_file_path.exists() and cached_file_path.is_file():
|
||||||
|
try:
|
||||||
|
with open(cached_file_path, "rb") as f:
|
||||||
|
f.read(1)
|
||||||
|
return str(cached_file_path)
|
||||||
|
except (OSError, IOError):
|
||||||
|
# File corrupted, fall through to HF cache
|
||||||
|
pass
|
||||||
|
|
||||||
|
# Load from HF cache
|
||||||
|
download_path = Path(
|
||||||
|
snapshot_download(
|
||||||
|
repo_id=repo_id,
|
||||||
|
allow_patterns=[f"*{filename}*"],
|
||||||
|
cache_dir=str(cache_path),
|
||||||
|
local_files_only=True,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
return LoraResolution._find_and_link_file(download_path, filename, cache_path)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _download_collection(repo_id: str, filename: str) -> str:
|
||||||
|
cache_path = MFLUX_LORA_CACHE_DIR
|
||||||
|
cache_path.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
|
print(f"Downloading LoRA '{filename}' from {repo_id}...")
|
||||||
|
download_path = Path(
|
||||||
|
snapshot_download(
|
||||||
|
repo_id=repo_id,
|
||||||
|
allow_patterns=[f"*{filename}*"],
|
||||||
|
cache_dir=str(cache_path),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
return LoraResolution._find_and_link_file(download_path, filename, cache_path)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _find_and_link_file(download_path: Path, filename: str, cache_path: Path) -> str:
|
||||||
|
found_files = list(download_path.glob(f"**/*{filename}*"))
|
||||||
|
for file in found_files:
|
||||||
|
if file.is_file() and file.suffix in [".safetensors", ".bin"]:
|
||||||
|
if not filename.endswith(file.suffix):
|
||||||
|
target_name = f"{filename}{file.suffix}"
|
||||||
|
else:
|
||||||
|
target_name = filename
|
||||||
|
|
||||||
|
target_path = cache_path / target_name
|
||||||
|
if not target_path.exists():
|
||||||
|
try:
|
||||||
|
target_path.symlink_to(file)
|
||||||
|
except (OSError, AttributeError):
|
||||||
|
shutil.copy2(file, target_path)
|
||||||
|
|
||||||
|
return str(target_path)
|
||||||
|
|
||||||
|
raise FileNotFoundError(f"Could not find LoRA file '{filename}' in downloaded files")
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def get_registry() -> dict[str, Path]:
|
||||||
|
return LoraResolution._registry.copy()
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def discover_files(library_paths: list[Path]) -> dict[str, Path]:
|
||||||
|
lora_files = {}
|
||||||
|
for library_path in reversed(library_paths):
|
||||||
|
if not library_path.exists() or not library_path.is_dir():
|
||||||
|
continue
|
||||||
|
for safetensor_path in library_path.rglob("*.safetensors"):
|
||||||
|
basename = safetensor_path.stem
|
||||||
|
if basename.isdigit() and safetensor_path.parent.name == "transformer":
|
||||||
|
continue
|
||||||
|
lora_files[basename] = safetensor_path.resolve()
|
||||||
|
return lora_files
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _initialize_registry() -> None:
|
||||||
|
library_path_env = os.environ.get("LORA_LIBRARY_PATH")
|
||||||
|
if library_path_env:
|
||||||
|
library_paths = [Path(p.strip()) for p in library_path_env.split(":") if p.strip()]
|
||||||
|
LoraResolution._registry = LoraResolution.discover_files(library_paths)
|
||||||
|
|
||||||
|
|
||||||
|
LoraResolution._initialize_registry()
|
||||||
157
src/mflux/models/common/resolution/path_resolution.py
Normal file
157
src/mflux/models/common/resolution/path_resolution.py
Normal file
@ -0,0 +1,157 @@
|
|||||||
|
import logging
|
||||||
|
import os
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
from huggingface_hub import snapshot_download
|
||||||
|
from huggingface_hub.constants import HF_HUB_CACHE
|
||||||
|
|
||||||
|
from mflux.models.common.resolution.actions import PathAction, Rule
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class PathResolution:
|
||||||
|
RULES = frozenset(
|
||||||
|
{
|
||||||
|
Rule(priority=0, name="none", check="is_none", action=PathAction.LOCAL),
|
||||||
|
Rule(priority=1, name="local", check="exists_locally", action=PathAction.LOCAL),
|
||||||
|
Rule(priority=2, name="hf_cached", check="is_hf_cached", action=PathAction.HUGGINGFACE_CACHED),
|
||||||
|
Rule(priority=3, name="hf_download", check="is_hf_format", action=PathAction.HUGGINGFACE),
|
||||||
|
Rule(priority=4, name="error", check="always", action=PathAction.ERROR),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def resolve(path: str | None, patterns: list[str] | None = None) -> Path | None:
|
||||||
|
if patterns is None:
|
||||||
|
patterns = ["*.safetensors"]
|
||||||
|
|
||||||
|
for rule in sorted(PathResolution.RULES, key=lambda r: r.priority):
|
||||||
|
if PathResolution._check(rule.check, path, patterns):
|
||||||
|
logger.debug(f"Path resolution: '{path}' → rule '{rule.name}' ({rule.action.value})")
|
||||||
|
return PathResolution._execute(rule.action, path, patterns)
|
||||||
|
|
||||||
|
raise ValueError(f"No rule matched for path: {path}")
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _is_hf_format(path: str | None) -> bool:
|
||||||
|
return path is not None and "/" in path and path.count("/") == 1 and not path.startswith(("./", "../", "~/"))
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _check(check: str, path: str | None, patterns: list[str]) -> bool:
|
||||||
|
if check == "is_none":
|
||||||
|
return path is None
|
||||||
|
if check == "exists_locally":
|
||||||
|
if path is None:
|
||||||
|
return False
|
||||||
|
local_path = Path(path).expanduser()
|
||||||
|
if not local_path.exists():
|
||||||
|
return False
|
||||||
|
# Warn if directory exists but contains no matching files
|
||||||
|
if local_path.is_dir():
|
||||||
|
has_matching_files = any(list(local_path.glob(p)) for p in patterns)
|
||||||
|
if not has_matching_files:
|
||||||
|
print(
|
||||||
|
f"⚠️ Directory '{path}' exists but contains no files matching {patterns}. "
|
||||||
|
f"Model loading may fail."
|
||||||
|
)
|
||||||
|
return True
|
||||||
|
if check == "is_hf_cached":
|
||||||
|
if not PathResolution._is_hf_format(path):
|
||||||
|
return False
|
||||||
|
# Check if we have a complete cached snapshot
|
||||||
|
return PathResolution._find_complete_cached_snapshot(path, patterns) is not None
|
||||||
|
if check == "is_hf_format":
|
||||||
|
return PathResolution._is_hf_format(path)
|
||||||
|
if check == "always":
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _execute(action: PathAction, path: str | None, patterns: list[str]) -> Path | None:
|
||||||
|
if action == PathAction.LOCAL:
|
||||||
|
return Path(path).expanduser() if path else None
|
||||||
|
if action == PathAction.HUGGINGFACE_CACHED:
|
||||||
|
# Find the best complete cached snapshot
|
||||||
|
cached_path = PathResolution._find_complete_cached_snapshot(path, patterns)
|
||||||
|
if cached_path:
|
||||||
|
return cached_path
|
||||||
|
# Fallback to standard snapshot_download (shouldn't happen if _check passed)
|
||||||
|
return Path(snapshot_download(repo_id=path, allow_patterns=patterns, local_files_only=True))
|
||||||
|
if action == PathAction.HUGGINGFACE:
|
||||||
|
print(f"Downloading model from HuggingFace: {path}...")
|
||||||
|
return Path(snapshot_download(repo_id=path, allow_patterns=patterns))
|
||||||
|
if action == PathAction.ERROR:
|
||||||
|
raise FileNotFoundError(
|
||||||
|
f"Model not found: '{path}'. "
|
||||||
|
f"If local path, make sure it exists. "
|
||||||
|
f"If HuggingFace repo, use 'org/model' format."
|
||||||
|
)
|
||||||
|
return None
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _find_complete_cached_snapshot(repo_id: str, patterns: list[str]) -> Path | None:
|
||||||
|
# Build the cache directory path for this repo
|
||||||
|
# HuggingFace cache structure: {cache_dir}/models--{org}--{model}/snapshots/{revision}/
|
||||||
|
repo_cache_name = f"models--{repo_id.replace('/', '--')}"
|
||||||
|
repo_cache_dir = Path(HF_HUB_CACHE) / repo_cache_name / "snapshots"
|
||||||
|
|
||||||
|
if not repo_cache_dir.exists():
|
||||||
|
return None
|
||||||
|
|
||||||
|
# Extract subdirectories that need safetensors files (e.g., "vae/*.safetensors" → "vae")
|
||||||
|
required_subdirs = PathResolution._get_required_subdirs_with_safetensors(patterns)
|
||||||
|
|
||||||
|
# Check each snapshot for completeness, prefer more recent ones
|
||||||
|
snapshots = sorted(repo_cache_dir.iterdir(), key=lambda p: p.stat().st_mtime, reverse=True)
|
||||||
|
|
||||||
|
for snapshot_path in snapshots:
|
||||||
|
if not snapshot_path.is_dir():
|
||||||
|
continue
|
||||||
|
if PathResolution._is_snapshot_complete(snapshot_path, required_subdirs):
|
||||||
|
logger.debug(f"Found complete cached snapshot: {snapshot_path}")
|
||||||
|
return snapshot_path
|
||||||
|
|
||||||
|
return None
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _get_required_subdirs_with_safetensors(patterns: list[str]) -> set[str]:
|
||||||
|
subdirs = set()
|
||||||
|
for pattern in patterns:
|
||||||
|
# Only care about safetensors patterns
|
||||||
|
if "*.safetensors" not in pattern:
|
||||||
|
continue
|
||||||
|
# Handle patterns like "vae/*.safetensors"
|
||||||
|
if "/" in pattern:
|
||||||
|
subdir = pattern.split("/")[0]
|
||||||
|
# Only add if it's a real subdir name (not a glob pattern itself)
|
||||||
|
if "*" not in subdir:
|
||||||
|
subdirs.add(subdir)
|
||||||
|
return subdirs
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _is_snapshot_complete(snapshot_path: Path, required_subdirs: set[str]) -> bool:
|
||||||
|
if not required_subdirs:
|
||||||
|
# No specific subdirs required, just check for any safetensors
|
||||||
|
return any(snapshot_path.glob("**/*.safetensors"))
|
||||||
|
|
||||||
|
for subdir in required_subdirs:
|
||||||
|
subdir_path = snapshot_path / subdir
|
||||||
|
if not subdir_path.exists():
|
||||||
|
return False
|
||||||
|
# Check if subdir has at least one safetensors file (following symlinks)
|
||||||
|
has_safetensors = False
|
||||||
|
for f in subdir_path.iterdir():
|
||||||
|
if f.name.endswith(".safetensors"):
|
||||||
|
# Verify the symlink target exists (handles broken symlinks)
|
||||||
|
if f.is_symlink():
|
||||||
|
if os.path.exists(f):
|
||||||
|
has_safetensors = True
|
||||||
|
break
|
||||||
|
else:
|
||||||
|
has_safetensors = True
|
||||||
|
break
|
||||||
|
if not has_safetensors:
|
||||||
|
return False
|
||||||
|
|
||||||
|
return True
|
||||||
@ -0,0 +1,52 @@
|
|||||||
|
import logging
|
||||||
|
|
||||||
|
from mflux.models.common.resolution.actions import QuantizationAction, Rule
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class QuantizationResolution:
|
||||||
|
RULES = frozenset(
|
||||||
|
{
|
||||||
|
Rule(priority=0, name="none", check="none_none", action=QuantizationAction.NONE),
|
||||||
|
Rule(priority=1, name="on_the_fly", check="none_any", action=QuantizationAction.REQUESTED),
|
||||||
|
Rule(priority=2, name="pre_quantized", check="any_none", action=QuantizationAction.STORED),
|
||||||
|
Rule(priority=3, name="conflict", check="any_any", action=QuantizationAction.STORED),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def resolve(stored: int | None, requested: int | None) -> tuple[int | None, str | None]:
|
||||||
|
for rule in sorted(QuantizationResolution.RULES, key=lambda r: r.priority):
|
||||||
|
if QuantizationResolution._check(rule.check, stored, requested):
|
||||||
|
logger.debug(
|
||||||
|
f"Quantization resolution: stored={stored}, requested={requested} "
|
||||||
|
f"→ rule '{rule.name}' ({rule.action.value})"
|
||||||
|
)
|
||||||
|
return QuantizationResolution._execute(rule, stored, requested)
|
||||||
|
|
||||||
|
raise ValueError(f"Unexpected quantization state: stored={stored}, requested={requested}")
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _check(check: str, stored: int | None, requested: int | None) -> bool:
|
||||||
|
if check == "none_none":
|
||||||
|
return stored is None and requested is None
|
||||||
|
if check == "none_any":
|
||||||
|
return stored is None and requested is not None
|
||||||
|
if check == "any_none":
|
||||||
|
return stored is not None and requested is None
|
||||||
|
if check == "any_any":
|
||||||
|
return stored is not None and requested is not None
|
||||||
|
return False
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _execute(rule: Rule, stored: int | None, requested: int | None) -> tuple[int | None, str | None]:
|
||||||
|
if rule.action == QuantizationAction.NONE:
|
||||||
|
return None, None
|
||||||
|
if rule.action == QuantizationAction.REQUESTED:
|
||||||
|
return requested, None
|
||||||
|
if rule.action == QuantizationAction.STORED:
|
||||||
|
warn = rule.name == "conflict" and stored != requested
|
||||||
|
warning = f"Model is pre-quantized at {stored}-bit. Ignoring -q {requested} flag." if warn else None
|
||||||
|
return stored, warning
|
||||||
|
return None, None
|
||||||
@ -4,23 +4,12 @@ import mlx.core as mx
|
|||||||
|
|
||||||
|
|
||||||
class BaseScheduler(ABC):
|
class BaseScheduler(ABC):
|
||||||
"""
|
|
||||||
Abstract base class for all schedulers.
|
|
||||||
"""
|
|
||||||
|
|
||||||
@property
|
@property
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def sigmas(self) -> mx.array:
|
def sigmas(self) -> mx.array: ...
|
||||||
"""
|
|
||||||
The sigma schedule for the diffusion process.
|
|
||||||
"""
|
|
||||||
...
|
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def step(self, model_output: mx.array, timestep: int, sample: mx.array, **kwargs) -> mx.array: ...
|
def step(self, noise: mx.array, timestep: int, latents: mx.array, **kwargs) -> mx.array: ...
|
||||||
|
|
||||||
def scale_model_input(self, latents: mx.array, t: int) -> 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
|
return latents
|
||||||
|
|||||||
@ -4,15 +4,15 @@ from typing import TYPE_CHECKING
|
|||||||
import mlx.core as mx
|
import mlx.core as mx
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from mflux.config.runtime_config import RuntimeConfig
|
from mflux.models.common.config.config import Config
|
||||||
|
|
||||||
from mflux.models.common.schedulers.base_scheduler import BaseScheduler
|
from mflux.models.common.schedulers.base_scheduler import BaseScheduler
|
||||||
|
|
||||||
|
|
||||||
class FlowMatchEulerDiscreteScheduler(BaseScheduler):
|
class FlowMatchEulerDiscreteScheduler(BaseScheduler):
|
||||||
def __init__(self, runtime_config: "RuntimeConfig"):
|
def __init__(self, config: "Config"):
|
||||||
self.runtime_config = runtime_config
|
self.config = config
|
||||||
self.model_config = runtime_config.model_config
|
self.model_config = config.model_config
|
||||||
self.num_train_timesteps = 1000
|
self.num_train_timesteps = 1000
|
||||||
self.shift_terminal = 0.02
|
self.shift_terminal = 0.02
|
||||||
self.base_shift = 0.5
|
self.base_shift = 0.5
|
||||||
@ -30,8 +30,8 @@ class FlowMatchEulerDiscreteScheduler(BaseScheduler):
|
|||||||
return self._timesteps
|
return self._timesteps
|
||||||
|
|
||||||
def _compute_mu(self) -> float:
|
def _compute_mu(self) -> float:
|
||||||
h_patches = self.runtime_config.height // 16
|
h_patches = self.config.height // 16
|
||||||
w_patches = self.runtime_config.width // 16
|
w_patches = self.config.width // 16
|
||||||
seq_len = h_patches * w_patches
|
seq_len = h_patches * w_patches
|
||||||
m = (self.max_shift - self.base_shift) / (self.max_image_seq_len - self.base_image_seq_len)
|
m = (self.max_shift - self.base_shift) / (self.max_image_seq_len - self.base_image_seq_len)
|
||||||
b = self.base_shift - m * self.base_image_seq_len
|
b = self.base_shift - m * self.base_image_seq_len
|
||||||
@ -49,7 +49,7 @@ class FlowMatchEulerDiscreteScheduler(BaseScheduler):
|
|||||||
return stretched
|
return stretched
|
||||||
|
|
||||||
def _compute_timesteps_and_sigmas(self) -> tuple[mx.array, mx.array]:
|
def _compute_timesteps_and_sigmas(self) -> tuple[mx.array, mx.array]:
|
||||||
num_steps = self.runtime_config.num_inference_steps
|
num_steps = self.config.num_inference_steps
|
||||||
sigma_min = 1.0 / self.num_train_timesteps
|
sigma_min = 1.0 / self.num_train_timesteps
|
||||||
sigma_max = 1.0
|
sigma_max = 1.0
|
||||||
timesteps_linear = [
|
timesteps_linear = [
|
||||||
@ -66,10 +66,9 @@ class FlowMatchEulerDiscreteScheduler(BaseScheduler):
|
|||||||
timesteps_arr = mx.array(timesteps, dtype=mx.float32)
|
timesteps_arr = mx.array(timesteps, dtype=mx.float32)
|
||||||
return sigmas_arr, timesteps_arr
|
return sigmas_arr, timesteps_arr
|
||||||
|
|
||||||
def step(self, model_output: mx.array, timestep: int, sample: mx.array, **kwargs) -> mx.array:
|
def step(self, noise: mx.array, timestep: int, latents: mx.array, **kwargs) -> mx.array:
|
||||||
dt = self._sigmas[timestep + 1] - self._sigmas[timestep]
|
dt = self._sigmas[timestep + 1] - self._sigmas[timestep]
|
||||||
prev_sample = sample + dt * model_output
|
return latents + dt * noise
|
||||||
return prev_sample
|
|
||||||
|
|
||||||
def scale_model_input(self, latents: mx.array, t: int) -> mx.array:
|
def scale_model_input(self, latents: mx.array, t: int) -> mx.array:
|
||||||
return latents
|
return latents
|
||||||
|
|||||||
@ -3,14 +3,14 @@ from typing import TYPE_CHECKING
|
|||||||
import mlx.core as mx
|
import mlx.core as mx
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from mflux.config.runtime_config import RuntimeConfig
|
from mflux.models.common.config.config import Config
|
||||||
|
|
||||||
from mflux.models.common.schedulers.base_scheduler import BaseScheduler
|
from mflux.models.common.schedulers.base_scheduler import BaseScheduler
|
||||||
|
|
||||||
|
|
||||||
class LinearScheduler(BaseScheduler):
|
class LinearScheduler(BaseScheduler):
|
||||||
def __init__(self, runtime_config: "RuntimeConfig"):
|
def __init__(self, config: "Config"):
|
||||||
self.runtime_config = runtime_config
|
self.config = config
|
||||||
self._sigmas = self._get_sigmas()
|
self._sigmas = self._get_sigmas()
|
||||||
self._timesteps = self._get_timesteps()
|
self._timesteps = self._get_timesteps()
|
||||||
|
|
||||||
@ -23,11 +23,11 @@ class LinearScheduler(BaseScheduler):
|
|||||||
return self._timesteps
|
return self._timesteps
|
||||||
|
|
||||||
def _get_sigmas(self) -> mx.array:
|
def _get_sigmas(self) -> mx.array:
|
||||||
model_config = self.runtime_config.model_config
|
model_config = self.config.model_config
|
||||||
sigmas = mx.linspace(
|
sigmas = mx.linspace(
|
||||||
1.0,
|
1.0,
|
||||||
1.0 / self.runtime_config.num_inference_steps,
|
1.0 / self.config.num_inference_steps,
|
||||||
self.runtime_config.num_inference_steps,
|
self.config.num_inference_steps,
|
||||||
)
|
)
|
||||||
sigmas = mx.array(sigmas).astype(mx.float32)
|
sigmas = mx.array(sigmas).astype(mx.float32)
|
||||||
sigmas = mx.concatenate([sigmas, mx.zeros(1)])
|
sigmas = mx.concatenate([sigmas, mx.zeros(1)])
|
||||||
@ -36,7 +36,7 @@ class LinearScheduler(BaseScheduler):
|
|||||||
x1 = 256
|
x1 = 256
|
||||||
m = (1.15 - y1) / (4096 - x1)
|
m = (1.15 - y1) / (4096 - x1)
|
||||||
b = y1 - m * x1
|
b = y1 - m * x1
|
||||||
mu = m * self.runtime_config.width * self.runtime_config.height / 256 + b
|
mu = m * self.config.width * self.config.height / 256 + b
|
||||||
mu = mx.array(mu)
|
mu = mx.array(mu)
|
||||||
shifted_sigmas = mx.exp(mu) / (mx.exp(mu) + (1 / sigmas - 1))
|
shifted_sigmas = mx.exp(mu) / (mx.exp(mu) + (1 / sigmas - 1))
|
||||||
shifted_sigmas[-1] = 0
|
shifted_sigmas[-1] = 0
|
||||||
@ -45,11 +45,11 @@ class LinearScheduler(BaseScheduler):
|
|||||||
return sigmas
|
return sigmas
|
||||||
|
|
||||||
def _get_timesteps(self) -> mx.array:
|
def _get_timesteps(self) -> mx.array:
|
||||||
num_steps = self.runtime_config.num_inference_steps
|
num_steps = self.config.num_inference_steps
|
||||||
timesteps = mx.arange(num_steps, dtype=mx.float32)
|
timesteps = mx.arange(num_steps, dtype=mx.float32)
|
||||||
|
|
||||||
return timesteps
|
return timesteps
|
||||||
|
|
||||||
def step(self, model_output: mx.array, timestep: int, sample: mx.array, **kwargs) -> mx.array:
|
def step(self, noise: mx.array, timestep: int, latents: mx.array, **kwargs) -> mx.array:
|
||||||
dt = self._sigmas[timestep + 1] - self._sigmas[timestep]
|
dt = self._sigmas[timestep + 1] - self._sigmas[timestep]
|
||||||
return sample + model_output * dt
|
return latents + noise * dt
|
||||||
|
|||||||
17
src/mflux/models/common/tokenizer/__init__.py
Normal file
17
src/mflux/models/common/tokenizer/__init__.py
Normal file
@ -0,0 +1,17 @@
|
|||||||
|
from mflux.models.common.tokenizer.tokenizer import (
|
||||||
|
BaseTokenizer,
|
||||||
|
LanguageTokenizer,
|
||||||
|
Tokenizer,
|
||||||
|
VisionLanguageTokenizer,
|
||||||
|
)
|
||||||
|
from mflux.models.common.tokenizer.tokenizer_loader import TokenizerLoader
|
||||||
|
from mflux.models.common.tokenizer.tokenizer_output import TokenizerOutput
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"Tokenizer",
|
||||||
|
"BaseTokenizer",
|
||||||
|
"LanguageTokenizer",
|
||||||
|
"VisionLanguageTokenizer",
|
||||||
|
"TokenizerLoader",
|
||||||
|
"TokenizerOutput",
|
||||||
|
]
|
||||||
188
src/mflux/models/common/tokenizer/tokenizer.py
Normal file
188
src/mflux/models/common/tokenizer/tokenizer.py
Normal file
@ -0,0 +1,188 @@
|
|||||||
|
from abc import ABC, abstractmethod
|
||||||
|
from typing import Protocol, runtime_checkable
|
||||||
|
|
||||||
|
import mlx.core as mx
|
||||||
|
import numpy as np
|
||||||
|
from PIL import Image
|
||||||
|
from transformers import PreTrainedTokenizer
|
||||||
|
|
||||||
|
from mflux.models.common.tokenizer.tokenizer_output import TokenizerOutput
|
||||||
|
|
||||||
|
|
||||||
|
@runtime_checkable
|
||||||
|
class Tokenizer(Protocol):
|
||||||
|
tokenizer: PreTrainedTokenizer
|
||||||
|
|
||||||
|
def tokenize(
|
||||||
|
self,
|
||||||
|
prompt: str | list[str],
|
||||||
|
images: list[Image.Image] | None = None,
|
||||||
|
max_length: int | None = None,
|
||||||
|
**kwargs,
|
||||||
|
) -> TokenizerOutput: ...
|
||||||
|
|
||||||
|
|
||||||
|
class BaseTokenizer(ABC):
|
||||||
|
def __init__(self, tokenizer: PreTrainedTokenizer, max_length: int = 512):
|
||||||
|
self.tokenizer = tokenizer
|
||||||
|
self.max_length = max_length
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def tokenize(
|
||||||
|
self,
|
||||||
|
prompt: str | list[str],
|
||||||
|
images: list[Image.Image] | None = None,
|
||||||
|
max_length: int | None = None,
|
||||||
|
**kwargs,
|
||||||
|
) -> TokenizerOutput: ...
|
||||||
|
|
||||||
|
|
||||||
|
class LanguageTokenizer(BaseTokenizer):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
tokenizer: PreTrainedTokenizer,
|
||||||
|
max_length: int = 512,
|
||||||
|
padding: str = "max_length",
|
||||||
|
return_attention_mask: bool = True,
|
||||||
|
template: str | None = None,
|
||||||
|
use_chat_template: bool = False,
|
||||||
|
chat_template_kwargs: dict | None = None,
|
||||||
|
add_special_tokens: bool = True,
|
||||||
|
):
|
||||||
|
super().__init__(tokenizer, max_length)
|
||||||
|
self.padding = padding
|
||||||
|
self.return_attention_mask = return_attention_mask
|
||||||
|
self.template = template
|
||||||
|
self.use_chat_template = use_chat_template
|
||||||
|
self.chat_template_kwargs = chat_template_kwargs or {}
|
||||||
|
self.add_special_tokens = add_special_tokens
|
||||||
|
|
||||||
|
def tokenize(
|
||||||
|
self,
|
||||||
|
prompt: str | list[str],
|
||||||
|
images: list[Image.Image] | None = None,
|
||||||
|
max_length: int | None = None,
|
||||||
|
**kwargs,
|
||||||
|
) -> TokenizerOutput:
|
||||||
|
max_length = max_length or self.max_length
|
||||||
|
|
||||||
|
if isinstance(prompt, str):
|
||||||
|
prompts = [prompt]
|
||||||
|
else:
|
||||||
|
prompts = list(prompt)
|
||||||
|
|
||||||
|
prompts = [p if p is not None else "" for p in prompts]
|
||||||
|
if all(p == "" for p in prompts):
|
||||||
|
batch_size = len(prompts)
|
||||||
|
input_ids = mx.array(np.empty((batch_size, 0), dtype=np.int32))
|
||||||
|
attention_mask = mx.array(np.empty((batch_size, 0), dtype=np.int32))
|
||||||
|
return TokenizerOutput(input_ids=input_ids, attention_mask=attention_mask)
|
||||||
|
|
||||||
|
if self.template or self.use_chat_template:
|
||||||
|
formatted_prompts = []
|
||||||
|
for p in prompts:
|
||||||
|
if self.template:
|
||||||
|
formatted = self.template.format(p)
|
||||||
|
elif self.use_chat_template:
|
||||||
|
formatted = self.tokenizer.apply_chat_template(
|
||||||
|
[{"role": "user", "content": p}],
|
||||||
|
tokenize=False,
|
||||||
|
add_generation_prompt=True,
|
||||||
|
**self.chat_template_kwargs,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
formatted = p
|
||||||
|
formatted_prompts.append(formatted)
|
||||||
|
prompts = formatted_prompts
|
||||||
|
|
||||||
|
tokens = self.tokenizer(
|
||||||
|
prompts,
|
||||||
|
padding=self.padding,
|
||||||
|
max_length=max_length,
|
||||||
|
truncation=True,
|
||||||
|
add_special_tokens=self.add_special_tokens,
|
||||||
|
return_length=False,
|
||||||
|
return_overflowing_tokens=False,
|
||||||
|
return_tensors="np",
|
||||||
|
)
|
||||||
|
|
||||||
|
input_ids = mx.array(tokens["input_ids"])
|
||||||
|
if self.return_attention_mask:
|
||||||
|
attention_mask = mx.array(tokens["attention_mask"])
|
||||||
|
else:
|
||||||
|
attention_mask = mx.ones_like(input_ids)
|
||||||
|
|
||||||
|
return TokenizerOutput(
|
||||||
|
input_ids=input_ids,
|
||||||
|
attention_mask=attention_mask,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class VisionLanguageTokenizer(BaseTokenizer):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
tokenizer: PreTrainedTokenizer,
|
||||||
|
processor,
|
||||||
|
max_length: int = 1024,
|
||||||
|
template: str | None = None,
|
||||||
|
image_token: str = "<|image_pad|>",
|
||||||
|
):
|
||||||
|
super().__init__(tokenizer, max_length)
|
||||||
|
self.processor = processor
|
||||||
|
self.template = template
|
||||||
|
self.image_token = image_token
|
||||||
|
|
||||||
|
def tokenize(
|
||||||
|
self,
|
||||||
|
prompt: str | list[str],
|
||||||
|
images: list[Image.Image] | None = None,
|
||||||
|
max_length: int | None = None,
|
||||||
|
**kwargs,
|
||||||
|
) -> TokenizerOutput:
|
||||||
|
max_length = max_length or self.max_length
|
||||||
|
|
||||||
|
if isinstance(prompt, str):
|
||||||
|
prompt = [prompt]
|
||||||
|
|
||||||
|
if self.template and images:
|
||||||
|
img_prompt = ""
|
||||||
|
for i in range(len(images)):
|
||||||
|
img_prompt += f"Picture {i + 1}: <|vision_start|>{self.image_token}<|vision_end|>"
|
||||||
|
formatted_text = self.template.format(img_prompt + prompt[0])
|
||||||
|
elif self.template:
|
||||||
|
formatted_text = self.template.format(prompt[0])
|
||||||
|
else:
|
||||||
|
formatted_text = prompt[0]
|
||||||
|
|
||||||
|
pixel_values = None
|
||||||
|
image_grid_thw = None
|
||||||
|
|
||||||
|
if images:
|
||||||
|
model_inputs = self.processor(
|
||||||
|
text=[formatted_text],
|
||||||
|
images=images,
|
||||||
|
padding=True,
|
||||||
|
return_tensors=None,
|
||||||
|
)
|
||||||
|
|
||||||
|
input_ids = model_inputs["input_ids"]
|
||||||
|
attention_mask = model_inputs["attention_mask"]
|
||||||
|
pixel_values = mx.array(model_inputs["pixel_values"])
|
||||||
|
image_grid_thw = mx.array(model_inputs["image_grid_thw"])
|
||||||
|
else:
|
||||||
|
tokens = self.tokenizer(
|
||||||
|
[formatted_text],
|
||||||
|
max_length=max_length,
|
||||||
|
padding=True,
|
||||||
|
truncation=True,
|
||||||
|
return_tensors="np",
|
||||||
|
)
|
||||||
|
input_ids = mx.array(tokens["input_ids"])
|
||||||
|
attention_mask = mx.array(tokens["attention_mask"])
|
||||||
|
|
||||||
|
return TokenizerOutput(
|
||||||
|
input_ids=input_ids,
|
||||||
|
attention_mask=attention_mask,
|
||||||
|
pixel_values=pixel_values,
|
||||||
|
image_grid_thw=image_grid_thw,
|
||||||
|
)
|
||||||
144
src/mflux/models/common/tokenizer/tokenizer_loader.py
Normal file
144
src/mflux/models/common/tokenizer/tokenizer_loader.py
Normal file
@ -0,0 +1,144 @@
|
|||||||
|
from pathlib import Path
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
import transformers
|
||||||
|
from huggingface_hub import snapshot_download
|
||||||
|
|
||||||
|
from mflux.models.common.tokenizer.tokenizer import (
|
||||||
|
BaseTokenizer,
|
||||||
|
LanguageTokenizer,
|
||||||
|
VisionLanguageTokenizer,
|
||||||
|
)
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from mflux.models.common.weights.loading.weight_definition import TokenizerDefinition
|
||||||
|
|
||||||
|
|
||||||
|
class TokenizerLoader:
|
||||||
|
@staticmethod
|
||||||
|
def load(
|
||||||
|
definition: "TokenizerDefinition",
|
||||||
|
model_path: str,
|
||||||
|
) -> BaseTokenizer:
|
||||||
|
tokenizer_path = TokenizerLoader._resolve_path(
|
||||||
|
model_path=model_path,
|
||||||
|
hf_subdir=definition.hf_subdir,
|
||||||
|
fallback_subdirs=definition.fallback_subdirs,
|
||||||
|
download_patterns=definition.download_patterns,
|
||||||
|
)
|
||||||
|
|
||||||
|
raw_tokenizer = TokenizerLoader._load_raw_tokenizer(
|
||||||
|
tokenizer_path=tokenizer_path,
|
||||||
|
tokenizer_class=definition.tokenizer_class,
|
||||||
|
)
|
||||||
|
|
||||||
|
return TokenizerLoader._create_tokenizer(
|
||||||
|
raw_tokenizer=raw_tokenizer,
|
||||||
|
definition=definition,
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def load_all(
|
||||||
|
definitions: list["TokenizerDefinition"],
|
||||||
|
model_path: str,
|
||||||
|
max_length_overrides: dict[str, int] | None = None,
|
||||||
|
) -> dict[str, BaseTokenizer]:
|
||||||
|
max_length_overrides = max_length_overrides or {}
|
||||||
|
result = {}
|
||||||
|
for d in definitions:
|
||||||
|
tokenizer = TokenizerLoader.load(d, model_path)
|
||||||
|
if d.name in max_length_overrides:
|
||||||
|
tokenizer.max_length = max_length_overrides[d.name]
|
||||||
|
result[d.name] = tokenizer
|
||||||
|
return result
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _resolve_path(
|
||||||
|
model_path: str,
|
||||||
|
hf_subdir: str,
|
||||||
|
fallback_subdirs: list[str] | None,
|
||||||
|
download_patterns: list[str] | None,
|
||||||
|
) -> Path:
|
||||||
|
expanded = Path(model_path).expanduser()
|
||||||
|
if expanded.exists():
|
||||||
|
root_path = expanded
|
||||||
|
elif "/" in model_path and model_path.count("/") == 1 and not model_path.startswith(("./", "../")):
|
||||||
|
patterns = download_patterns or [f"{hf_subdir}/**"]
|
||||||
|
root_path = Path(
|
||||||
|
snapshot_download(
|
||||||
|
repo_id=model_path,
|
||||||
|
allow_patterns=patterns,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
raise FileNotFoundError(
|
||||||
|
f"Model not found: '{model_path}'. "
|
||||||
|
f"If local path, make sure it exists. "
|
||||||
|
f"If HuggingFace repo, use 'org/model' format."
|
||||||
|
)
|
||||||
|
|
||||||
|
tokenizer_path = root_path / hf_subdir
|
||||||
|
if tokenizer_path.exists():
|
||||||
|
return tokenizer_path
|
||||||
|
|
||||||
|
if fallback_subdirs:
|
||||||
|
for subdir in fallback_subdirs:
|
||||||
|
if subdir == ".":
|
||||||
|
if TokenizerLoader._has_tokenizer_files(root_path):
|
||||||
|
return root_path
|
||||||
|
else:
|
||||||
|
fallback_path = root_path / subdir
|
||||||
|
if fallback_path.exists():
|
||||||
|
return fallback_path
|
||||||
|
|
||||||
|
return tokenizer_path
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _has_tokenizer_files(path: Path) -> bool:
|
||||||
|
tokenizer_indicators = ["vocab.json", "tokenizer.json", "tokenizer_config.json"]
|
||||||
|
return any((path / f).exists() for f in tokenizer_indicators)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _load_raw_tokenizer(
|
||||||
|
tokenizer_path: Path,
|
||||||
|
tokenizer_class: str,
|
||||||
|
):
|
||||||
|
if hasattr(transformers, tokenizer_class):
|
||||||
|
cls = getattr(transformers, tokenizer_class)
|
||||||
|
else:
|
||||||
|
raise ValueError(f"Unknown tokenizer class: {tokenizer_class}")
|
||||||
|
|
||||||
|
return cls.from_pretrained(
|
||||||
|
pretrained_model_name_or_path=str(tokenizer_path),
|
||||||
|
local_files_only=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _create_tokenizer(
|
||||||
|
raw_tokenizer,
|
||||||
|
definition: "TokenizerDefinition",
|
||||||
|
) -> BaseTokenizer:
|
||||||
|
encoder_class = definition.encoder_class
|
||||||
|
|
||||||
|
if encoder_class is VisionLanguageTokenizer:
|
||||||
|
if definition.processor_class is None:
|
||||||
|
raise ValueError("VisionLanguageTokenizer requires processor_class in definition")
|
||||||
|
processor = definition.processor_class(tokenizer=raw_tokenizer)
|
||||||
|
return VisionLanguageTokenizer(
|
||||||
|
tokenizer=raw_tokenizer,
|
||||||
|
processor=processor,
|
||||||
|
max_length=definition.max_length,
|
||||||
|
template=definition.template,
|
||||||
|
image_token=definition.image_token,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
# Default to LanguageTokenizer for all text-only cases
|
||||||
|
return LanguageTokenizer(
|
||||||
|
tokenizer=raw_tokenizer,
|
||||||
|
max_length=definition.max_length,
|
||||||
|
padding=definition.padding,
|
||||||
|
template=definition.template,
|
||||||
|
use_chat_template=definition.use_chat_template,
|
||||||
|
chat_template_kwargs=definition.chat_template_kwargs or {},
|
||||||
|
add_special_tokens=definition.add_special_tokens,
|
||||||
|
)
|
||||||
11
src/mflux/models/common/tokenizer/tokenizer_output.py
Normal file
11
src/mflux/models/common/tokenizer/tokenizer_output.py
Normal file
@ -0,0 +1,11 @@
|
|||||||
|
from dataclasses import dataclass
|
||||||
|
|
||||||
|
import mlx.core as mx
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class TokenizerOutput:
|
||||||
|
input_ids: mx.array
|
||||||
|
attention_mask: mx.array
|
||||||
|
pixel_values: mx.array | None = None
|
||||||
|
image_grid_thw: mx.array | None = None
|
||||||
@ -1,5 +1,14 @@
|
|||||||
"""Common weight mapping utilities."""
|
from mflux.models.common.weights.loading.loaded_weights import LoadedWeights, MetaData
|
||||||
|
from mflux.models.common.weights.loading.weight_applier import WeightApplier
|
||||||
|
from mflux.models.common.weights.loading.weight_definition import ComponentDefinition
|
||||||
|
from mflux.models.common.weights.loading.weight_loader import WeightLoader
|
||||||
|
from mflux.models.common.weights.saving.model_saver import ModelSaver
|
||||||
|
|
||||||
from mflux.models.common.weights.model_saver import ModelSaver
|
__all__ = [
|
||||||
|
"ComponentDefinition",
|
||||||
__all__ = ["ModelSaver"]
|
"LoadedWeights",
|
||||||
|
"MetaData",
|
||||||
|
"ModelSaver",
|
||||||
|
"WeightApplier",
|
||||||
|
"WeightLoader",
|
||||||
|
]
|
||||||
|
|||||||
40
src/mflux/models/common/weights/loading/loaded_weights.py
Normal file
40
src/mflux/models/common/weights/loading/loaded_weights.py
Normal file
@ -0,0 +1,40 @@
|
|||||||
|
from dataclasses import dataclass
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class MetaData:
|
||||||
|
quantization_level: int | None = None
|
||||||
|
mflux_version: str | None = None
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class LoadedWeights:
|
||||||
|
components: dict[str, dict]
|
||||||
|
meta_data: MetaData
|
||||||
|
|
||||||
|
def __getattr__(self, name: str) -> dict | None:
|
||||||
|
if name in ("components", "meta_data"):
|
||||||
|
return object.__getattribute__(self, name)
|
||||||
|
return self.components.get(name)
|
||||||
|
|
||||||
|
def num_transformer_blocks(self, component_name: str = "transformer") -> int:
|
||||||
|
transformer = self.components.get(component_name)
|
||||||
|
if transformer is None:
|
||||||
|
for comp in self.components.values():
|
||||||
|
if isinstance(comp, dict) and "transformer_blocks" in comp:
|
||||||
|
transformer = comp
|
||||||
|
break
|
||||||
|
if transformer and "transformer_blocks" in transformer:
|
||||||
|
return len(transformer["transformer_blocks"])
|
||||||
|
return 0
|
||||||
|
|
||||||
|
def num_single_transformer_blocks(self, component_name: str = "transformer") -> int:
|
||||||
|
transformer = self.components.get(component_name)
|
||||||
|
if transformer is None:
|
||||||
|
for comp in self.components.values():
|
||||||
|
if isinstance(comp, dict) and "single_transformer_blocks" in comp:
|
||||||
|
transformer = comp
|
||||||
|
break
|
||||||
|
if transformer and "single_transformer_blocks" in transformer:
|
||||||
|
return len(transformer["single_transformer_blocks"])
|
||||||
|
return 0
|
||||||
103
src/mflux/models/common/weights/loading/weight_applier.py
Normal file
103
src/mflux/models/common/weights/loading/weight_applier.py
Normal file
@ -0,0 +1,103 @@
|
|||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
import mlx.nn as nn
|
||||||
|
|
||||||
|
from mflux.models.common.resolution.quantization_resolution import QuantizationResolution
|
||||||
|
from mflux.models.common.weights.loading.loaded_weights import LoadedWeights
|
||||||
|
from mflux.models.common.weights.loading.weight_definition import ComponentDefinition
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from mflux.models.common.weights.loading.weight_definition import WeightDefinitionType
|
||||||
|
|
||||||
|
|
||||||
|
class WeightApplier:
|
||||||
|
@staticmethod
|
||||||
|
def apply_and_quantize_single(
|
||||||
|
weights: LoadedWeights,
|
||||||
|
model: nn.Module,
|
||||||
|
component: ComponentDefinition,
|
||||||
|
quantize_arg: int | None,
|
||||||
|
quantization_predicate=None,
|
||||||
|
) -> int | None:
|
||||||
|
stored_q = weights.meta_data.quantization_level
|
||||||
|
component_weights = weights.components.get(component.name)
|
||||||
|
|
||||||
|
if component_weights is None:
|
||||||
|
raise ValueError(f"No weights found for component: {component.name}")
|
||||||
|
|
||||||
|
if quantization_predicate is None:
|
||||||
|
|
||||||
|
def quantization_predicate(path, module):
|
||||||
|
return hasattr(module, "to_quantized")
|
||||||
|
|
||||||
|
bits, warning = QuantizationResolution.resolve(stored=stored_q, requested=quantize_arg)
|
||||||
|
|
||||||
|
if warning:
|
||||||
|
print(f"⚠️ {warning}")
|
||||||
|
|
||||||
|
if bits is None:
|
||||||
|
model.update(component_weights, strict=False)
|
||||||
|
elif stored_q is None:
|
||||||
|
model.update(component_weights, strict=False)
|
||||||
|
if not component.skip_quantization:
|
||||||
|
nn.quantize(model, class_predicate=quantization_predicate, bits=bits)
|
||||||
|
else:
|
||||||
|
if not component.skip_quantization:
|
||||||
|
nn.quantize(model, class_predicate=quantization_predicate, bits=bits)
|
||||||
|
model.update(component_weights, strict=False)
|
||||||
|
|
||||||
|
return bits
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def apply_and_quantize(
|
||||||
|
weights: LoadedWeights,
|
||||||
|
models: dict[str, nn.Module],
|
||||||
|
quantize_arg: int | None,
|
||||||
|
weight_definition: "WeightDefinitionType",
|
||||||
|
) -> int | None:
|
||||||
|
stored_q = weights.meta_data.quantization_level
|
||||||
|
components = {c.name: c for c in weight_definition.get_components()}
|
||||||
|
|
||||||
|
bits, warning = QuantizationResolution.resolve(stored=stored_q, requested=quantize_arg)
|
||||||
|
|
||||||
|
if warning:
|
||||||
|
print(f"⚠️ {warning}")
|
||||||
|
|
||||||
|
if bits is None:
|
||||||
|
WeightApplier._set_weights(weights, models, components)
|
||||||
|
elif stored_q is None:
|
||||||
|
WeightApplier._set_weights(weights, models, components)
|
||||||
|
WeightApplier._quantize(models, bits, components, weight_definition)
|
||||||
|
else:
|
||||||
|
WeightApplier._quantize(models, bits, components, weight_definition)
|
||||||
|
WeightApplier._set_weights(weights, models, components)
|
||||||
|
|
||||||
|
return bits
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _set_weights(
|
||||||
|
weights: LoadedWeights,
|
||||||
|
models: dict[str, nn.Module],
|
||||||
|
components: dict | None = None,
|
||||||
|
) -> None:
|
||||||
|
for name, model in models.items():
|
||||||
|
component_weights = weights.components.get(name)
|
||||||
|
if component_weights is not None:
|
||||||
|
if components is not None:
|
||||||
|
component = components.get(name)
|
||||||
|
if component is not None and component.weight_subkey is not None:
|
||||||
|
component_weights = component_weights.get(component.weight_subkey, component_weights)
|
||||||
|
model.update(component_weights, strict=False)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _quantize(
|
||||||
|
models: dict[str, nn.Module],
|
||||||
|
bits: int,
|
||||||
|
components: dict,
|
||||||
|
weight_definition: "WeightDefinitionType",
|
||||||
|
) -> None:
|
||||||
|
for name, model in models.items():
|
||||||
|
component = components.get(name)
|
||||||
|
if component and component.skip_quantization:
|
||||||
|
continue
|
||||||
|
nn.quantize(model, class_predicate=weight_definition.quantization_predicate, bits=bits)
|
||||||
59
src/mflux/models/common/weights/loading/weight_definition.py
Normal file
59
src/mflux/models/common/weights/loading/weight_definition.py
Normal file
@ -0,0 +1,59 @@
|
|||||||
|
from dataclasses import dataclass, field
|
||||||
|
from typing import TYPE_CHECKING, Callable, List, TypeAlias
|
||||||
|
|
||||||
|
import mlx.core as mx
|
||||||
|
|
||||||
|
from mflux.models.common.weights.mapping.weight_mapping import WeightTarget
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from mflux.models.common.tokenizer.tokenizer import BaseTokenizer
|
||||||
|
from mflux.models.depth_pro.weights.depth_pro_weight_definition import DepthProWeightDefinition
|
||||||
|
from mflux.models.fibo.weights.fibo_weight_definition import FIBOWeightDefinition
|
||||||
|
from mflux.models.fibo_vlm.weights.fibo_vlm_weight_definition import FIBOVLMWeightDefinition
|
||||||
|
from mflux.models.flux.weights.flux_weight_definition import FluxWeightDefinition
|
||||||
|
from mflux.models.qwen.weights.qwen_weight_definition import QwenWeightDefinition
|
||||||
|
from mflux.models.z_image.weights.z_image_weight_definition import ZImageWeightDefinition
|
||||||
|
|
||||||
|
WeightDefinitionType: TypeAlias = type[
|
||||||
|
FluxWeightDefinition
|
||||||
|
| FIBOWeightDefinition
|
||||||
|
| FIBOVLMWeightDefinition
|
||||||
|
| QwenWeightDefinition
|
||||||
|
| ZImageWeightDefinition
|
||||||
|
| DepthProWeightDefinition
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class ComponentDefinition:
|
||||||
|
name: str
|
||||||
|
hf_subdir: str
|
||||||
|
mapping_getter: Callable[[], List[WeightTarget]] | None = None
|
||||||
|
model_attr: str | None = None
|
||||||
|
num_blocks: int | None = None
|
||||||
|
num_layers: int | None = None
|
||||||
|
loading_mode: str = "mlx_native"
|
||||||
|
precision: mx.Dtype | None = None
|
||||||
|
skip_quantization: bool = False
|
||||||
|
bulk_transform: Callable[[mx.array], mx.array] | None = None
|
||||||
|
weight_subkey: str | None = None
|
||||||
|
download_url: str | None = None
|
||||||
|
weight_prefix_filters: List[str] | None = None
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class TokenizerDefinition:
|
||||||
|
name: str
|
||||||
|
hf_subdir: str
|
||||||
|
tokenizer_class: str = "AutoTokenizer"
|
||||||
|
fallback_subdirs: List[str] | None = None
|
||||||
|
download_patterns: List[str] | None = None
|
||||||
|
encoder_class: type["BaseTokenizer"] | None = None
|
||||||
|
max_length: int = 512
|
||||||
|
padding: str = "max_length"
|
||||||
|
template: str | None = None
|
||||||
|
use_chat_template: bool = False
|
||||||
|
chat_template_kwargs: dict | None = field(default_factory=dict)
|
||||||
|
add_special_tokens: bool = True
|
||||||
|
processor_class: type | None = None
|
||||||
|
image_token: str = "<|image_pad|>"
|
||||||
317
src/mflux/models/common/weights/loading/weight_loader.py
Normal file
317
src/mflux/models/common/weights/loading/weight_loader.py
Normal file
@ -0,0 +1,317 @@
|
|||||||
|
import json
|
||||||
|
import logging
|
||||||
|
import urllib.error
|
||||||
|
import urllib.request
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
import mlx.core as mx
|
||||||
|
import torch
|
||||||
|
from huggingface_hub import snapshot_download
|
||||||
|
from mlx.utils import tree_unflatten
|
||||||
|
from safetensors.torch import load_file as torch_load_file
|
||||||
|
|
||||||
|
from mflux.cli.defaults.defaults import MFLUX_CACHE_DIR
|
||||||
|
from mflux.models.common.resolution.path_resolution import PathResolution
|
||||||
|
from mflux.models.common.weights.loading.loaded_weights import LoadedWeights, MetaData
|
||||||
|
from mflux.models.common.weights.loading.weight_definition import ComponentDefinition
|
||||||
|
from mflux.models.common.weights.mapping.weight_mapper import WeightMapper
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from mflux.models.common.weights.loading.weight_definition import WeightDefinitionType
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class WeightLoader:
|
||||||
|
@staticmethod
|
||||||
|
def load_single(
|
||||||
|
component: ComponentDefinition,
|
||||||
|
repo_id: str,
|
||||||
|
file_pattern: str = "*.safetensors",
|
||||||
|
) -> LoadedWeights:
|
||||||
|
root_path = Path(snapshot_download(repo_id=repo_id, allow_patterns=[file_pattern, "config.json"]))
|
||||||
|
weights, q_level, version = WeightLoader._load_component(root_path, component)
|
||||||
|
return LoadedWeights(
|
||||||
|
components={component.name: weights},
|
||||||
|
meta_data=MetaData(quantization_level=q_level, mflux_version=version),
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def load(
|
||||||
|
weight_definition: "WeightDefinitionType",
|
||||||
|
model_path: str | None = None,
|
||||||
|
) -> LoadedWeights:
|
||||||
|
root_path = PathResolution.resolve(
|
||||||
|
path=model_path,
|
||||||
|
patterns=weight_definition.get_download_patterns(),
|
||||||
|
)
|
||||||
|
|
||||||
|
# 2. Load each component (with caching for shared sources)
|
||||||
|
components = {}
|
||||||
|
quantization_level = None
|
||||||
|
mflux_version = None
|
||||||
|
raw_weights_cache: dict[tuple, dict] = {} # Cache by (path, loading_mode)
|
||||||
|
|
||||||
|
for component in weight_definition.get_components():
|
||||||
|
weights, q_level, version = WeightLoader._load_component(root_path, component, raw_weights_cache)
|
||||||
|
components[component.name] = weights
|
||||||
|
|
||||||
|
# Track metadata from first component that has it
|
||||||
|
if quantization_level is None and q_level is not None:
|
||||||
|
quantization_level = q_level
|
||||||
|
mflux_version = version
|
||||||
|
|
||||||
|
return LoadedWeights(
|
||||||
|
components=components,
|
||||||
|
meta_data=MetaData(
|
||||||
|
quantization_level=quantization_level,
|
||||||
|
mflux_version=mflux_version,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _load_component(
|
||||||
|
root_path: Path | None,
|
||||||
|
component: ComponentDefinition,
|
||||||
|
raw_weights_cache: dict[tuple, dict] | None = None,
|
||||||
|
) -> tuple[dict, int | None, str | None]:
|
||||||
|
# Handle direct URL downloads (e.g., Apple CDN for DepthPro)
|
||||||
|
if component.download_url is not None:
|
||||||
|
file_path = WeightLoader._download_from_url(component.download_url, component.name)
|
||||||
|
raw_weights = WeightLoader._load_weights_file(file_path, component.loading_mode)
|
||||||
|
else:
|
||||||
|
if root_path is None:
|
||||||
|
raise ValueError(f"No root_path and no download_url for component: {component.name}")
|
||||||
|
component_path = root_path / component.hf_subdir
|
||||||
|
|
||||||
|
# Try mflux saved format first
|
||||||
|
weights, q_level, version = WeightLoader._try_load_mflux_format(component_path)
|
||||||
|
if weights is not None:
|
||||||
|
return weights, q_level, version
|
||||||
|
|
||||||
|
# Check cache for shared loading (e.g., FIBO VLM decoder + visual from same source)
|
||||||
|
cache_key = (str(component_path), component.loading_mode)
|
||||||
|
if raw_weights_cache is not None and cache_key in raw_weights_cache:
|
||||||
|
raw_weights = raw_weights_cache[cache_key]
|
||||||
|
else:
|
||||||
|
# Fall back to HuggingFace format with mapping
|
||||||
|
raw_weights = WeightLoader._load_safetensors(component_path, component.loading_mode)
|
||||||
|
# Cache for potential reuse by other components
|
||||||
|
if raw_weights_cache is not None:
|
||||||
|
raw_weights_cache[cache_key] = raw_weights
|
||||||
|
|
||||||
|
# Apply prefix filtering if specified (e.g., filter "model.language_model" vs "model.visual")
|
||||||
|
if component.weight_prefix_filters is not None:
|
||||||
|
raw_weights = {
|
||||||
|
k: v
|
||||||
|
for k, v in raw_weights.items()
|
||||||
|
if any(k.startswith(prefix) for prefix in component.weight_prefix_filters)
|
||||||
|
}
|
||||||
|
|
||||||
|
# Apply precision conversion if specified
|
||||||
|
if component.precision is not None:
|
||||||
|
raw_weights = WeightLoader._convert_precision(raw_weights, component.precision)
|
||||||
|
|
||||||
|
# Passthrough mode: apply bulk transform and unflatten (no key mapping)
|
||||||
|
if component.mapping_getter is None:
|
||||||
|
if component.bulk_transform is not None:
|
||||||
|
raw_weights = {k: component.bulk_transform(v) for k, v in raw_weights.items()}
|
||||||
|
return tree_unflatten(list(raw_weights.items())), None, None
|
||||||
|
|
||||||
|
# Standard mode: apply declarative weight mapping
|
||||||
|
mapped_weights = WeightMapper.apply_mapping(
|
||||||
|
hf_weights=raw_weights,
|
||||||
|
mapping=component.mapping_getter(),
|
||||||
|
num_blocks=component.num_blocks,
|
||||||
|
num_layers=component.num_layers,
|
||||||
|
)
|
||||||
|
return mapped_weights, None, None
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _try_load_mflux_format(path: Path) -> tuple[dict | None, int | None, str | None]:
|
||||||
|
if not path.exists():
|
||||||
|
return None, None, None
|
||||||
|
|
||||||
|
shard_files = sorted(f for f in path.glob("*.safetensors") if not f.name.startswith("._"))
|
||||||
|
if not shard_files:
|
||||||
|
return None, None, None
|
||||||
|
|
||||||
|
# Check metadata on first file
|
||||||
|
data = mx.load(str(shard_files[0]), return_metadata=True)
|
||||||
|
if len(data) <= 1:
|
||||||
|
return None, None, None
|
||||||
|
|
||||||
|
quantization_level_str = data[1].get("quantization_level")
|
||||||
|
mflux_version = data[1].get("mflux_version")
|
||||||
|
|
||||||
|
# If no mflux metadata, this isn't our format
|
||||||
|
if quantization_level_str is None and mflux_version is None:
|
||||||
|
return None, None, None
|
||||||
|
|
||||||
|
# Convert quantization level from string to int
|
||||||
|
quantization_level = int(quantization_level_str) if quantization_level_str is not None else None
|
||||||
|
|
||||||
|
# Load all shards
|
||||||
|
all_weights: dict[str, mx.array] = {}
|
||||||
|
for shard in shard_files:
|
||||||
|
shard_data = mx.load(str(shard), return_metadata=True)
|
||||||
|
all_weights.update(dict(shard_data[0].items()))
|
||||||
|
|
||||||
|
unflattened = tree_unflatten(list(all_weights.items()))
|
||||||
|
return unflattened, quantization_level, mflux_version
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _download_from_url(url: str, component_name: str) -> Path:
|
||||||
|
cache_dir = MFLUX_CACHE_DIR / component_name
|
||||||
|
cache_dir.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
|
# Extract filename from URL
|
||||||
|
filename = url.split("/")[-1]
|
||||||
|
file_path = cache_dir / filename
|
||||||
|
|
||||||
|
if not file_path.exists():
|
||||||
|
logger.info(f"Downloading {component_name} weights from {url}...")
|
||||||
|
try:
|
||||||
|
urllib.request.urlretrieve(url, file_path)
|
||||||
|
logger.info(f"Downloaded to {file_path}")
|
||||||
|
except (urllib.error.URLError, urllib.error.HTTPError) as e:
|
||||||
|
logger.error(f"Failed to download: {e}")
|
||||||
|
logger.info(f"Please manually download from: {url}")
|
||||||
|
raise FileNotFoundError(f"Model file not found at {file_path}") from e
|
||||||
|
|
||||||
|
return file_path
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _load_weights_file(file_path: Path, loading_mode: str) -> dict[str, mx.array]:
|
||||||
|
if loading_mode == "torch_checkpoint":
|
||||||
|
return WeightLoader._load_torch_checkpoint(file_path)
|
||||||
|
elif loading_mode in ("mlx_native", "single"):
|
||||||
|
data = mx.load(str(file_path), return_metadata=True)
|
||||||
|
return dict(data[0].items())
|
||||||
|
else:
|
||||||
|
raise ValueError(f"Unsupported loading mode for single file: {loading_mode}")
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _load_torch_checkpoint(file_path: Path) -> dict[str, mx.array]:
|
||||||
|
pt_weights = torch.load(file_path, map_location="cpu", weights_only=False)
|
||||||
|
return {k: mx.array(v.numpy()) for k, v in pt_weights.items() if isinstance(v, torch.Tensor)}
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _load_safetensors(path: Path, loading_mode: str) -> dict[str, mx.array]:
|
||||||
|
if loading_mode == "mlx_native":
|
||||||
|
return WeightLoader._load_mlx_native(path)
|
||||||
|
elif loading_mode == "torch_convert":
|
||||||
|
return WeightLoader._load_torch_convert(path)
|
||||||
|
elif loading_mode == "multi_json":
|
||||||
|
return WeightLoader._load_multi_json(path)
|
||||||
|
elif loading_mode == "torch_bfloat16":
|
||||||
|
return WeightLoader._load_torch_bfloat16(path)
|
||||||
|
elif loading_mode == "single":
|
||||||
|
return WeightLoader._load_single(path)
|
||||||
|
elif loading_mode == "multi_glob":
|
||||||
|
return WeightLoader._load_multi_glob(path)
|
||||||
|
else:
|
||||||
|
raise ValueError(f"Unknown loading mode: {loading_mode}")
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _load_mlx_native(path: Path) -> dict[str, mx.array]:
|
||||||
|
shard_files = sorted(f for f in path.glob("*.safetensors") if not f.name.startswith("._"))
|
||||||
|
if not shard_files:
|
||||||
|
raise FileNotFoundError(f"No safetensors files found in {path}")
|
||||||
|
|
||||||
|
all_weights: dict[str, mx.array] = {}
|
||||||
|
for shard in shard_files:
|
||||||
|
weights = mx.load(str(shard))
|
||||||
|
all_weights.update(weights)
|
||||||
|
|
||||||
|
return all_weights
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _load_torch_convert(path: Path) -> dict[str, mx.array]:
|
||||||
|
shard_files = sorted(f for f in path.glob("*.safetensors") if not f.name.startswith("._"))
|
||||||
|
if not shard_files:
|
||||||
|
raise FileNotFoundError(f"No safetensors files found in {path}")
|
||||||
|
|
||||||
|
all_weights: dict[str, mx.array] = {}
|
||||||
|
for shard in shard_files:
|
||||||
|
torch_weights = torch_load_file(str(shard))
|
||||||
|
for key, tensor in torch_weights.items():
|
||||||
|
if tensor.dtype == torch.bfloat16:
|
||||||
|
tensor = tensor.to(torch.float16)
|
||||||
|
all_weights[key] = mx.array(tensor.numpy())
|
||||||
|
|
||||||
|
return all_weights
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _load_multi_json(path: Path) -> dict[str, mx.array]:
|
||||||
|
index_path = path / "model.safetensors.index.json"
|
||||||
|
with open(index_path) as f:
|
||||||
|
index = json.load(f)
|
||||||
|
|
||||||
|
# Group weights by file
|
||||||
|
files_to_load: dict[str, list[str]] = {}
|
||||||
|
for param_name, file_name in index["weight_map"].items():
|
||||||
|
if file_name not in files_to_load:
|
||||||
|
files_to_load[file_name] = []
|
||||||
|
files_to_load[file_name].append(param_name)
|
||||||
|
|
||||||
|
all_weights: dict[str, mx.array] = {}
|
||||||
|
for file_name, param_names in files_to_load.items():
|
||||||
|
file_path = path / file_name
|
||||||
|
|
||||||
|
# Use mx.load which handles bfloat16 natively
|
||||||
|
file_weights = mx.load(str(file_path))
|
||||||
|
|
||||||
|
for param_name in param_names:
|
||||||
|
if param_name in file_weights:
|
||||||
|
all_weights[param_name] = file_weights[param_name]
|
||||||
|
|
||||||
|
return all_weights
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _load_torch_bfloat16(path: Path) -> dict[str, mx.array]:
|
||||||
|
index_path = path / "model.safetensors.index.json"
|
||||||
|
with open(index_path) as f:
|
||||||
|
index = json.load(f)
|
||||||
|
|
||||||
|
weight_files = sorted(set(index["weight_map"].values()))
|
||||||
|
|
||||||
|
all_weights: dict[str, mx.array] = {}
|
||||||
|
for wf in weight_files:
|
||||||
|
file_path = path / wf
|
||||||
|
data = torch_load_file(str(file_path))
|
||||||
|
for k, v in data.items():
|
||||||
|
if v.dtype == torch.bfloat16:
|
||||||
|
v = v.to(torch.float16)
|
||||||
|
np_arr = v.detach().cpu().numpy()
|
||||||
|
all_weights[k] = mx.array(np_arr)
|
||||||
|
|
||||||
|
return all_weights
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _load_single(path: Path) -> dict[str, mx.array]:
|
||||||
|
safetensors_files = [f for f in path.glob("*.safetensors") if not f.name.startswith("._")]
|
||||||
|
if not safetensors_files:
|
||||||
|
raise FileNotFoundError(f"No safetensors files found in {path}")
|
||||||
|
|
||||||
|
weights_file = safetensors_files[0]
|
||||||
|
data = mx.load(str(weights_file), return_metadata=True)
|
||||||
|
return dict(data[0].items())
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _load_multi_glob(path: Path) -> dict[str, mx.array]:
|
||||||
|
shard_files = sorted(f for f in path.glob("*.safetensors") if not f.name.startswith("._"))
|
||||||
|
if not shard_files:
|
||||||
|
raise FileNotFoundError(f"No safetensors files found in {path}")
|
||||||
|
|
||||||
|
all_weights: dict[str, mx.array] = {}
|
||||||
|
for shard in shard_files:
|
||||||
|
data, _ = mx.load(str(shard), return_metadata=True)
|
||||||
|
all_weights.update(dict(data.items()))
|
||||||
|
|
||||||
|
return all_weights
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _convert_precision(weights: dict[str, mx.array], precision: mx.Dtype) -> dict[str, mx.array]:
|
||||||
|
return {k: v.astype(precision) for k, v in weights.items()}
|
||||||
@ -1,6 +1,10 @@
|
|||||||
"""Weight mapping utilities for transforming HuggingFace weights to MLX structure."""
|
|
||||||
|
|
||||||
from mflux.models.common.weights.mapping.weight_mapper import WeightMapper
|
from mflux.models.common.weights.mapping.weight_mapper import WeightMapper
|
||||||
from mflux.models.common.weights.mapping.weight_mapping import WeightMapping, WeightTarget
|
from mflux.models.common.weights.mapping.weight_mapping import WeightMapping, WeightTarget
|
||||||
|
from mflux.models.common.weights.mapping.weight_transforms import WeightTransforms
|
||||||
|
|
||||||
__all__ = ["WeightMapping", "WeightTarget", "WeightMapper"]
|
__all__ = [
|
||||||
|
"WeightMapping",
|
||||||
|
"WeightTarget",
|
||||||
|
"WeightMapper",
|
||||||
|
"WeightTransforms",
|
||||||
|
]
|
||||||
|
|||||||
@ -1,9 +1,3 @@
|
|||||||
"""
|
|
||||||
Weight Mapper - Applies declarative weight mappings to transform HF weights to MLX structure.
|
|
||||||
|
|
||||||
Similar to LoRALoader, but for weight mapping instead of LoRA application.
|
|
||||||
"""
|
|
||||||
|
|
||||||
import re
|
import re
|
||||||
from typing import Callable, Dict, List, Optional
|
from typing import Callable, Dict, List, Optional
|
||||||
|
|
||||||
@ -13,8 +7,6 @@ from mflux.models.common.weights.mapping.weight_mapping import WeightTarget
|
|||||||
|
|
||||||
|
|
||||||
class WeightMapper:
|
class WeightMapper:
|
||||||
"""Maps HuggingFace weights to MLX nested structure using declarative mappings."""
|
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def apply_mapping(
|
def apply_mapping(
|
||||||
hf_weights: Dict[str, mx.array],
|
hf_weights: Dict[str, mx.array],
|
||||||
@ -22,18 +14,6 @@ class WeightMapper:
|
|||||||
num_blocks: Optional[int] = None,
|
num_blocks: Optional[int] = None,
|
||||||
num_layers: Optional[int] = None,
|
num_layers: Optional[int] = None,
|
||||||
) -> Dict:
|
) -> Dict:
|
||||||
"""
|
|
||||||
Apply weight mapping to transform HF weights to MLX structure.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
hf_weights: Raw HuggingFace weights (flat dict with dot-notation keys)
|
|
||||||
mapping: List of WeightTarget mappings
|
|
||||||
num_blocks: Number of transformer blocks (auto-detected if None)
|
|
||||||
num_layers: Number of text encoder layers (auto-detected if None)
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Nested dict structure matching MLX model
|
|
||||||
"""
|
|
||||||
# Auto-detect number of blocks if not provided
|
# Auto-detect number of blocks if not provided
|
||||||
if num_blocks is None:
|
if num_blocks is None:
|
||||||
num_blocks = WeightMapper._detect_num_blocks(hf_weights)
|
num_blocks = WeightMapper._detect_num_blocks(hf_weights)
|
||||||
@ -42,7 +22,7 @@ class WeightMapper:
|
|||||||
if num_layers is None:
|
if num_layers is None:
|
||||||
num_layers = WeightMapper._detect_num_layers(hf_weights)
|
num_layers = WeightMapper._detect_num_layers(hf_weights)
|
||||||
|
|
||||||
# Build flat mapping: HF pattern -> (MLX path, transform)
|
# Build flat mapping: HF pattern -> [(MLX path, transform), ...] (supports one-to-many)
|
||||||
flat_mapping = WeightMapper._build_flat_mapping(mapping, num_blocks, num_layers)
|
flat_mapping = WeightMapper._build_flat_mapping(mapping, num_blocks, num_layers)
|
||||||
|
|
||||||
# Map weights
|
# Map weights
|
||||||
@ -51,18 +31,19 @@ class WeightMapper:
|
|||||||
skipped_count = 0
|
skipped_count = 0
|
||||||
|
|
||||||
for hf_key, hf_tensor in hf_weights.items():
|
for hf_key, hf_tensor in hf_weights.items():
|
||||||
# Try to find matching mapping
|
# Try to find matching mappings (can be multiple targets for one source)
|
||||||
mlx_path, transform = WeightMapper._find_mapping(hf_key, flat_mapping)
|
targets = flat_mapping.get(hf_key, [])
|
||||||
|
|
||||||
if mlx_path:
|
if targets:
|
||||||
# Apply transform if specified
|
for mlx_path, transform in targets:
|
||||||
tensor = hf_tensor
|
# Apply transform if specified
|
||||||
if transform:
|
tensor = hf_tensor
|
||||||
tensor = transform(tensor)
|
if transform:
|
||||||
|
tensor = transform(tensor)
|
||||||
|
|
||||||
# Build nested structure
|
# Build nested structure
|
||||||
WeightMapper._set_nested_value(mapped_weights, mlx_path, tensor)
|
WeightMapper._set_nested_value(mapped_weights, mlx_path, tensor)
|
||||||
mapped_count += 1
|
mapped_count += 1
|
||||||
else:
|
else:
|
||||||
# Weight not in mapping - might be intentionally skipped (e.g., lm_head)
|
# Weight not in mapping - might be intentionally skipped (e.g., lm_head)
|
||||||
# or optional weight (e.g., conv_shortcut) - that's OK
|
# or optional weight (e.g., conv_shortcut) - that's OK
|
||||||
@ -75,7 +56,6 @@ class WeightMapper:
|
|||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _detect_num_blocks(hf_weights: Dict[str, mx.array]) -> int:
|
def _detect_num_blocks(hf_weights: Dict[str, mx.array]) -> int:
|
||||||
"""Detect number of transformer blocks from weight keys."""
|
|
||||||
block_numbers = set()
|
block_numbers = set()
|
||||||
for key in hf_weights.keys():
|
for key in hf_weights.keys():
|
||||||
# Match pattern: transformer_blocks.{number}.something
|
# Match pattern: transformer_blocks.{number}.something
|
||||||
@ -89,7 +69,6 @@ class WeightMapper:
|
|||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _detect_num_layers(hf_weights: Dict[str, mx.array]) -> int:
|
def _detect_num_layers(hf_weights: Dict[str, mx.array]) -> int:
|
||||||
"""Detect number of text encoder layers from weight keys."""
|
|
||||||
layer_numbers = set()
|
layer_numbers = set()
|
||||||
for key in hf_weights.keys():
|
for key in hf_weights.keys():
|
||||||
# Match pattern: model.layers.{number}.something
|
# Match pattern: model.layers.{number}.something
|
||||||
@ -104,89 +83,85 @@ class WeightMapper:
|
|||||||
@staticmethod
|
@staticmethod
|
||||||
def _build_flat_mapping(
|
def _build_flat_mapping(
|
||||||
mapping: List[WeightTarget], num_blocks: int = 0, num_layers: int = 28
|
mapping: List[WeightTarget], num_blocks: int = 0, num_layers: int = 28
|
||||||
) -> Dict[str, tuple[str, Optional[Callable[[mx.array], mx.array]]]]:
|
) -> Dict[str, List[tuple[str, Optional[Callable[[mx.array], mx.array]]]]]:
|
||||||
"""
|
flat: Dict[str, List[tuple[str, Optional[Callable[[mx.array], mx.array]]]]] = {}
|
||||||
Build flat mapping from declarative targets.
|
|
||||||
|
|
||||||
Returns:
|
def add_mapping(hf_key: str, mlx_path: str, transform: Optional[Callable[[mx.array], mx.array]]):
|
||||||
Dict mapping HF pattern -> (MLX path, transform)
|
if hf_key not in flat:
|
||||||
"""
|
flat[hf_key] = []
|
||||||
flat = {}
|
flat[hf_key].append((mlx_path, transform))
|
||||||
|
|
||||||
for target in mapping:
|
for target in mapping:
|
||||||
# Expand placeholders for each pattern
|
# Expand placeholders for each pattern
|
||||||
for hf_pattern in target.hf_patterns:
|
for hf_pattern in target.from_pattern:
|
||||||
# Check which placeholders are present
|
# Check which placeholders are present in BOTH patterns
|
||||||
has_block = "{block}" in hf_pattern or "{block}" in target.mlx_path
|
hf_has_block = "{block}" in hf_pattern
|
||||||
has_i = "{i}" in hf_pattern or "{i}" in target.mlx_path
|
to_has_block = "{block}" in target.to_pattern
|
||||||
has_res = "{res}" in hf_pattern or "{res}" in target.mlx_path
|
has_i = "{i}" in hf_pattern or "{i}" in target.to_pattern
|
||||||
has_layer = "{layer}" in hf_pattern or "{layer}" in target.mlx_path
|
has_res = "{res}" in hf_pattern or "{res}" in target.to_pattern
|
||||||
|
has_layer = "{layer}" in hf_pattern or "{layer}" in target.to_pattern
|
||||||
|
|
||||||
# Handle multiple placeholders together
|
# Handle multiple placeholders together
|
||||||
if has_block and has_res:
|
if (hf_has_block or to_has_block) and has_res:
|
||||||
# Up blocks: expand both {block} and {res}
|
# Up blocks: expand both {block} and {res}
|
||||||
max_blocks = num_blocks if num_blocks > 0 else 4 # Default 4 for up_blocks
|
max_blocks = num_blocks if num_blocks > 0 else 4 # Default 4 for up_blocks
|
||||||
for block_num in range(max_blocks):
|
for block_num in range(max_blocks):
|
||||||
for res in range(3): # 3 resnets per up_block
|
for res in range(3): # 3 resnets per up_block
|
||||||
concrete_hf = hf_pattern.replace("{block}", str(block_num)).replace("{res}", str(res))
|
concrete_hf = hf_pattern.replace("{block}", str(block_num)).replace("{res}", str(res))
|
||||||
concrete_mlx = target.mlx_path.replace("{block}", str(block_num)).replace("{res}", str(res))
|
concrete_mlx = target.to_pattern.replace("{block}", str(block_num)).replace(
|
||||||
flat[concrete_hf] = (concrete_mlx, target.transform)
|
"{res}", str(res)
|
||||||
elif has_block:
|
)
|
||||||
# Expand {block} only (for transformer blocks or visual blocks)
|
add_mapping(concrete_hf, concrete_mlx, target.transform)
|
||||||
# Check if target has max_blocks override
|
elif hf_has_block and to_has_block:
|
||||||
|
# Both have {block} - standard one-to-one expansion
|
||||||
if target.max_blocks is not None:
|
if target.max_blocks is not None:
|
||||||
max_blocks = target.max_blocks
|
max_blocks = target.max_blocks
|
||||||
# Check if this is for visual blocks (32 blocks) or transformer blocks
|
elif "visual.blocks" in hf_pattern or "visual.blocks" in target.to_pattern:
|
||||||
elif "visual.blocks" in hf_pattern or "visual.blocks" in target.mlx_path:
|
|
||||||
max_blocks = 32 # Visual blocks are always 32
|
max_blocks = 32 # Visual blocks are always 32
|
||||||
else:
|
else:
|
||||||
max_blocks = num_blocks if num_blocks > 0 else 4 # Default 4 for up_blocks
|
max_blocks = num_blocks if num_blocks > 0 else 4
|
||||||
for block_num in range(max_blocks):
|
for block_num in range(max_blocks):
|
||||||
concrete_hf = hf_pattern.replace("{block}", str(block_num))
|
concrete_hf = hf_pattern.replace("{block}", str(block_num))
|
||||||
concrete_mlx = target.mlx_path.replace("{block}", str(block_num))
|
concrete_mlx = target.to_pattern.replace("{block}", str(block_num))
|
||||||
flat[concrete_hf] = (concrete_mlx, target.transform)
|
add_mapping(concrete_hf, concrete_mlx, target.transform)
|
||||||
|
elif to_has_block and not hf_has_block:
|
||||||
|
# One-to-many: single HF key maps to multiple MLX targets (e.g., relative_attention_bias)
|
||||||
|
if target.max_blocks is not None:
|
||||||
|
max_blocks = target.max_blocks
|
||||||
|
else:
|
||||||
|
max_blocks = num_blocks if num_blocks > 0 else 24 # Default for T5
|
||||||
|
for block_num in range(max_blocks):
|
||||||
|
concrete_mlx = target.to_pattern.replace("{block}", str(block_num))
|
||||||
|
add_mapping(hf_pattern, concrete_mlx, target.transform)
|
||||||
elif has_layer:
|
elif has_layer:
|
||||||
# Expand {layer} for text encoder layers or visual blocks
|
# Expand {layer} for text encoder layers or visual blocks
|
||||||
max_layers = num_layers if num_layers > 0 else 28 # Default 28 for text encoder
|
max_layers = num_layers if num_layers > 0 else 28 # Default 28 for text encoder
|
||||||
for layer_num in range(max_layers):
|
for layer_num in range(max_layers):
|
||||||
concrete_hf = hf_pattern.replace("{layer}", str(layer_num))
|
concrete_hf = hf_pattern.replace("{layer}", str(layer_num))
|
||||||
concrete_mlx = target.mlx_path.replace("{layer}", str(layer_num))
|
concrete_mlx = target.to_pattern.replace("{layer}", str(layer_num))
|
||||||
flat[concrete_hf] = (concrete_mlx, target.transform)
|
add_mapping(concrete_hf, concrete_mlx, target.transform)
|
||||||
elif has_i:
|
elif has_i:
|
||||||
# Expand {i} only (for mid_block resnets)
|
# Expand {i} only (for mid_block resnets)
|
||||||
for i in range(2): # 2 resnets in mid_block
|
for i in range(2): # 2 resnets in mid_block
|
||||||
concrete_hf = hf_pattern.replace("{i}", str(i))
|
concrete_hf = hf_pattern.replace("{i}", str(i))
|
||||||
concrete_mlx = target.mlx_path.replace("{i}", str(i))
|
concrete_mlx = target.to_pattern.replace("{i}", str(i))
|
||||||
flat[concrete_hf] = (concrete_mlx, target.transform)
|
add_mapping(concrete_hf, concrete_mlx, target.transform)
|
||||||
elif has_res:
|
elif has_res:
|
||||||
# This shouldn't happen for VAE (encoder down_blocks are explicit)
|
# This shouldn't happen for VAE (encoder down_blocks are explicit)
|
||||||
# But handle it just in case
|
# But handle it just in case
|
||||||
if "up_block" in hf_pattern:
|
if "up_block" in hf_pattern:
|
||||||
for res in range(3):
|
for res in range(3):
|
||||||
concrete_hf = hf_pattern.replace("{res}", str(res))
|
concrete_hf = hf_pattern.replace("{res}", str(res))
|
||||||
concrete_mlx = target.mlx_path.replace("{res}", str(res))
|
concrete_mlx = target.to_pattern.replace("{res}", str(res))
|
||||||
flat[concrete_hf] = (concrete_mlx, target.transform)
|
add_mapping(concrete_hf, concrete_mlx, target.transform)
|
||||||
else:
|
else:
|
||||||
# No placeholder, use as-is
|
# No placeholder, use as-is
|
||||||
flat[hf_pattern] = (target.mlx_path, target.transform)
|
add_mapping(hf_pattern, target.to_pattern, target.transform)
|
||||||
|
|
||||||
return flat
|
return flat
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _find_mapping(
|
|
||||||
hf_key: str, flat_mapping: Dict[str, tuple[str, Optional[Callable[[mx.array], mx.array]]]]
|
|
||||||
) -> tuple[Optional[str], Optional[Callable[[mx.array], mx.array]]]:
|
|
||||||
"""Find MLX path and transform for a given HF key."""
|
|
||||||
return flat_mapping.get(hf_key, (None, None))
|
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _set_nested_value(d: Dict, path: str, value: mx.array):
|
def _set_nested_value(d: Dict, path: str, value: mx.array):
|
||||||
"""
|
|
||||||
Set value in nested dict using dot-notation path.
|
|
||||||
|
|
||||||
Creates nested structure as needed.
|
|
||||||
Handles both dict keys and list indices.
|
|
||||||
"""
|
|
||||||
parts = path.split(".")
|
parts = path.split(".")
|
||||||
current = d
|
current = d
|
||||||
i = 0
|
i = 0
|
||||||
|
|||||||
@ -1,9 +1,3 @@
|
|||||||
"""
|
|
||||||
Base classes for declarative weight mapping.
|
|
||||||
|
|
||||||
Similar to LoRA mapping, but for transforming HuggingFace weights to MLX structure.
|
|
||||||
"""
|
|
||||||
|
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import Callable, List, Optional, Protocol
|
from typing import Callable, List, Optional, Protocol
|
||||||
|
|
||||||
@ -12,23 +6,14 @@ import mlx.core as mx
|
|||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class WeightTarget:
|
class WeightTarget:
|
||||||
"""
|
to_pattern: str
|
||||||
Declarative weight mapping target.
|
from_pattern: List[str]
|
||||||
|
transform: Optional[Callable[[mx.array], mx.array]] = None
|
||||||
Maps HuggingFace weight names to MLX nested structure paths.
|
required: bool = True
|
||||||
"""
|
max_blocks: Optional[int] = None
|
||||||
|
|
||||||
mlx_path: str # MLX nested path, e.g., "transformer_blocks.{block}.attn.to_q.weight"
|
|
||||||
hf_patterns: List[str] # HuggingFace naming patterns, e.g., ["transformer_blocks.{block}.attn.to_q.weight"]
|
|
||||||
transform: Optional[Callable[[mx.array], mx.array]] = None # Optional transform (reshape, transpose, etc.)
|
|
||||||
required: bool = True # If False, weight is optional (may not exist in all models)
|
|
||||||
max_blocks: Optional[int] = None # Override num_blocks for this target (e.g., for fixed-size lists)
|
|
||||||
|
|
||||||
|
|
||||||
class WeightMapping(Protocol):
|
class WeightMapping(Protocol):
|
||||||
"""Protocol for weight mapping classes."""
|
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def get_mapping() -> List[WeightTarget]:
|
def get_mapping() -> List[WeightTarget]:
|
||||||
"""Return list of weight mapping targets."""
|
|
||||||
return []
|
return []
|
||||||
|
|||||||
33
src/mflux/models/common/weights/mapping/weight_transforms.py
Normal file
33
src/mflux/models/common/weights/mapping/weight_transforms.py
Normal file
@ -0,0 +1,33 @@
|
|||||||
|
import mlx.core as mx
|
||||||
|
|
||||||
|
|
||||||
|
class WeightTransforms:
|
||||||
|
@staticmethod
|
||||||
|
def reshape_gamma_to_1d(tensor: mx.array) -> mx.array:
|
||||||
|
if len(tensor.shape) > 1:
|
||||||
|
return mx.reshape(tensor, (tensor.shape[0],))
|
||||||
|
return tensor
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def transpose_patch_embed(tensor: mx.array) -> mx.array:
|
||||||
|
if len(tensor.shape) == 5:
|
||||||
|
return tensor.transpose(0, 2, 3, 4, 1)
|
||||||
|
return tensor
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def transpose_conv3d_weight(tensor: mx.array) -> mx.array:
|
||||||
|
if len(tensor.shape) == 5:
|
||||||
|
return tensor.transpose(0, 2, 3, 4, 1)
|
||||||
|
return tensor
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def transpose_conv2d_weight(tensor: mx.array) -> mx.array:
|
||||||
|
if len(tensor.shape) == 4:
|
||||||
|
return tensor.transpose(0, 2, 3, 1)
|
||||||
|
return tensor
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def transpose_conv_transpose2d_weight(tensor: mx.array) -> mx.array:
|
||||||
|
if len(tensor.shape) == 4:
|
||||||
|
return tensor.transpose(1, 2, 3, 0)
|
||||||
|
return tensor
|
||||||
@ -1,114 +0,0 @@
|
|||||||
from pathlib import Path
|
|
||||||
from typing import Any
|
|
||||||
|
|
||||||
import mlx.core as mx
|
|
||||||
from mlx import nn
|
|
||||||
from mlx.utils import tree_flatten
|
|
||||||
from transformers import PreTrainedTokenizer
|
|
||||||
|
|
||||||
from mflux.utils.version_util import VersionUtil
|
|
||||||
|
|
||||||
|
|
||||||
class ModelSaver:
|
|
||||||
@staticmethod
|
|
||||||
def save_model(
|
|
||||||
model: Any,
|
|
||||||
bits: int,
|
|
||||||
base_path: str,
|
|
||||||
tokenizers: list[tuple[str, str]] | None = None,
|
|
||||||
components: list[tuple[str, str]] | None = None,
|
|
||||||
) -> None:
|
|
||||||
# Default tokenizers: try common patterns
|
|
||||||
if tokenizers is None:
|
|
||||||
tokenizers = ModelSaver._detect_tokenizers(model)
|
|
||||||
|
|
||||||
# Default components: try common patterns
|
|
||||||
if components is None:
|
|
||||||
components = ModelSaver._detect_components(model)
|
|
||||||
|
|
||||||
# Save tokenizers
|
|
||||||
for attr_path, subdir in tokenizers:
|
|
||||||
tokenizer = ModelSaver._get_nested_attr(model, attr_path)
|
|
||||||
if tokenizer is not None:
|
|
||||||
ModelSaver._save_tokenizer(base_path, tokenizer, subdir)
|
|
||||||
|
|
||||||
# Save model components
|
|
||||||
for attr_name, subdir in components:
|
|
||||||
component = getattr(model, attr_name, None)
|
|
||||||
if component is not None:
|
|
||||||
ModelSaver._save_weights(base_path, bits, component, subdir)
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _detect_tokenizers(model: Any) -> list[tuple[str, str]]:
|
|
||||||
tokenizers = []
|
|
||||||
if hasattr(model, "clip_tokenizer") and hasattr(model.clip_tokenizer, "tokenizer"):
|
|
||||||
tokenizers.append(("clip_tokenizer.tokenizer", "tokenizer"))
|
|
||||||
if hasattr(model, "t5_tokenizer") and hasattr(model.t5_tokenizer, "tokenizer"):
|
|
||||||
tokenizers.append(("t5_tokenizer.tokenizer", "tokenizer_2"))
|
|
||||||
if hasattr(model, "qwen_tokenizer") and hasattr(model.qwen_tokenizer, "tokenizer"):
|
|
||||||
tokenizers.append(("qwen_tokenizer.tokenizer", "tokenizer"))
|
|
||||||
if hasattr(model, "fibo_tokenizer") and hasattr(model.fibo_tokenizer, "tokenizer"):
|
|
||||||
tokenizers.append(("fibo_tokenizer.tokenizer", "tokenizer"))
|
|
||||||
return tokenizers
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _detect_components(model: Any) -> list[tuple[str, str]]:
|
|
||||||
components = []
|
|
||||||
if hasattr(model, "vae"):
|
|
||||||
components.append(("vae", "vae"))
|
|
||||||
if hasattr(model, "transformer"):
|
|
||||||
components.append(("transformer", "transformer"))
|
|
||||||
if hasattr(model, "clip_text_encoder"):
|
|
||||||
components.append(("clip_text_encoder", "text_encoder"))
|
|
||||||
if hasattr(model, "t5_text_encoder"):
|
|
||||||
components.append(("t5_text_encoder", "text_encoder_2"))
|
|
||||||
if hasattr(model, "text_encoder"):
|
|
||||||
# Only add if we haven't already added clip_text_encoder or t5_text_encoder
|
|
||||||
if not any(c[1] == "text_encoder" for c in components):
|
|
||||||
components.append(("text_encoder", "text_encoder"))
|
|
||||||
return components
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _get_nested_attr(obj: Any, attr_path: str) -> Any:
|
|
||||||
attrs = attr_path.split(".")
|
|
||||||
result = obj
|
|
||||||
for attr in attrs:
|
|
||||||
if not hasattr(result, attr):
|
|
||||||
return None
|
|
||||||
result = getattr(result, attr)
|
|
||||||
return result
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _save_tokenizer(base_path: str, tokenizer: PreTrainedTokenizer, subdir: str) -> None:
|
|
||||||
path = Path(base_path) / subdir
|
|
||||||
path.mkdir(parents=True, exist_ok=True)
|
|
||||||
tokenizer.save_pretrained(path)
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _save_weights(base_path: str, bits: int, model: nn.Module, subdir: str) -> None:
|
|
||||||
path = Path(base_path) / subdir
|
|
||||||
path.mkdir(parents=True, exist_ok=True)
|
|
||||||
weights = ModelSaver._split_weights(base_path, dict(tree_flatten(model.parameters())))
|
|
||||||
for i, weight in enumerate(weights):
|
|
||||||
mx.save_safetensors(
|
|
||||||
str(path / f"{i}.safetensors"),
|
|
||||||
weight,
|
|
||||||
{
|
|
||||||
"quantization_level": str(bits),
|
|
||||||
"mflux_version": VersionUtil.get_mflux_version(),
|
|
||||||
},
|
|
||||||
)
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _split_weights(base_path: str, weights: dict, max_file_size_gb: int = 2) -> list:
|
|
||||||
max_file_size_bytes = max_file_size_gb << 30
|
|
||||||
shards = []
|
|
||||||
shard, shard_size = {}, 0
|
|
||||||
for k, v in weights.items():
|
|
||||||
if shard_size + v.nbytes > max_file_size_bytes:
|
|
||||||
shards.append(shard)
|
|
||||||
shard, shard_size = {}, 0
|
|
||||||
shard[k] = v
|
|
||||||
shard_size += v.nbytes
|
|
||||||
shards.append(shard)
|
|
||||||
return shards
|
|
||||||
96
src/mflux/models/common/weights/saving/model_saver.py
Normal file
96
src/mflux/models/common/weights/saving/model_saver.py
Normal file
@ -0,0 +1,96 @@
|
|||||||
|
import json
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import TYPE_CHECKING, Any
|
||||||
|
|
||||||
|
import mlx.core as mx
|
||||||
|
from mlx import nn
|
||||||
|
from mlx.utils import tree_flatten
|
||||||
|
from tqdm import tqdm
|
||||||
|
from transformers import PreTrainedTokenizer
|
||||||
|
|
||||||
|
from mflux.utils.version_util import VersionUtil
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from mflux.models.common.weights.loading.weight_definition import WeightDefinitionType
|
||||||
|
|
||||||
|
|
||||||
|
class ModelSaver:
|
||||||
|
@staticmethod
|
||||||
|
def save_model(
|
||||||
|
model: Any,
|
||||||
|
bits: int,
|
||||||
|
base_path: str,
|
||||||
|
weight_definition: "WeightDefinitionType",
|
||||||
|
) -> None:
|
||||||
|
# Save tokenizers from model.tokenizers dict
|
||||||
|
tokenizer_defs = weight_definition.get_tokenizers()
|
||||||
|
for t in tokenizer_defs:
|
||||||
|
if hasattr(model, "tokenizers") and t.name in model.tokenizers:
|
||||||
|
tokenizer_wrapper = model.tokenizers[t.name]
|
||||||
|
if hasattr(tokenizer_wrapper, "tokenizer"):
|
||||||
|
ModelSaver._save_tokenizer(base_path, tokenizer_wrapper.tokenizer, t.hf_subdir)
|
||||||
|
|
||||||
|
# Save model components with progress bar
|
||||||
|
components = [(c.model_attr or c.name, c.hf_subdir) for c in weight_definition.get_components()]
|
||||||
|
for attr_name, subdir in tqdm(components, desc="Saving components", unit="component"):
|
||||||
|
component = getattr(model, attr_name, None)
|
||||||
|
if component is not None:
|
||||||
|
ModelSaver._save_weights(base_path, bits, component, subdir)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _save_tokenizer(base_path: str, tokenizer: PreTrainedTokenizer, subdir: str) -> None:
|
||||||
|
path = Path(base_path) / subdir
|
||||||
|
path.mkdir(parents=True, exist_ok=True)
|
||||||
|
tokenizer.save_pretrained(path)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _save_weights(base_path: str, bits: int, model: nn.Module, subdir: str) -> None:
|
||||||
|
path = Path(base_path) / subdir
|
||||||
|
path.mkdir(parents=True, exist_ok=True)
|
||||||
|
weights = dict(tree_flatten(model.parameters()))
|
||||||
|
shards = ModelSaver._split_weights(weights)
|
||||||
|
|
||||||
|
# Build weight_map for index.json (maps each weight key to its shard file)
|
||||||
|
weight_map = {}
|
||||||
|
shard_iter = tqdm(enumerate(shards), total=len(shards), desc=f" {subdir}", unit="shard", leave=False)
|
||||||
|
for i, shard in shard_iter:
|
||||||
|
shard_filename = f"{i}.safetensors"
|
||||||
|
mx.save_safetensors(
|
||||||
|
str(path / shard_filename),
|
||||||
|
shard,
|
||||||
|
{
|
||||||
|
"quantization_level": str(bits),
|
||||||
|
"mflux_version": VersionUtil.get_mflux_version(),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
# Record which file each weight belongs to
|
||||||
|
for key in shard.keys():
|
||||||
|
weight_map[key] = shard_filename
|
||||||
|
|
||||||
|
# Write model.safetensors.index.json for HuggingFace compatibility
|
||||||
|
# This ensures the saved model works even if custom metadata is stripped
|
||||||
|
index_data = {
|
||||||
|
"metadata": {
|
||||||
|
"quantization_level": str(bits),
|
||||||
|
"mflux_version": VersionUtil.get_mflux_version(),
|
||||||
|
},
|
||||||
|
"weight_map": weight_map,
|
||||||
|
}
|
||||||
|
with open(path / "model.safetensors.index.json", "w") as f:
|
||||||
|
json.dump(index_data, f, indent=2)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _split_weights(weights: dict, max_file_size_gb: int = 2) -> list[dict]:
|
||||||
|
max_file_size_bytes = max_file_size_gb << 30
|
||||||
|
shards: list[dict] = []
|
||||||
|
shard: dict = {}
|
||||||
|
shard_size = 0
|
||||||
|
for k, v in weights.items():
|
||||||
|
if shard_size + v.nbytes > max_file_size_bytes:
|
||||||
|
shards.append(shard)
|
||||||
|
shard, shard_size = {}, 0
|
||||||
|
shard[k] = v
|
||||||
|
shard_size += v.nbytes
|
||||||
|
if shard: # Don't append empty shard
|
||||||
|
shards.append(shard)
|
||||||
|
return shards
|
||||||
@ -1,7 +1,7 @@
|
|||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
from mflux.models.depth_pro.depth_pro import DepthPro
|
from mflux.cli.parser.parsers import CommandLineParser
|
||||||
from mflux.ui.cli.parsers import CommandLineParser
|
from mflux.models.depth_pro.model.depth_pro import DepthPro
|
||||||
|
|
||||||
|
|
||||||
def main():
|
def main():
|
||||||
@ -1,25 +1,21 @@
|
|||||||
import mlx.nn as nn
|
from mflux.models.common.weights.loading.weight_applier import WeightApplier
|
||||||
|
from mflux.models.common.weights.loading.weight_loader import WeightLoader
|
||||||
from mflux.models.depth_pro.model.depth_pro_model import DepthProModel
|
from mflux.models.depth_pro.model.depth_pro_model import DepthProModel
|
||||||
from mflux.models.depth_pro.weights.weight_handler_depth_pro import WeightHandlerDepthPro
|
from mflux.models.depth_pro.weights.depth_pro_weight_definition import DepthProWeightDefinition
|
||||||
|
|
||||||
|
|
||||||
class DepthProInitializer:
|
class DepthProInitializer:
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def init(depth_pro_model: DepthProModel, quantize: int | None = None) -> None:
|
def init(model: DepthProModel, quantize: int | None = None) -> None:
|
||||||
# 1. Load the weights
|
# 1. Load weights using unified loader (handles download from Apple CDN)
|
||||||
depth_pro_weights = WeightHandlerDepthPro.load_weights()
|
weights = WeightLoader.load(weight_definition=DepthProWeightDefinition)
|
||||||
WeightHandlerDepthPro.reposition_encoder_weights(depth_pro_weights, "upsample_latent0")
|
|
||||||
WeightHandlerDepthPro.reposition_encoder_weights(depth_pro_weights, "upsample_latent1")
|
|
||||||
WeightHandlerDepthPro.reposition_encoder_weights(depth_pro_weights, "upsample0")
|
|
||||||
WeightHandlerDepthPro.reposition_encoder_weights(depth_pro_weights, "upsample1")
|
|
||||||
WeightHandlerDepthPro.reposition_encoder_weights(depth_pro_weights, "upsample2")
|
|
||||||
WeightHandlerDepthPro.reposition_head_weights(depth_pro_weights)
|
|
||||||
WeightHandlerDepthPro.reshape_transposed_convolution_weights(depth_pro_weights)
|
|
||||||
|
|
||||||
# 2. Assign the weights to the model
|
# 2. Apply weights and quantize using unified applier
|
||||||
depth_pro_model.update(depth_pro_weights.weights, strict=False)
|
WeightApplier.apply_and_quantize(
|
||||||
|
weights=weights,
|
||||||
# 3. Optionally quantize the model
|
quantize_arg=quantize,
|
||||||
if quantize:
|
weight_definition=DepthProWeightDefinition,
|
||||||
nn.quantize(depth_pro_model, bits=quantize)
|
models={
|
||||||
|
"depth_pro": model,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|||||||
@ -1,23 +0,0 @@
|
|||||||
import mlx.core as mx
|
|
||||||
import mlx.nn as nn
|
|
||||||
|
|
||||||
|
|
||||||
class ConvUtils:
|
|
||||||
@staticmethod
|
|
||||||
def apply_conv(x: mx.array, conv_module: nn.Module) -> mx.array:
|
|
||||||
"""Apply a convolution with channel format conversion.
|
|
||||||
|
|
||||||
MLX expects channels-last format (B,H,W,C) for convolutions,
|
|
||||||
but tensors are generally in channels-first format (B,C,H,W).
|
|
||||||
This helper handles the conversion automatically.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
x: Input tensor in channels-first format (B,C,H,W)
|
|
||||||
conv_module: Convolution module to apply
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Output tensor in channels-first format (B,C,H,W)
|
|
||||||
"""
|
|
||||||
x = mx.transpose(x, (0, 2, 3, 1))
|
|
||||||
x = conv_module(x)
|
|
||||||
return mx.transpose(x, (0, 3, 1, 2))
|
|
||||||
@ -1,8 +1,8 @@
|
|||||||
import mlx.core as mx
|
import mlx.core as mx
|
||||||
import mlx.nn as nn
|
import mlx.nn as nn
|
||||||
|
|
||||||
from mflux.models.depth_pro.model.conv_utils import ConvUtils
|
from mflux.models.depth_pro.model.decoder.residual_block import ResidualBlock
|
||||||
from mflux.models.depth_pro.model.residual_block import ResidualBlock
|
from mflux.models.depth_pro.model.depth_pro_util import DepthProUtil
|
||||||
|
|
||||||
|
|
||||||
class FeatureFusionBlock2d(nn.Module):
|
class FeatureFusionBlock2d(nn.Module):
|
||||||
@ -21,6 +21,6 @@ class FeatureFusionBlock2d(nn.Module):
|
|||||||
x = x + res
|
x = x + res
|
||||||
x = self.resnet2(x)
|
x = self.resnet2(x)
|
||||||
if self.use_deconv:
|
if self.use_deconv:
|
||||||
x = ConvUtils.apply_conv(x, self.deconv)
|
x = DepthProUtil.apply_conv(x, self.deconv)
|
||||||
x = ConvUtils.apply_conv(x, self.out_conv)
|
x = DepthProUtil.apply_conv(x, self.out_conv)
|
||||||
return x
|
return x
|
||||||
@ -1,8 +1,8 @@
|
|||||||
import mlx.core as mx
|
import mlx.core as mx
|
||||||
import mlx.nn as nn
|
import mlx.nn as nn
|
||||||
|
|
||||||
from mflux.models.depth_pro.model.conv_utils import ConvUtils
|
from mflux.models.depth_pro.model.decoder.feature_fusion_block_2d import FeatureFusionBlock2d
|
||||||
from mflux.models.depth_pro.model.feature_fusion_block_2d import FeatureFusionBlock2d
|
from mflux.models.depth_pro.model.depth_pro_util import DepthProUtil
|
||||||
|
|
||||||
|
|
||||||
class MultiresConvDecoder(nn.Module):
|
class MultiresConvDecoder(nn.Module):
|
||||||
@ -32,20 +32,20 @@ class MultiresConvDecoder(nn.Module):
|
|||||||
x_global_features: mx.array,
|
x_global_features: mx.array,
|
||||||
) -> mx.array:
|
) -> mx.array:
|
||||||
# Process global features:
|
# Process global features:
|
||||||
features = ConvUtils.apply_conv(x_global_features, self.convs[4])
|
features = DepthProUtil.apply_conv(x_global_features, self.convs[4])
|
||||||
features = self.fusions[4](features)
|
features = self.fusions[4](features)
|
||||||
|
|
||||||
# Process remaining levels with skip connections:
|
# Process remaining levels with skip connections:
|
||||||
x1_skip_features = ConvUtils.apply_conv(x1_features, self.convs[3])
|
x1_skip_features = DepthProUtil.apply_conv(x1_features, self.convs[3])
|
||||||
features = self.fusions[3](features, x1_skip_features)
|
features = self.fusions[3](features, x1_skip_features)
|
||||||
|
|
||||||
x0_skip_features = ConvUtils.apply_conv(x0_features, self.convs[2])
|
x0_skip_features = DepthProUtil.apply_conv(x0_features, self.convs[2])
|
||||||
features = self.fusions[2](features, x0_skip_features)
|
features = self.fusions[2](features, x0_skip_features)
|
||||||
|
|
||||||
x1_skip_latents = ConvUtils.apply_conv(x1_latent, self.convs[1])
|
x1_skip_latents = DepthProUtil.apply_conv(x1_latent, self.convs[1])
|
||||||
features = self.fusions[1](features, x1_skip_latents)
|
features = self.fusions[1](features, x1_skip_latents)
|
||||||
|
|
||||||
x0_skip_latents = ConvUtils.apply_conv(x0_latent, self.convs[0])
|
x0_skip_latents = DepthProUtil.apply_conv(x0_latent, self.convs[0])
|
||||||
features = self.fusions[0](features, x0_skip_latents)
|
features = self.fusions[0](features, x0_skip_latents)
|
||||||
|
|
||||||
return features
|
return features
|
||||||
@ -1,7 +1,7 @@
|
|||||||
import mlx.core as mx
|
import mlx.core as mx
|
||||||
import mlx.nn as nn
|
import mlx.nn as nn
|
||||||
|
|
||||||
from mflux.models.depth_pro.model.conv_utils import ConvUtils
|
from mflux.models.depth_pro.model.depth_pro_util import DepthProUtil
|
||||||
|
|
||||||
|
|
||||||
class ResidualBlock(nn.Module):
|
class ResidualBlock(nn.Module):
|
||||||
@ -30,7 +30,7 @@ class ResidualBlock(nn.Module):
|
|||||||
|
|
||||||
def __call__(self, x: mx.array) -> mx.array:
|
def __call__(self, x: mx.array) -> mx.array:
|
||||||
delta_x = nn.relu(x)
|
delta_x = nn.relu(x)
|
||||||
delta_x = ConvUtils.apply_conv(delta_x, self.residual[1])
|
delta_x = DepthProUtil.apply_conv(delta_x, self.residual[1])
|
||||||
delta_x = nn.relu(delta_x)
|
delta_x = nn.relu(delta_x)
|
||||||
delta_x = ConvUtils.apply_conv(delta_x, self.residual[3])
|
delta_x = DepthProUtil.apply_conv(delta_x, self.residual[3])
|
||||||
return x + delta_x
|
return x + delta_x
|
||||||
@ -1,9 +1,9 @@
|
|||||||
import mlx.core as mx
|
import mlx.core as mx
|
||||||
import mlx.nn as nn
|
import mlx.nn as nn
|
||||||
|
|
||||||
from mflux.models.depth_pro.model.depth_pro_encoder import DepthProEncoder
|
from mflux.models.depth_pro.model.decoder.multires_conv_decoder import MultiresConvDecoder
|
||||||
from mflux.models.depth_pro.model.fov_head import FOVHead
|
from mflux.models.depth_pro.model.encoder.depth_pro_encoder import DepthProEncoder
|
||||||
from mflux.models.depth_pro.model.multires_conv_decoder import MultiresConvDecoder
|
from mflux.models.depth_pro.model.head.fov_head import FOVHead
|
||||||
|
|
||||||
|
|
||||||
class DepthProModel(nn.Module):
|
class DepthProModel(nn.Module):
|
||||||
|
|||||||
@ -1,6 +1,7 @@
|
|||||||
import math
|
import math
|
||||||
|
|
||||||
import mlx.core as mx
|
import mlx.core as mx
|
||||||
|
import mlx.nn as nn
|
||||||
import numpy as np
|
import numpy as np
|
||||||
from PIL import Image
|
from PIL import Image
|
||||||
|
|
||||||
@ -64,3 +65,9 @@ class DepthProUtil:
|
|||||||
final_result_np = result_proc
|
final_result_np = result_proc
|
||||||
|
|
||||||
return mx.array(final_result_np)
|
return mx.array(final_result_np)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def apply_conv(x: mx.array, conv_module: nn.Module) -> mx.array:
|
||||||
|
x = mx.transpose(x, (0, 2, 3, 1))
|
||||||
|
x = conv_module(x)
|
||||||
|
return mx.transpose(x, (0, 3, 1, 2))
|
||||||
|
|||||||
@ -3,9 +3,9 @@ import math
|
|||||||
import mlx.core as mx
|
import mlx.core as mx
|
||||||
import mlx.nn as nn
|
import mlx.nn as nn
|
||||||
|
|
||||||
from mflux.models.depth_pro.model.conv_utils import ConvUtils
|
from mflux.models.depth_pro.model.depth_pro_util import DepthProUtil
|
||||||
from mflux.models.depth_pro.model.dino_v2.dino_vision_transformer import DinoVisionTransformer
|
from mflux.models.depth_pro.model.dino_v2.dino_vision_transformer import DinoVisionTransformer
|
||||||
from mflux.models.depth_pro.model.upsample_block import UpSampleBlock
|
from mflux.models.depth_pro.model.encoder.upsample_block import UpSampleBlock
|
||||||
|
|
||||||
|
|
||||||
class DepthProEncoder(nn.Module):
|
class DepthProEncoder(nn.Module):
|
||||||
@ -56,9 +56,9 @@ class DepthProEncoder(nn.Module):
|
|||||||
# 4. Apply the image encoder model.
|
# 4. Apply the image encoder model.
|
||||||
x_global_features, _, _ = self.image_encoder(x2)
|
x_global_features, _, _ = self.image_encoder(x2)
|
||||||
x_global_features = DepthProEncoder._reshape_feature(embeddings=x_global_features, width=24, height=24)
|
x_global_features = DepthProEncoder._reshape_feature(embeddings=x_global_features, width=24, height=24)
|
||||||
x_global_features = ConvUtils.apply_conv(x_global_features, self.upsample_lowres)
|
x_global_features = DepthProUtil.apply_conv(x_global_features, self.upsample_lowres)
|
||||||
x_global_features = mx.concatenate((x2_features, x_global_features), axis=1)
|
x_global_features = mx.concatenate((x2_features, x_global_features), axis=1)
|
||||||
x_global_features = ConvUtils.apply_conv(x_global_features, self.fuse_lowres)
|
x_global_features = DepthProUtil.apply_conv(x_global_features, self.fuse_lowres)
|
||||||
|
|
||||||
return (
|
return (
|
||||||
x_latent0_features,
|
x_latent0_features,
|
||||||
@ -1,7 +1,7 @@
|
|||||||
import mlx.core as mx
|
import mlx.core as mx
|
||||||
import mlx.nn as nn
|
import mlx.nn as nn
|
||||||
|
|
||||||
from mflux.models.depth_pro.model.conv_utils import ConvUtils
|
from mflux.models.depth_pro.model.depth_pro_util import DepthProUtil
|
||||||
|
|
||||||
|
|
||||||
class UpSampleBlock(nn.Module):
|
class UpSampleBlock(nn.Module):
|
||||||
@ -47,5 +47,5 @@ class UpSampleBlock(nn.Module):
|
|||||||
|
|
||||||
def __call__(self, x: mx.array) -> mx.array:
|
def __call__(self, x: mx.array) -> mx.array:
|
||||||
for layer in self.layers:
|
for layer in self.layers:
|
||||||
x = ConvUtils.apply_conv(x, layer)
|
x = DepthProUtil.apply_conv(x, layer)
|
||||||
return x
|
return x
|
||||||
0
src/mflux/models/depth_pro/model/head/__init__.py
Normal file
0
src/mflux/models/depth_pro/model/head/__init__.py
Normal file
@ -1,7 +1,7 @@
|
|||||||
import mlx.core as mx
|
import mlx.core as mx
|
||||||
import mlx.nn as nn
|
import mlx.nn as nn
|
||||||
|
|
||||||
from mflux.models.depth_pro.model.conv_utils import ConvUtils
|
from mflux.models.depth_pro.model.depth_pro_util import DepthProUtil
|
||||||
|
|
||||||
|
|
||||||
class FOVHead(nn.Module):
|
class FOVHead(nn.Module):
|
||||||
@ -16,10 +16,10 @@ class FOVHead(nn.Module):
|
|||||||
]
|
]
|
||||||
|
|
||||||
def __call__(self, x: mx.array) -> mx.array:
|
def __call__(self, x: mx.array) -> mx.array:
|
||||||
x = ConvUtils.apply_conv(x, self.convs[0])
|
x = DepthProUtil.apply_conv(x, self.convs[0])
|
||||||
x = ConvUtils.apply_conv(x, self.convs[1])
|
x = DepthProUtil.apply_conv(x, self.convs[1])
|
||||||
x = ConvUtils.apply_conv(x, self.convs[2])
|
x = DepthProUtil.apply_conv(x, self.convs[2])
|
||||||
x = nn.relu(x)
|
x = nn.relu(x)
|
||||||
x = ConvUtils.apply_conv(x, self.convs[4])
|
x = DepthProUtil.apply_conv(x, self.convs[4])
|
||||||
x = nn.relu(x)
|
x = nn.relu(x)
|
||||||
return x
|
return x
|
||||||
@ -0,0 +1,34 @@
|
|||||||
|
from typing import List
|
||||||
|
|
||||||
|
import mlx.nn as nn
|
||||||
|
|
||||||
|
from mflux.models.common.weights.loading.weight_definition import ComponentDefinition, TokenizerDefinition
|
||||||
|
from mflux.models.depth_pro.weights.depth_pro_weight_mapping import DepthProWeightMapping
|
||||||
|
|
||||||
|
|
||||||
|
class DepthProWeightDefinition:
|
||||||
|
@staticmethod
|
||||||
|
def get_components() -> List[ComponentDefinition]:
|
||||||
|
return [
|
||||||
|
ComponentDefinition(
|
||||||
|
name="depth_pro",
|
||||||
|
hf_subdir="",
|
||||||
|
loading_mode="torch_checkpoint",
|
||||||
|
mapping_getter=DepthProWeightMapping.get_mapping,
|
||||||
|
download_url="https://ml-site.cdn-apple.com/models/depth-pro/depth_pro.pt",
|
||||||
|
),
|
||||||
|
]
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def get_tokenizers() -> List[TokenizerDefinition]:
|
||||||
|
return []
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def get_download_patterns() -> List[str]:
|
||||||
|
return []
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def quantization_predicate(path: str, module) -> bool:
|
||||||
|
if isinstance(module, nn.Conv2d):
|
||||||
|
return False
|
||||||
|
return hasattr(module, "to_quantized")
|
||||||
307
src/mflux/models/depth_pro/weights/depth_pro_weight_mapping.py
Normal file
307
src/mflux/models/depth_pro/weights/depth_pro_weight_mapping.py
Normal file
@ -0,0 +1,307 @@
|
|||||||
|
from typing import List
|
||||||
|
|
||||||
|
from mflux.models.common.weights.mapping.weight_mapping import WeightMapping, WeightTarget
|
||||||
|
from mflux.models.common.weights.mapping.weight_transforms import WeightTransforms
|
||||||
|
|
||||||
|
|
||||||
|
class DepthProWeightMapping(WeightMapping):
|
||||||
|
@staticmethod
|
||||||
|
def get_mapping() -> List[WeightTarget]:
|
||||||
|
return (
|
||||||
|
DepthProWeightMapping._get_dino_encoder_mapping("patch_encoder")
|
||||||
|
+ DepthProWeightMapping._get_dino_encoder_mapping("image_encoder")
|
||||||
|
+ DepthProWeightMapping._get_upsample_block_mapping("upsample_latent0", num_layers=4)
|
||||||
|
+ DepthProWeightMapping._get_upsample_block_mapping("upsample_latent1", num_layers=3)
|
||||||
|
+ DepthProWeightMapping._get_upsample_block_mapping("upsample0", num_layers=2)
|
||||||
|
+ DepthProWeightMapping._get_upsample_block_mapping("upsample1", num_layers=2)
|
||||||
|
+ DepthProWeightMapping._get_upsample_block_mapping("upsample2", num_layers=2)
|
||||||
|
+ DepthProWeightMapping._get_encoder_conv_mapping()
|
||||||
|
+ DepthProWeightMapping._get_decoder_mapping()
|
||||||
|
+ DepthProWeightMapping._get_head_mapping()
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _get_dino_encoder_mapping(encoder_name: str) -> List[WeightTarget]:
|
||||||
|
prefix = f"encoder.{encoder_name}"
|
||||||
|
return [
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern=f"{prefix}.cls_token",
|
||||||
|
from_pattern=[f"{prefix}.cls_token"],
|
||||||
|
),
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern=f"{prefix}.pos_embed",
|
||||||
|
from_pattern=[f"{prefix}.pos_embed"],
|
||||||
|
),
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern=f"{prefix}.patch_embed.proj.weight",
|
||||||
|
from_pattern=[f"{prefix}.patch_embed.proj.weight"],
|
||||||
|
transform=WeightTransforms.transpose_conv2d_weight,
|
||||||
|
),
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern=f"{prefix}.patch_embed.proj.bias",
|
||||||
|
from_pattern=[f"{prefix}.patch_embed.proj.bias"],
|
||||||
|
),
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern=f"{prefix}.norm.weight",
|
||||||
|
from_pattern=[f"{prefix}.norm.weight"],
|
||||||
|
),
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern=f"{prefix}.norm.bias",
|
||||||
|
from_pattern=[f"{prefix}.norm.bias"],
|
||||||
|
),
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern=f"{prefix}.blocks.{{block}}.norm1.weight",
|
||||||
|
from_pattern=[f"{prefix}.blocks.{{block}}.norm1.weight"],
|
||||||
|
max_blocks=24,
|
||||||
|
),
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern=f"{prefix}.blocks.{{block}}.norm1.bias",
|
||||||
|
from_pattern=[f"{prefix}.blocks.{{block}}.norm1.bias"],
|
||||||
|
max_blocks=24,
|
||||||
|
),
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern=f"{prefix}.blocks.{{block}}.attn.qkv.weight",
|
||||||
|
from_pattern=[f"{prefix}.blocks.{{block}}.attn.qkv.weight"],
|
||||||
|
max_blocks=24,
|
||||||
|
),
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern=f"{prefix}.blocks.{{block}}.attn.qkv.bias",
|
||||||
|
from_pattern=[f"{prefix}.blocks.{{block}}.attn.qkv.bias"],
|
||||||
|
max_blocks=24,
|
||||||
|
),
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern=f"{prefix}.blocks.{{block}}.attn.proj.weight",
|
||||||
|
from_pattern=[f"{prefix}.blocks.{{block}}.attn.proj.weight"],
|
||||||
|
max_blocks=24,
|
||||||
|
),
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern=f"{prefix}.blocks.{{block}}.attn.proj.bias",
|
||||||
|
from_pattern=[f"{prefix}.blocks.{{block}}.attn.proj.bias"],
|
||||||
|
max_blocks=24,
|
||||||
|
),
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern=f"{prefix}.blocks.{{block}}.ls1.gamma",
|
||||||
|
from_pattern=[f"{prefix}.blocks.{{block}}.ls1.gamma"],
|
||||||
|
transform=WeightTransforms.reshape_gamma_to_1d,
|
||||||
|
max_blocks=24,
|
||||||
|
),
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern=f"{prefix}.blocks.{{block}}.norm2.weight",
|
||||||
|
from_pattern=[f"{prefix}.blocks.{{block}}.norm2.weight"],
|
||||||
|
max_blocks=24,
|
||||||
|
),
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern=f"{prefix}.blocks.{{block}}.norm2.bias",
|
||||||
|
from_pattern=[f"{prefix}.blocks.{{block}}.norm2.bias"],
|
||||||
|
max_blocks=24,
|
||||||
|
),
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern=f"{prefix}.blocks.{{block}}.mlp.fc1.weight",
|
||||||
|
from_pattern=[f"{prefix}.blocks.{{block}}.mlp.fc1.weight"],
|
||||||
|
max_blocks=24,
|
||||||
|
),
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern=f"{prefix}.blocks.{{block}}.mlp.fc1.bias",
|
||||||
|
from_pattern=[f"{prefix}.blocks.{{block}}.mlp.fc1.bias"],
|
||||||
|
max_blocks=24,
|
||||||
|
),
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern=f"{prefix}.blocks.{{block}}.mlp.fc2.weight",
|
||||||
|
from_pattern=[f"{prefix}.blocks.{{block}}.mlp.fc2.weight"],
|
||||||
|
max_blocks=24,
|
||||||
|
),
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern=f"{prefix}.blocks.{{block}}.mlp.fc2.bias",
|
||||||
|
from_pattern=[f"{prefix}.blocks.{{block}}.mlp.fc2.bias"],
|
||||||
|
max_blocks=24,
|
||||||
|
),
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern=f"{prefix}.blocks.{{block}}.ls2.gamma",
|
||||||
|
from_pattern=[f"{prefix}.blocks.{{block}}.ls2.gamma"],
|
||||||
|
transform=WeightTransforms.reshape_gamma_to_1d,
|
||||||
|
max_blocks=24,
|
||||||
|
),
|
||||||
|
]
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _get_upsample_block_mapping(block_name: str, num_layers: int) -> List[WeightTarget]:
|
||||||
|
prefix = f"encoder.{block_name}"
|
||||||
|
targets = []
|
||||||
|
|
||||||
|
targets.append(
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern=f"{prefix}.layers.0.weight",
|
||||||
|
from_pattern=[f"{prefix}.0.weight"],
|
||||||
|
transform=WeightTransforms.transpose_conv2d_weight,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
targets.extend(
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern=f"{prefix}.layers.{layer}.weight",
|
||||||
|
from_pattern=[f"{prefix}.{layer}.weight"],
|
||||||
|
transform=WeightTransforms.transpose_conv_transpose2d_weight,
|
||||||
|
)
|
||||||
|
for layer in range(1, num_layers)
|
||||||
|
)
|
||||||
|
|
||||||
|
return targets
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _get_encoder_conv_mapping() -> List[WeightTarget]:
|
||||||
|
return [
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern="encoder.upsample_lowres.weight",
|
||||||
|
from_pattern=["encoder.upsample_lowres.weight"],
|
||||||
|
transform=WeightTransforms.transpose_conv_transpose2d_weight,
|
||||||
|
),
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern="encoder.upsample_lowres.bias",
|
||||||
|
from_pattern=["encoder.upsample_lowres.bias"],
|
||||||
|
),
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern="encoder.fuse_lowres.weight",
|
||||||
|
from_pattern=["encoder.fuse_lowres.weight"],
|
||||||
|
transform=WeightTransforms.transpose_conv2d_weight,
|
||||||
|
),
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern="encoder.fuse_lowres.bias",
|
||||||
|
from_pattern=["encoder.fuse_lowres.bias"],
|
||||||
|
),
|
||||||
|
]
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _get_decoder_mapping() -> List[WeightTarget]:
|
||||||
|
return (
|
||||||
|
DepthProWeightMapping._get_decoder_convs_mapping()
|
||||||
|
+ DepthProWeightMapping._get_fusion_block_mapping(0)
|
||||||
|
+ DepthProWeightMapping._get_fusion_block_mapping(1)
|
||||||
|
+ DepthProWeightMapping._get_fusion_block_mapping(2)
|
||||||
|
+ DepthProWeightMapping._get_fusion_block_mapping(3)
|
||||||
|
+ DepthProWeightMapping._get_fusion_block_mapping(4)
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _get_decoder_convs_mapping() -> List[WeightTarget]:
|
||||||
|
return [
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern="decoder.convs.1.weight",
|
||||||
|
from_pattern=["decoder.convs.1.weight"],
|
||||||
|
transform=WeightTransforms.transpose_conv2d_weight,
|
||||||
|
),
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern="decoder.convs.2.weight",
|
||||||
|
from_pattern=["decoder.convs.2.weight"],
|
||||||
|
transform=WeightTransforms.transpose_conv2d_weight,
|
||||||
|
),
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern="decoder.convs.3.weight",
|
||||||
|
from_pattern=["decoder.convs.3.weight"],
|
||||||
|
transform=WeightTransforms.transpose_conv2d_weight,
|
||||||
|
),
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern="decoder.convs.4.weight",
|
||||||
|
from_pattern=["decoder.convs.4.weight"],
|
||||||
|
transform=WeightTransforms.transpose_conv2d_weight,
|
||||||
|
),
|
||||||
|
]
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _get_fusion_block_mapping(i: int) -> List[WeightTarget]:
|
||||||
|
targets = [
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern=f"decoder.fusions.{i}.resnet1.residual.1.weight",
|
||||||
|
from_pattern=[f"decoder.fusions.{i}.resnet1.residual.1.weight"],
|
||||||
|
transform=WeightTransforms.transpose_conv2d_weight,
|
||||||
|
),
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern=f"decoder.fusions.{i}.resnet1.residual.1.bias",
|
||||||
|
from_pattern=[f"decoder.fusions.{i}.resnet1.residual.1.bias"],
|
||||||
|
),
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern=f"decoder.fusions.{i}.resnet1.residual.3.weight",
|
||||||
|
from_pattern=[f"decoder.fusions.{i}.resnet1.residual.3.weight"],
|
||||||
|
transform=WeightTransforms.transpose_conv2d_weight,
|
||||||
|
),
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern=f"decoder.fusions.{i}.resnet1.residual.3.bias",
|
||||||
|
from_pattern=[f"decoder.fusions.{i}.resnet1.residual.3.bias"],
|
||||||
|
),
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern=f"decoder.fusions.{i}.resnet2.residual.1.weight",
|
||||||
|
from_pattern=[f"decoder.fusions.{i}.resnet2.residual.1.weight"],
|
||||||
|
transform=WeightTransforms.transpose_conv2d_weight,
|
||||||
|
),
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern=f"decoder.fusions.{i}.resnet2.residual.1.bias",
|
||||||
|
from_pattern=[f"decoder.fusions.{i}.resnet2.residual.1.bias"],
|
||||||
|
),
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern=f"decoder.fusions.{i}.resnet2.residual.3.weight",
|
||||||
|
from_pattern=[f"decoder.fusions.{i}.resnet2.residual.3.weight"],
|
||||||
|
transform=WeightTransforms.transpose_conv2d_weight,
|
||||||
|
),
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern=f"decoder.fusions.{i}.resnet2.residual.3.bias",
|
||||||
|
from_pattern=[f"decoder.fusions.{i}.resnet2.residual.3.bias"],
|
||||||
|
),
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern=f"decoder.fusions.{i}.out_conv.weight",
|
||||||
|
from_pattern=[f"decoder.fusions.{i}.out_conv.weight"],
|
||||||
|
transform=WeightTransforms.transpose_conv2d_weight,
|
||||||
|
),
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern=f"decoder.fusions.{i}.out_conv.bias",
|
||||||
|
from_pattern=[f"decoder.fusions.{i}.out_conv.bias"],
|
||||||
|
),
|
||||||
|
]
|
||||||
|
if i > 0:
|
||||||
|
targets.append(
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern=f"decoder.fusions.{i}.deconv.weight",
|
||||||
|
from_pattern=[f"decoder.fusions.{i}.deconv.weight"],
|
||||||
|
transform=WeightTransforms.transpose_conv_transpose2d_weight,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
return targets
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _get_head_mapping() -> List[WeightTarget]:
|
||||||
|
return [
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern="head.convs.0.weight",
|
||||||
|
from_pattern=["head.0.weight"],
|
||||||
|
transform=WeightTransforms.transpose_conv2d_weight,
|
||||||
|
),
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern="head.convs.0.bias",
|
||||||
|
from_pattern=["head.0.bias"],
|
||||||
|
),
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern="head.convs.1.weight",
|
||||||
|
from_pattern=["head.1.weight"],
|
||||||
|
transform=WeightTransforms.transpose_conv_transpose2d_weight,
|
||||||
|
),
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern="head.convs.1.bias",
|
||||||
|
from_pattern=["head.1.bias"],
|
||||||
|
),
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern="head.convs.2.weight",
|
||||||
|
from_pattern=["head.2.weight"],
|
||||||
|
transform=WeightTransforms.transpose_conv2d_weight,
|
||||||
|
),
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern="head.convs.2.bias",
|
||||||
|
from_pattern=["head.2.bias"],
|
||||||
|
),
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern="head.convs.4.weight",
|
||||||
|
from_pattern=["head.4.weight"],
|
||||||
|
transform=WeightTransforms.transpose_conv2d_weight,
|
||||||
|
),
|
||||||
|
WeightTarget(
|
||||||
|
to_pattern="head.convs.4.bias",
|
||||||
|
from_pattern=["head.4.bias"],
|
||||||
|
),
|
||||||
|
]
|
||||||
@ -1,131 +0,0 @@
|
|||||||
import logging
|
|
||||||
import urllib.error
|
|
||||||
import urllib.request
|
|
||||||
|
|
||||||
import mlx.core as mx
|
|
||||||
import torch
|
|
||||||
from mlx.utils import tree_unflatten
|
|
||||||
|
|
||||||
from mflux.models.flux.weights.weight_handler import MetaData
|
|
||||||
from mflux.models.flux.weights.weight_util import WeightUtil
|
|
||||||
from mflux.ui.defaults import MFLUX_CACHE_DIR
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
|
||||||
|
|
||||||
|
|
||||||
class WeightHandlerDepthPro:
|
|
||||||
def __init__(self, weights: dict, meta_data: MetaData):
|
|
||||||
self.weights = weights
|
|
||||||
self.meta_data = meta_data
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def load_weights() -> "WeightHandlerDepthPro":
|
|
||||||
model_path = WeightHandlerDepthPro._download_or_get_cached_weights()
|
|
||||||
pt_weights = torch.load(model_path, map_location="cpu")
|
|
||||||
weights = WeightHandlerDepthPro._to_mlx_weights(pt_weights)
|
|
||||||
weights = [WeightUtil.reshape_weights(k, v) for k, v in weights.items()]
|
|
||||||
weights = WeightUtil.flatten(weights)
|
|
||||||
weights = tree_unflatten(weights)
|
|
||||||
return WeightHandlerDepthPro(
|
|
||||||
weights=weights,
|
|
||||||
meta_data=MetaData(quantization_level=None)
|
|
||||||
) # fmt:off
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _to_mlx_weights(pt_weights) -> dict:
|
|
||||||
mlx_weights = {}
|
|
||||||
for key, value in pt_weights.items():
|
|
||||||
if isinstance(value, torch.Tensor):
|
|
||||||
mlx_weights[key] = mx.array(value.numpy())
|
|
||||||
else:
|
|
||||||
mlx_weights[key] = value
|
|
||||||
return mlx_weights
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _download_or_get_cached_weights():
|
|
||||||
APPLE_MODEL_URL = "https://ml-site.cdn-apple.com/models/depth-pro/depth_pro.pt"
|
|
||||||
|
|
||||||
# 1. Create cache directory for the model
|
|
||||||
cache_dir = MFLUX_CACHE_DIR / "depth_pro"
|
|
||||||
|
|
||||||
cache_dir.mkdir(parents=True, exist_ok=True)
|
|
||||||
model_path = cache_dir / "depth_pro.pt"
|
|
||||||
|
|
||||||
# 2. Download if model doesn't exist
|
|
||||||
if not model_path.exists():
|
|
||||||
logger.info("Downloading Depth Pro model from Apple...")
|
|
||||||
try:
|
|
||||||
urllib.request.urlretrieve(APPLE_MODEL_URL, model_path)
|
|
||||||
logger.info(f"Downloaded model to {model_path}")
|
|
||||||
except (urllib.error.URLError, urllib.error.HTTPError) as e:
|
|
||||||
logger.error(f"Failed to download model: {e}")
|
|
||||||
logger.info(f"Please manually download from: {APPLE_MODEL_URL}")
|
|
||||||
if not model_path.exists():
|
|
||||||
raise FileNotFoundError(f"Model file not found at {model_path}")
|
|
||||||
|
|
||||||
return model_path
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def reposition_encoder_weights(depth_pro_weights, name):
|
|
||||||
tmp = depth_pro_weights.weights["encoder"][name]
|
|
||||||
depth_pro_weights.weights["encoder"][name] = {}
|
|
||||||
depth_pro_weights.weights["encoder"][name]["layers"] = tmp
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def reposition_head_weights(depth_pro_weights):
|
|
||||||
tmp = depth_pro_weights.weights["head"]
|
|
||||||
depth_pro_weights.weights["head"] = {}
|
|
||||||
depth_pro_weights.weights["head"]["convs"] = tmp
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def reshape_transposed_convolution_weights(depth_pro_weights):
|
|
||||||
WeightHandlerDepthPro._reshape_upsample(depth_pro_weights, "upsample_latent0", 1)
|
|
||||||
WeightHandlerDepthPro._reshape_upsample(depth_pro_weights, "upsample_latent0", 2)
|
|
||||||
WeightHandlerDepthPro._reshape_upsample(depth_pro_weights, "upsample_latent0", 3)
|
|
||||||
|
|
||||||
WeightHandlerDepthPro._reshape_upsample(depth_pro_weights, "upsample_latent1", 1)
|
|
||||||
WeightHandlerDepthPro._reshape_upsample(depth_pro_weights, "upsample_latent1", 2)
|
|
||||||
|
|
||||||
WeightHandlerDepthPro._reshape_upsample(depth_pro_weights, "upsample0", 1)
|
|
||||||
WeightHandlerDepthPro._reshape_upsample(depth_pro_weights, "upsample1", 1)
|
|
||||||
WeightHandlerDepthPro._reshape_upsample(depth_pro_weights, "upsample2", 1)
|
|
||||||
|
|
||||||
WeightHandlerDepthPro._reshape_upsample_lowres(depth_pro_weights)
|
|
||||||
|
|
||||||
WeightHandlerDepthPro._reshape_deconv(depth_pro_weights, 1)
|
|
||||||
WeightHandlerDepthPro._reshape_deconv(depth_pro_weights, 2)
|
|
||||||
WeightHandlerDepthPro._reshape_deconv(depth_pro_weights, 3)
|
|
||||||
WeightHandlerDepthPro._reshape_deconv(depth_pro_weights, 4)
|
|
||||||
|
|
||||||
WeightHandlerDepthPro._reshape_head(depth_pro_weights, 1)
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _reshape_upsample(depth_pro_weights, name, layer):
|
|
||||||
tmp = depth_pro_weights.weights["encoder"][name]["layers"][layer]["weight"]
|
|
||||||
tmp = WeightHandlerDepthPro._reshape(tmp)
|
|
||||||
depth_pro_weights.weights["encoder"][name]["layers"][layer]["weight"] = tmp
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _reshape_upsample_lowres(depth_pro_weights):
|
|
||||||
tmp = depth_pro_weights.weights["encoder"]["upsample_lowres"]["weight"]
|
|
||||||
tmp = WeightHandlerDepthPro._reshape(tmp)
|
|
||||||
depth_pro_weights.weights["encoder"]["upsample_lowres"]["weight"] = tmp
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _reshape_deconv(depth_pro_weights, layer):
|
|
||||||
tmp = depth_pro_weights.weights["decoder"]["fusions"][layer]["deconv"]["weight"]
|
|
||||||
tmp = WeightHandlerDepthPro._reshape(tmp)
|
|
||||||
depth_pro_weights.weights["decoder"]["fusions"][layer]["deconv"]["weight"] = tmp
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _reshape_head(depth_pro_weights, layer):
|
|
||||||
tmp = depth_pro_weights.weights["head"]["convs"][layer]["weight"]
|
|
||||||
tmp = WeightHandlerDepthPro._reshape(tmp)
|
|
||||||
depth_pro_weights.weights["head"]["convs"][layer]["weight"] = tmp
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _reshape(tensor):
|
|
||||||
tensor = tensor.transpose(0, 3, 1, 2)
|
|
||||||
tensor = tensor.transpose(1, 0, 2, 3)
|
|
||||||
tensor = tensor.transpose(0, 2, 3, 1)
|
|
||||||
return tensor
|
|
||||||
0
src/mflux/models/fibo/cli/__init__.py
Normal file
0
src/mflux/models/fibo/cli/__init__.py
Normal file
@ -1,17 +1,12 @@
|
|||||||
import gc
|
|
||||||
import json
|
|
||||||
|
|
||||||
import mlx.core as mx
|
|
||||||
|
|
||||||
from mflux.callbacks.callback_manager import CallbackManager
|
from mflux.callbacks.callback_manager import CallbackManager
|
||||||
from mflux.config.config import Config
|
from mflux.cli.defaults import defaults as ui_defaults
|
||||||
from mflux.config.model_config import ModelConfig
|
from mflux.cli.parser.parsers import CommandLineParser
|
||||||
|
from mflux.models.fibo.latent_creator.fibo_latent_creator import FiboLatentCreator
|
||||||
from mflux.models.fibo.variants.txt2img.fibo import FIBO
|
from mflux.models.fibo.variants.txt2img.fibo import FIBO
|
||||||
from mflux.models.fibo_vlm.model.fibo_vlm import FiboVLM
|
from mflux.models.fibo.variants.txt2img.util import FiboUtil
|
||||||
from mflux.ui import defaults as ui_defaults
|
from mflux.utils.dimension_resolver import DimensionResolver
|
||||||
from mflux.ui.cli.parsers import CommandLineParser
|
|
||||||
from mflux.ui.prompt_utils import PromptUtils
|
|
||||||
from mflux.utils.exceptions import PromptFileReadError, StopImageGenerationException
|
from mflux.utils.exceptions import PromptFileReadError, StopImageGenerationException
|
||||||
|
from mflux.utils.prompt_util import PromptUtil
|
||||||
|
|
||||||
|
|
||||||
def main():
|
def main():
|
||||||
@ -20,7 +15,7 @@ def main():
|
|||||||
parser.add_general_arguments()
|
parser.add_general_arguments()
|
||||||
parser.add_model_arguments(require_model_arg=False)
|
parser.add_model_arguments(require_model_arg=False)
|
||||||
parser.add_lora_arguments()
|
parser.add_lora_arguments()
|
||||||
parser.add_image_generator_arguments(supports_metadata_config=True)
|
parser.add_image_generator_arguments(supports_metadata_config=True, supports_dimension_scale_factor=True)
|
||||||
parser.add_image_to_image_arguments(required=False)
|
parser.add_image_to_image_arguments(required=False)
|
||||||
parser.add_output_arguments()
|
parser.add_output_arguments()
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
@ -29,34 +24,42 @@ def main():
|
|||||||
if args.guidance is None:
|
if args.guidance is None:
|
||||||
args.guidance = ui_defaults.GUIDANCE_SCALE
|
args.guidance = ui_defaults.GUIDANCE_SCALE
|
||||||
|
|
||||||
json_prompt = _get_json_prompt(args)
|
json_prompt = FiboUtil.get_json_prompt(args, quantize=args.quantize)
|
||||||
|
|
||||||
# 1. Load the FIBO model
|
# 1. Load the FIBO model
|
||||||
fibo = FIBO(
|
fibo = FIBO(
|
||||||
model_config=ModelConfig.fibo(),
|
|
||||||
quantize=args.quantize,
|
quantize=args.quantize,
|
||||||
local_path=args.path,
|
model_path=args.model_path,
|
||||||
)
|
)
|
||||||
|
|
||||||
# 2. Register callbacks
|
# 2. Register callbacks
|
||||||
memory_saver = CallbackManager.register_callbacks(args=args, model=fibo)
|
memory_saver = CallbackManager.register_callbacks(
|
||||||
|
args=args,
|
||||||
|
model=fibo,
|
||||||
|
latent_creator=FiboLatentCreator,
|
||||||
|
)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
# Resolve dimensions (supports ScaleFactor like "2x" when --image-path is provided)
|
||||||
|
width, height = DimensionResolver.resolve(
|
||||||
|
width=args.width,
|
||||||
|
height=args.height,
|
||||||
|
reference_image_path=args.image_path,
|
||||||
|
)
|
||||||
|
|
||||||
for seed in args.seed:
|
for seed in args.seed:
|
||||||
# 3. Generate an image for each seed value
|
# 3. Generate an image for each seed value
|
||||||
image = fibo.generate_image(
|
image = fibo.generate_image(
|
||||||
seed=seed,
|
seed=seed,
|
||||||
prompt=json_prompt,
|
prompt=json_prompt,
|
||||||
negative_prompt=PromptUtils.get_effective_negative_prompt(args),
|
width=width,
|
||||||
config=Config(
|
height=height,
|
||||||
num_inference_steps=args.steps,
|
guidance=args.guidance,
|
||||||
height=args.height,
|
image_path=args.image_path,
|
||||||
width=args.width,
|
num_inference_steps=args.steps,
|
||||||
guidance=args.guidance,
|
image_strength=args.image_strength,
|
||||||
image_path=args.image_path,
|
scheduler="flow_match_euler_discrete",
|
||||||
image_strength=args.image_strength,
|
negative_prompt=PromptUtil.read_negative_prompt(args),
|
||||||
scheduler="flow_match_euler_discrete",
|
|
||||||
),
|
|
||||||
)
|
)
|
||||||
|
|
||||||
# 4. Save the image
|
# 4. Save the image
|
||||||
@ -68,20 +71,5 @@ def main():
|
|||||||
print(memory_saver.memory_stats())
|
print(memory_saver.memory_stats())
|
||||||
|
|
||||||
|
|
||||||
def _get_json_prompt(args):
|
|
||||||
prompt = PromptUtils.get_effective_prompt(args)
|
|
||||||
|
|
||||||
try:
|
|
||||||
json.loads(prompt)
|
|
||||||
json_prompt = prompt
|
|
||||||
except json.JSONDecodeError:
|
|
||||||
vlm = FiboVLM()
|
|
||||||
json_prompt = vlm.generate(prompt=prompt, seed=42)
|
|
||||||
del vlm
|
|
||||||
gc.collect()
|
|
||||||
mx.clear_cache()
|
|
||||||
return json_prompt
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
main()
|
main()
|
||||||
@ -1,47 +1,66 @@
|
|||||||
from mflux.config.model_config import ModelConfig
|
from mflux.callbacks.callback_registry import CallbackRegistry
|
||||||
|
from mflux.models.common.config import ModelConfig
|
||||||
|
from mflux.models.common.tokenizer import TokenizerLoader
|
||||||
|
from mflux.models.common.weights.loading.loaded_weights import LoadedWeights
|
||||||
|
from mflux.models.common.weights.loading.weight_applier import WeightApplier
|
||||||
|
from mflux.models.common.weights.loading.weight_loader import WeightLoader
|
||||||
from mflux.models.fibo.model.fibo_text_encoder import SmolLM3_3B_TextEncoder
|
from mflux.models.fibo.model.fibo_text_encoder import SmolLM3_3B_TextEncoder
|
||||||
from mflux.models.fibo.model.fibo_transformer import FiboTransformer
|
from mflux.models.fibo.model.fibo_transformer import FiboTransformer
|
||||||
from mflux.models.fibo.model.fibo_vae.wan_2_2_vae import Wan2_2_VAE
|
from mflux.models.fibo.model.fibo_vae.wan_2_2_vae import Wan2_2_VAE
|
||||||
from mflux.models.fibo.tokenizer import FiboTokenizerHandler
|
from mflux.models.fibo.weights.fibo_weight_definition import FIBOWeightDefinition
|
||||||
from mflux.models.fibo.weights.fibo_weight_handler import FIBOWeightHandler
|
|
||||||
from mflux.models.fibo.weights.fibo_weight_util import FIBOWeightUtil
|
|
||||||
|
|
||||||
|
|
||||||
class FIBOInitializer:
|
class FIBOInitializer:
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def init(
|
def init(
|
||||||
fibo_model,
|
model,
|
||||||
model_config: ModelConfig | None = None,
|
model_config: ModelConfig,
|
||||||
quantize: int | None = None,
|
quantize: int | None = None,
|
||||||
local_path: str | None = None,
|
model_path: str | None = None,
|
||||||
|
lora_paths: list[str] | None = None, # noqa: ARG004
|
||||||
|
lora_scales: list[float] | None = None, # noqa: ARG004
|
||||||
) -> None:
|
) -> None:
|
||||||
# 1. Load VAE weights
|
path = model_path if model_path else model_config.model_name
|
||||||
weights = FIBOWeightHandler.load_regular_weights(
|
FIBOInitializer._init_config(model, model_config)
|
||||||
repo_id=model_config.model_name,
|
weights = FIBOInitializer._load_weights(path)
|
||||||
local_path=local_path,
|
FIBOInitializer._init_tokenizers(model, path)
|
||||||
|
FIBOInitializer._init_models(model)
|
||||||
|
FIBOInitializer._apply_weights(model, weights, quantize)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _init_config(model, model_config: ModelConfig) -> None:
|
||||||
|
model.model_config = model_config
|
||||||
|
model.callbacks = CallbackRegistry()
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _load_weights(model_path: str) -> LoadedWeights:
|
||||||
|
return WeightLoader.load(
|
||||||
|
weight_definition=FIBOWeightDefinition,
|
||||||
|
model_path=model_path,
|
||||||
)
|
)
|
||||||
|
|
||||||
# 2. Initialize tokenizers
|
@staticmethod
|
||||||
tokenizer_handler = FiboTokenizerHandler(
|
def _init_tokenizers(model, model_path: str) -> None:
|
||||||
repo_id=model_config.model_name,
|
model.tokenizers = TokenizerLoader.load_all(
|
||||||
local_path=local_path,
|
definitions=FIBOWeightDefinition.get_tokenizers(),
|
||||||
)
|
model_path=model_path,
|
||||||
fibo_model.fibo_tokenizer = tokenizer_handler.fibo
|
|
||||||
|
|
||||||
# 3. Initialize all models
|
|
||||||
fibo_model.vae = Wan2_2_VAE()
|
|
||||||
fibo_model.text_encoder = SmolLM3_3B_TextEncoder()
|
|
||||||
fibo_model.transformer = FiboTransformer(
|
|
||||||
in_channels=48,
|
|
||||||
num_layers=8,
|
|
||||||
num_single_layers=38,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
# 4. Apply weights and quantize VAE, transformer, and text encoder
|
@staticmethod
|
||||||
fibo_model.bits = FIBOWeightUtil.set_weights_and_quantize(
|
def _init_models(model) -> None:
|
||||||
quantize_arg=quantize,
|
model.vae = Wan2_2_VAE()
|
||||||
|
model.text_encoder = SmolLM3_3B_TextEncoder()
|
||||||
|
model.transformer = FiboTransformer()
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _apply_weights(model, weights: LoadedWeights, quantize: int | None) -> None:
|
||||||
|
model.bits = WeightApplier.apply_and_quantize(
|
||||||
weights=weights,
|
weights=weights,
|
||||||
vae=fibo_model.vae,
|
quantize_arg=quantize,
|
||||||
transformer=fibo_model.transformer,
|
weight_definition=FIBOWeightDefinition,
|
||||||
text_encoder=fibo_model.text_encoder,
|
models={
|
||||||
|
"vae": model.vae,
|
||||||
|
"transformer": model.transformer,
|
||||||
|
"text_encoder": model.text_encoder,
|
||||||
|
},
|
||||||
)
|
)
|
||||||
|
|||||||
@ -17,3 +17,13 @@ class FiboLatentCreator:
|
|||||||
batch_size, channels, latent_height, latent_width = latents.shape
|
batch_size, channels, latent_height, latent_width = latents.shape
|
||||||
latents = mx.transpose(latents, (0, 2, 3, 1))
|
latents = mx.transpose(latents, (0, 2, 3, 1))
|
||||||
return mx.reshape(latents, (batch_size, latent_height * latent_width, channels))
|
return mx.reshape(latents, (batch_size, latent_height * latent_width, channels))
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def unpack_latents(latents: mx.array, height: int, width: int) -> mx.array:
|
||||||
|
batch_size, seq_len, channels = latents.shape
|
||||||
|
vae_scale_factor = 16
|
||||||
|
latent_height = height // vae_scale_factor
|
||||||
|
latent_width = width // vae_scale_factor
|
||||||
|
latents = mx.reshape(latents, (batch_size, latent_height, latent_width, channels))
|
||||||
|
latents = mx.transpose(latents, (0, 3, 1, 2))
|
||||||
|
return latents
|
||||||
|
|||||||
@ -1 +0,0 @@
|
|||||||
"""FIBO model components."""
|
|
||||||
@ -3,8 +3,8 @@ from typing import List, Union
|
|||||||
|
|
||||||
import mlx.core as mx
|
import mlx.core as mx
|
||||||
|
|
||||||
|
from mflux.models.common.tokenizer import Tokenizer
|
||||||
from mflux.models.fibo.model.fibo_text_encoder.smol_lm3_3b_text_encoder import SmolLM3_3B_TextEncoder
|
from mflux.models.fibo.model.fibo_text_encoder.smol_lm3_3b_text_encoder import SmolLM3_3B_TextEncoder
|
||||||
from mflux.models.fibo.tokenizer.fibo_tokenizer import TokenizerFibo
|
|
||||||
|
|
||||||
|
|
||||||
class PromptEncoder:
|
class PromptEncoder:
|
||||||
@ -12,7 +12,7 @@ class PromptEncoder:
|
|||||||
def encode_prompt(
|
def encode_prompt(
|
||||||
prompt: str,
|
prompt: str,
|
||||||
negative_prompt: str | None,
|
negative_prompt: str | None,
|
||||||
tokenizer: TokenizerFibo,
|
tokenizer: Tokenizer,
|
||||||
text_encoder: SmolLM3_3B_TextEncoder,
|
text_encoder: SmolLM3_3B_TextEncoder,
|
||||||
) -> tuple[str, mx.array, List[mx.array]]:
|
) -> tuple[str, mx.array, List[mx.array]]:
|
||||||
# 0. Set default negative prompt if not provided
|
# 0. Set default negative prompt if not provided
|
||||||
@ -87,7 +87,7 @@ class PromptEncoder:
|
|||||||
def _get_prompt_embeds(
|
def _get_prompt_embeds(
|
||||||
prompt: Union[str, List[str]],
|
prompt: Union[str, List[str]],
|
||||||
text_encoder: SmolLM3_3B_TextEncoder,
|
text_encoder: SmolLM3_3B_TextEncoder,
|
||||||
tokenizer: TokenizerFibo,
|
tokenizer: Tokenizer,
|
||||||
num_images_per_prompt: int = 1,
|
num_images_per_prompt: int = 1,
|
||||||
max_sequence_length: int = 2048,
|
max_sequence_length: int = 2048,
|
||||||
tokenization_prefix: str | None = None,
|
tokenization_prefix: str | None = None,
|
||||||
@ -97,13 +97,9 @@ class PromptEncoder:
|
|||||||
raise ValueError("`prompt` must be a non-empty string or list of strings.")
|
raise ValueError("`prompt` must be a non-empty string or list of strings.")
|
||||||
|
|
||||||
# 1) Tokenize and convert to MX
|
# 1) Tokenize and convert to MX
|
||||||
input_ids_mx, attention_mask_mx = tokenizer.tokenize(
|
tokenizer_output = tokenizer.tokenize(prompt=prompts, max_length=max_sequence_length)
|
||||||
prompts=prompts,
|
input_ids_mx = tokenizer_output.input_ids
|
||||||
max_length=max_sequence_length,
|
attention_mask_mx = tokenizer_output.attention_mask
|
||||||
padding="longest",
|
|
||||||
truncation=True,
|
|
||||||
add_special_tokens=True,
|
|
||||||
)
|
|
||||||
|
|
||||||
# 2) Run MLX text encoder and collect hidden states
|
# 2) Run MLX text encoder and collect hidden states
|
||||||
hidden_states_list = text_encoder(
|
hidden_states_list = text_encoder(
|
||||||
|
|||||||
@ -58,9 +58,6 @@ class SmolLM3_3B_SelfAttention(nn.Module):
|
|||||||
k = k.astype(hidden_states.dtype)
|
k = k.astype(hidden_states.dtype)
|
||||||
v = v.astype(hidden_states.dtype)
|
v = v.astype(hidden_states.dtype)
|
||||||
|
|
||||||
# Force evaluation to ensure computation is complete
|
|
||||||
mx.eval(q, k, v)
|
|
||||||
|
|
||||||
# Reshape to (batch, heads, seq, head_dim)
|
# Reshape to (batch, heads, seq, head_dim)
|
||||||
q = q.reshape(batch_size, seq_len, self.num_attention_heads, self.head_dim).transpose(0, 2, 1, 3)
|
q = q.reshape(batch_size, seq_len, self.num_attention_heads, self.head_dim).transpose(0, 2, 1, 3)
|
||||||
k = k.reshape(batch_size, seq_len, self.num_key_value_heads, self.head_dim).transpose(0, 2, 1, 3)
|
k = k.reshape(batch_size, seq_len, self.num_key_value_heads, self.head_dim).transpose(0, 2, 1, 3)
|
||||||
|
|||||||
@ -1,7 +1,7 @@
|
|||||||
import mlx.core as mx
|
import mlx.core as mx
|
||||||
from mlx import nn
|
from mlx import nn
|
||||||
|
|
||||||
from mflux.config.runtime_config import RuntimeConfig
|
from mflux.models.common.config.config import Config
|
||||||
from mflux.models.fibo.model.fibo_transformer.fibo_embed_nd import FiboEmbedND
|
from mflux.models.fibo.model.fibo_transformer.fibo_embed_nd import FiboEmbedND
|
||||||
from mflux.models.fibo.model.fibo_transformer.joint_transformer_block import FiboJointTransformerBlock
|
from mflux.models.fibo.model.fibo_transformer.joint_transformer_block import FiboJointTransformerBlock
|
||||||
from mflux.models.fibo.model.fibo_transformer.single_transformer_block import FiboSingleTransformerBlock
|
from mflux.models.fibo.model.fibo_transformer.single_transformer_block import FiboSingleTransformerBlock
|
||||||
@ -31,7 +31,7 @@ class FiboTransformer(nn.Module):
|
|||||||
def __call__(
|
def __call__(
|
||||||
self,
|
self,
|
||||||
t: int,
|
t: int,
|
||||||
config: RuntimeConfig,
|
config: Config,
|
||||||
hidden_states: mx.array,
|
hidden_states: mx.array,
|
||||||
encoder_hidden_states: mx.array,
|
encoder_hidden_states: mx.array,
|
||||||
text_encoder_layers: list[mx.array],
|
text_encoder_layers: list[mx.array],
|
||||||
@ -150,7 +150,7 @@ class FiboTransformer(nn.Module):
|
|||||||
@staticmethod
|
@staticmethod
|
||||||
def _compute_time_embeddings(
|
def _compute_time_embeddings(
|
||||||
t: int,
|
t: int,
|
||||||
config: RuntimeConfig,
|
config: Config,
|
||||||
batch_size: int,
|
batch_size: int,
|
||||||
dtype: mx.Dtype,
|
dtype: mx.Dtype,
|
||||||
time_embed: BriaFiboTimestepProjEmbeddings,
|
time_embed: BriaFiboTimestepProjEmbeddings,
|
||||||
@ -163,7 +163,7 @@ class FiboTransformer(nn.Module):
|
|||||||
def _compute_rotary_embeddings(
|
def _compute_rotary_embeddings(
|
||||||
encoder_hidden_states: mx.array,
|
encoder_hidden_states: mx.array,
|
||||||
pos_embed: FiboEmbedND,
|
pos_embed: FiboEmbedND,
|
||||||
config: RuntimeConfig,
|
config: Config,
|
||||||
dtype: mx.Dtype,
|
dtype: mx.Dtype,
|
||||||
) -> mx.array:
|
) -> mx.array:
|
||||||
max_tokens = encoder_hidden_states.shape[1]
|
max_tokens = encoder_hidden_states.shape[1]
|
||||||
@ -209,7 +209,7 @@ class FiboTransformer(nn.Module):
|
|||||||
@staticmethod
|
@staticmethod
|
||||||
def _compute_attention_mask(
|
def _compute_attention_mask(
|
||||||
batch_size: int,
|
batch_size: int,
|
||||||
config: RuntimeConfig,
|
config: Config,
|
||||||
encoder_hidden_states: mx.array,
|
encoder_hidden_states: mx.array,
|
||||||
max_tokens: int,
|
max_tokens: int,
|
||||||
) -> mx.array:
|
) -> mx.array:
|
||||||
|
|||||||
@ -1 +0,0 @@
|
|||||||
"""FIBO VAE components."""
|
|
||||||
@ -1,5 +1,3 @@
|
|||||||
"""FIBO VAE common/shared components."""
|
|
||||||
|
|
||||||
from mflux.models.fibo.model.fibo_vae.common.wan_2_2_attention_block import Wan2_2_AttentionBlock
|
from mflux.models.fibo.model.fibo_vae.common.wan_2_2_attention_block import Wan2_2_AttentionBlock
|
||||||
from mflux.models.fibo.model.fibo_vae.common.wan_2_2_causal_conv_3d import Wan2_2_CausalConv3d
|
from mflux.models.fibo.model.fibo_vae.common.wan_2_2_causal_conv_3d import Wan2_2_CausalConv3d
|
||||||
from mflux.models.fibo.model.fibo_vae.common.wan_2_2_mid_block import Wan2_2_MidBlock
|
from mflux.models.fibo.model.fibo_vae.common.wan_2_2_mid_block import Wan2_2_MidBlock
|
||||||
|
|||||||
@ -1,5 +1,3 @@
|
|||||||
"""FIBO VAE decoder components."""
|
|
||||||
|
|
||||||
from mflux.models.fibo.model.fibo_vae.decoder.wan_2_2_decoder_3d import Wan2_2_Decoder3d
|
from mflux.models.fibo.model.fibo_vae.decoder.wan_2_2_decoder_3d import Wan2_2_Decoder3d
|
||||||
|
|
||||||
__all__ = ["Wan2_2_Decoder3d"]
|
__all__ = ["Wan2_2_Decoder3d"]
|
||||||
|
|||||||
Some files were not shown because too many files have changed in this diff Show More
Loading…
Reference in New Issue
Block a user