diff --git a/README.md b/README.md
index 70ebb67..84358cb 100644
--- a/README.md
+++ b/README.md
@@ -23,6 +23,7 @@ Run the powerful [FLUX](https://blackforestlabs.ai/#get-flux), [Qwen Image](http
- [🦙 Qwen Models](#-qwen-models)
* [🖼️ Qwen Image](#%EF%B8%8F-qwen-image)
* [✏️ Qwen Image Edit](#%EF%B8%8F-qwen-image-edit)
+ * [Qwen Image Layered](#qwen-image-layered)
- [🌀 FIBO](#-fibo)
- [⚡ Z-Image](#-z-image)
- [🔌 LoRA](#-lora)
@@ -374,6 +375,35 @@ See the [Qwen Image](#-qwen-image) section for more details on this feature.
+#### Qwen Image Layered Command-Line Arguments
+
+
+Click to expand Qwen Image Layered arguments
+
+The `mflux-generate-qwen-layered` command decomposes an input image into separate RGBA layers:
+
+- **`--image`** (required, `str`): Path to the input image to decompose into layers.
+
+- **`--layers`** (optional, `int`, default: `4`): Number of output layers to generate. Each layer will contain distinct visual elements from the source image.
+
+- **`--steps`** (optional, `int`, default: `50`): Number of inference steps. More steps generally produce better results.
+
+- **`--resolution`** (optional, `int`, default: `640`): Target resolution bucket (`640` or `1024`). The image will be resized to this resolution while maintaining aspect ratio.
+
+- **`--guidance`** (optional, `float`, default: `4.0`): Guidance scale for the decomposition.
+
+- **`--cfg-normalize`** (optional, flag): Enable CFG normalization for more stable guidance.
+
+- **`--prompt`** (optional, `str`): Text description of the image (auto-generated if not provided).
+
+- **`--negative-prompt`** (optional, `str`, default: `"blurry, bad quality"`): Negative prompt for quality guidance.
+
+- **`--output-dir`** (optional, `str`, default: `"."`): Directory where layer images will be saved.
+
+See the [Qwen Image Layered](#-qwen-image-layered) section for more details on this feature.
+
+
+
#### Qwen Image Edit Command-Line Arguments
@@ -1117,6 +1147,55 @@ mflux-generate-qwen-edit \
⚠️ *Note: The Qwen Image Edit model requires downloading the `Qwen/Qwen-Image-Edit-2509` model weights (~58GB for the full model, or use quantization for smaller sizes).*
+#### Qwen Image Layered
+
+**Qwen Image Layered** is a specialized image decomposition model that separates an input image into semantically disentangled RGBA layers. This enables powerful layer-based editing workflows where each layer can be independently manipulated and recomposed—similar to working with Photoshop layers but generated automatically from any image.
+
+The model uses a custom RGBA-VAE and a Layer3D RoPE transformer architecture to understand the semantic structure of images and separate them into distinct compositional layers (foreground, background, objects, etc.).
+
+**Example: Basic Image Decomposition**
+
+```sh
+mflux-generate-qwen-layered \
+ --image "input.png" \
+ --layers 4 \
+ --steps 50 \
+ --resolution 640 \
+ --guidance 4.0 \
+ --output-dir "./layers" \
+ -q 6
+```
+
+This will generate 4 RGBA layer files (`layer_0.png`, `layer_1.png`, etc.) in the output directory.
+
+**Example: Fast Preview with Fewer Steps**
+
+```sh
+mflux-generate-qwen-layered \
+ --image "photo.jpg" \
+ --layers 2 \
+ --steps 10 \
+ --resolution 640 \
+ -q 6 \
+ --output-dir "./preview"
+```
+
+**Use Cases:**
+- **Layer-based editing**: Edit individual layers independently and recompose
+- **Background removal**: Extract foreground objects with transparency
+- **Image compositing**: Combine layers from multiple decomposed images
+- **Animation**: Animate individual layers for parallax or motion effects
+- **Asset extraction**: Extract clean assets from complex scenes
+
+**Tips for Qwen Image Layered:**
+1. **Resolution**: Use 640 for faster processing, 1024 for higher quality results
+2. **Number of layers**: Start with 2-4 layers; more layers require more VRAM and time
+3. **Quantization**: 6-bit quantization (`-q 6`) significantly reduces memory usage (~29GB vs ~55GB BF16)
+4. **Steps**: 50 steps provides best quality; 10-20 steps work for quick previews
+5. **Output format**: All layers are saved as RGBA PNG files with transparency
+
+⚠️ *Note: The Qwen Image Layered model requires local weights from `Qwen/Qwen-Image-Layered` (~55GB for the full model in BF16, ~29GB with 6-bit quantization). This is a research model optimized for 48GB+ Apple Silicon Macs.*
+
---
### 🌀 FIBO
diff --git a/pyproject.toml b/pyproject.toml
index 90a2296..adf5b65 100644
--- a/pyproject.toml
+++ b/pyproject.toml
@@ -80,6 +80,7 @@ mflux-generate-redux = "mflux.models.flux.cli.flux_generate_redux:main"
mflux-generate-kontext = "mflux.models.flux.cli.flux_generate_kontext:main"
mflux-generate-qwen = "mflux.models.qwen.cli.qwen_image_generate:main"
mflux-generate-qwen-edit = "mflux.models.qwen.cli.qwen_image_edit_generate:main"
+mflux-generate-qwen-layered = "mflux.models.qwen_layered.cli.qwen_image_layered_generate:main"
mflux-generate-fibo = "mflux.models.fibo.cli.fibo_generate:main"
mflux-generate-z-image-turbo = "mflux.models.z_image.cli.z_image_turbo_generate:main"
mflux-refine-fibo = "mflux.models.fibo_vlm.cli.fibo_refine:main"
diff --git a/src/mflux/models/common/config/model_config.py b/src/mflux/models/common/config/model_config.py
index dc11af8..7f9caac 100644
--- a/src/mflux/models/common/config/model_config.py
+++ b/src/mflux/models/common/config/model_config.py
@@ -98,6 +98,11 @@ class ModelConfig:
def qwen_image_edit() -> "ModelConfig":
return AVAILABLE_MODELS["qwen-image-edit"]
+ @staticmethod
+ @lru_cache
+ def qwen_image_layered() -> "ModelConfig":
+ return AVAILABLE_MODELS["qwen-image-layered"]
+
@staticmethod
@lru_cache
def fibo() -> "ModelConfig":
@@ -308,4 +313,16 @@ AVAILABLE_MODELS = {
requires_sigma_shift=True,
priority=14,
),
+ "qwen-image-layered": ModelConfig(
+ aliases=["qwen-image-layered", "qwen-layered"],
+ model_name="Qwen/Qwen-Image-Layered",
+ base_model=None,
+ controlnet_model=None,
+ custom_transformer_model=None,
+ num_train_steps=None,
+ max_sequence_length=None,
+ supports_guidance=None,
+ requires_sigma_shift=None,
+ priority=15,
+ ),
}
diff --git a/src/mflux/models/qwen_layered/__init__.py b/src/mflux/models/qwen_layered/__init__.py
new file mode 100644
index 0000000..fae5f12
--- /dev/null
+++ b/src/mflux/models/qwen_layered/__init__.py
@@ -0,0 +1 @@
+# Qwen-Image-Layered model support for mflux
diff --git a/src/mflux/models/qwen_layered/cli/__init__.py b/src/mflux/models/qwen_layered/cli/__init__.py
new file mode 100644
index 0000000..1e184fe
--- /dev/null
+++ b/src/mflux/models/qwen_layered/cli/__init__.py
@@ -0,0 +1 @@
+# CLI module
diff --git a/src/mflux/models/qwen_layered/cli/qwen_image_layered_generate.py b/src/mflux/models/qwen_layered/cli/qwen_image_layered_generate.py
new file mode 100644
index 0000000..f1beb27
--- /dev/null
+++ b/src/mflux/models/qwen_layered/cli/qwen_image_layered_generate.py
@@ -0,0 +1,161 @@
+#!/usr/bin/env python
+"""CLI for Qwen-Image-Layered: Image decomposition into RGBA layers."""
+
+import argparse
+import sys
+from pathlib import Path
+
+
+def main():
+ parser = argparse.ArgumentParser(
+ description="Decompose an image into RGBA layers using Qwen-Image-Layered.",
+ formatter_class=argparse.RawDescriptionHelpFormatter,
+ epilog="""
+Examples:
+ # Basic decomposition into 4 layers
+ mflux-generate-qwen-layered --image input.png --layers 4
+
+ # With 6-bit quantization (recommended for 48GB Macs)
+ mflux-generate-qwen-layered --image input.png --layers 4 -q 6
+
+ # Custom settings
+ mflux-generate-qwen-layered --image input.png --layers 8 --steps 50 --resolution 640 -q 6
+ """,
+ )
+
+ # Required arguments
+ parser.add_argument(
+ "--image",
+ type=Path,
+ required=True,
+ help="Path to input image to decompose",
+ )
+
+ # Optional arguments
+ parser.add_argument(
+ "--layers",
+ type=int,
+ default=4,
+ help="Number of output layers (default: 4)",
+ )
+ parser.add_argument(
+ "--steps",
+ type=int,
+ default=50,
+ help="Number of inference steps (default: 50)",
+ )
+ parser.add_argument(
+ "--resolution",
+ type=int,
+ default=640,
+ choices=[640, 1024],
+ help="Target resolution bucket (default: 640)",
+ )
+ parser.add_argument(
+ "--guidance",
+ type=float,
+ default=4.0,
+ help="Guidance scale (default: 4.0)",
+ )
+ parser.add_argument(
+ "--seed",
+ type=int,
+ default=42,
+ help="Random seed (default: 42)",
+ )
+ parser.add_argument(
+ "--prompt",
+ type=str,
+ default=None,
+ help="Optional text prompt describing the image",
+ )
+ parser.add_argument(
+ "--negative-prompt",
+ type=str,
+ default=" ",
+ help="Optional negative prompt",
+ )
+ parser.add_argument(
+ "-q",
+ "--quantize",
+ type=int,
+ choices=[4, 6, 8],
+ default=None,
+ help="Quantization bits (4, 6, or 8). Recommended: 6 for 48GB Macs",
+ )
+ parser.add_argument(
+ "--model-path",
+ type=str,
+ default=None,
+ help="Path to local model weights (otherwise downloads from HuggingFace)",
+ )
+ parser.add_argument(
+ "--output",
+ type=str,
+ default="layer_{i}.png",
+ help="Output filename pattern. Use {i} for layer index (default: layer_{i}.png)",
+ )
+ parser.add_argument(
+ "--output-dir",
+ type=Path,
+ default=Path("."),
+ help="Output directory (default: current directory)",
+ )
+
+ args = parser.parse_args()
+
+ # Validate input
+ if not args.image.exists():
+ print(f"Error: Input image not found: {args.image}")
+ sys.exit(1)
+
+ # Import here to avoid slow startup for --help
+ from mflux.models.qwen_layered.variants.i2l.qwen_image_layered import QwenImageLayered
+
+ print("=" * 60)
+ print("Qwen-Image-Layered: Image Decomposition")
+ print("=" * 60)
+ print(f" Input: {args.image}")
+ print(f" Layers: {args.layers}")
+ print(f" Steps: {args.steps}")
+ print(f" Resolution: {args.resolution}")
+ print(f" Guidance: {args.guidance}")
+ print(f" Seed: {args.seed}")
+ print(f" Quantization: {args.quantize or 'BF16 (full precision)'}")
+ print("=" * 60)
+
+ # Initialize model
+ print("\nLoading model...")
+ model = QwenImageLayered(
+ quantize=args.quantize,
+ model_path=args.model_path,
+ )
+
+ # Run decomposition
+ print("\nStarting decomposition...")
+ layers = model.decompose(
+ seed=args.seed,
+ image_path=args.image,
+ num_layers=args.layers,
+ num_inference_steps=args.steps,
+ guidance=args.guidance,
+ resolution=args.resolution,
+ prompt=args.prompt,
+ negative_prompt=args.negative_prompt,
+ )
+
+ # Save output layers
+ args.output_dir.mkdir(parents=True, exist_ok=True)
+ print(f"\nSaving {len(layers)} layers to {args.output_dir}/")
+
+ for i, layer in enumerate(layers):
+ filename = args.output.format(i=i)
+ output_path = args.output_dir / filename
+ layer.save(output_path)
+ print(f" Saved {output_path}")
+
+ print("\nDone!")
+
+
+if __name__ == "__main__":
+ main()
diff --git a/src/mflux/models/qwen_layered/latent_creator/__init__.py b/src/mflux/models/qwen_layered/latent_creator/__init__.py
new file mode 100644
index 0000000..697dafa
--- /dev/null
+++ b/src/mflux/models/qwen_layered/latent_creator/__init__.py
@@ -0,0 +1 @@
+# Latent creator module
diff --git a/src/mflux/models/qwen_layered/latent_creator/qwen_layered_latent_creator.py b/src/mflux/models/qwen_layered/latent_creator/qwen_layered_latent_creator.py
new file mode 100644
index 0000000..368bd0a
--- /dev/null
+++ b/src/mflux/models/qwen_layered/latent_creator/qwen_layered_latent_creator.py
@@ -0,0 +1,208 @@
+import mlx.core as mx
+
+
+class QwenLayeredLatentCreator:
+ """
+ Latent creator for Qwen-Image-Layered.
+
+ Handles multi-layer latent creation, packing, and unpacking.
+ The layer dimension is treated as a temporal dimension in the 3D VAE.
+ """
+
+ @staticmethod
+ def create_noise(seed: int, num_layers: int, height: int, width: int) -> mx.array:
+ """
+ Create initial noise for N output layers.
+
+ Args:
+ seed: Random seed
+ num_layers: Number of output layers
+ height: Image height
+ width: Image width
+
+ Returns:
+ Noise tensor [1, N, 16, H/8, W/8]
+ """
+ mx.random.seed(seed)
+ latent_height = height // 8
+ latent_width = width // 8
+ num_channels = 16
+
+ # Create noise for all layers
+ noise = mx.random.normal(
+ shape=(1, num_layers, num_channels, latent_height, latent_width)
+ )
+ return noise
+
+ @staticmethod
+ def pack_latents(
+ latents: mx.array,
+ num_layers: int,
+ height: int,
+ width: int,
+ patch_size: int = 2,
+ ) -> mx.array:
+ """
+ Pack N-layer latents for transformer input.
+
+ Args:
+ latents: [B, N, 16, H, W] - multi-layer latents
+ num_layers: Number of layers
+ height: Image height
+ width: Image width
+ patch_size: Patch size for patchification (default 2)
+
+ Returns:
+ Packed latents [B, N*H*W/(patch_size^2), C] for transformer
+ """
+ batch_size = latents.shape[0]
+ num_channels = latents.shape[2]
+ latent_height = height // 8
+ latent_width = width // 8
+
+ # [B, N, 16, H, W] -> [B, N, H, W, 16]
+ latents = mx.transpose(latents, (0, 1, 3, 4, 2))
+
+ # Patchify: [B, N, H, W, 16] -> [B, N, H/p, W/p, p*p*16]
+ patched_height = latent_height // patch_size
+ patched_width = latent_width // patch_size
+
+ # Reshape for patchification
+ latents = latents.reshape(
+ batch_size,
+ num_layers,
+ patched_height,
+ patch_size,
+ patched_width,
+ patch_size,
+ num_channels,
+ )
+ # [B, N, H/p, p, W/p, p, C] -> [B, N, H/p, W/p, p, p, C]
+ latents = mx.transpose(latents, (0, 1, 2, 4, 3, 5, 6))
+ # [B, N, H/p, W/p, p*p*C]
+ latents = latents.reshape(
+ batch_size,
+ num_layers,
+ patched_height,
+ patched_width,
+ patch_size * patch_size * num_channels,
+ )
+
+ # Flatten spatial and layer dimensions: [B, N*H/p*W/p, C]
+ latents = latents.reshape(
+ batch_size,
+ num_layers * patched_height * patched_width,
+ patch_size * patch_size * num_channels,
+ )
+
+ return latents
+
+ @staticmethod
+ def unpack_latents(
+ latents: mx.array,
+ num_layers: int,
+ height: int,
+ width: int,
+ patch_size: int = 2,
+ out_channels: int = 16,
+ ) -> mx.array:
+ """
+ Unpack transformer output to N-layer latents.
+
+ Args:
+ latents: [B, N*H*W/(patch_size^2), p*p*out_channels] - packed latents
+ num_layers: Number of layers
+ height: Image height
+ width: Image width
+ patch_size: Patch size used in packing
+ out_channels: Number of output channels (default 16)
+
+ Returns:
+ Unpacked latents [B, N, out_channels, H/8, W/8]
+ """
+ batch_size = latents.shape[0]
+ latent_height = height // 8
+ latent_width = width // 8
+ patched_height = latent_height // patch_size
+ patched_width = latent_width // patch_size
+
+ # [B, N*H/p*W/p, p*p*C] -> [B, N, H/p, W/p, p*p*C]
+ latents = latents.reshape(
+ batch_size,
+ num_layers,
+ patched_height,
+ patched_width,
+ patch_size * patch_size * out_channels,
+ )
+
+ # Unpatchify: [B, N, H/p, W/p, p, p, C]
+ latents = latents.reshape(
+ batch_size,
+ num_layers,
+ patched_height,
+ patched_width,
+ patch_size,
+ patch_size,
+ out_channels,
+ )
+ # [B, N, H/p, W/p, p, p, C] -> [B, N, H/p, p, W/p, p, C]
+ latents = mx.transpose(latents, (0, 1, 2, 4, 3, 5, 6))
+ # [B, N, H, W, C]
+ latents = latents.reshape(
+ batch_size,
+ num_layers,
+ latent_height,
+ latent_width,
+ out_channels,
+ )
+ # [B, N, C, H, W]
+ latents = mx.transpose(latents, (0, 1, 4, 2, 3))
+
+ return latents
+
+ @staticmethod
+ def pack_condition_image(
+ latents: mx.array,
+ height: int,
+ width: int,
+ patch_size: int = 2,
+ ) -> mx.array:
+ """
+ Pack condition image latents (single layer) for transformer.
+
+ Args:
+ latents: [B, 16, H, W] - single layer latents
+ height: Image height
+ width: Image width
+ patch_size: Patch size
+
+ Returns:
+ Packed latents [B, H*W/(patch_size^2), p*p*16]
+ """
+ batch_size = latents.shape[0]
+ num_channels = latents.shape[1]
+ latent_height = height // 8
+ latent_width = width // 8
+ patched_height = latent_height // patch_size
+ patched_width = latent_width // patch_size
+
+ # [B, 16, H, W] -> [B, H, W, 16]
+ latents = mx.transpose(latents, (0, 2, 3, 1))
+
+ # Patchify
+ latents = latents.reshape(
+ batch_size,
+ patched_height,
+ patch_size,
+ patched_width,
+ patch_size,
+ num_channels,
+ )
+ latents = mx.transpose(latents, (0, 1, 3, 2, 4, 5))
+ latents = latents.reshape(
+ batch_size,
+ patched_height * patched_width,
+ patch_size * patch_size * num_channels,
+ )
+
+ return latents
diff --git a/src/mflux/models/qwen_layered/model/__init__.py b/src/mflux/models/qwen_layered/model/__init__.py
new file mode 100644
index 0000000..1a18dc5
--- /dev/null
+++ b/src/mflux/models/qwen_layered/model/__init__.py
@@ -0,0 +1 @@
+# Qwen-Layered VAE model
diff --git a/src/mflux/models/qwen_layered/model/qwen_layered_transformer/__init__.py b/src/mflux/models/qwen_layered/model/qwen_layered_transformer/__init__.py
new file mode 100644
index 0000000..d9ae873
--- /dev/null
+++ b/src/mflux/models/qwen_layered/model/qwen_layered_transformer/__init__.py
@@ -0,0 +1 @@
+# Qwen-Layered Transformer components
diff --git a/src/mflux/models/qwen_layered/model/qwen_layered_transformer/qwen_layered_rope.py b/src/mflux/models/qwen_layered/model/qwen_layered_transformer/qwen_layered_rope.py
new file mode 100644
index 0000000..d6abb8b
--- /dev/null
+++ b/src/mflux/models/qwen_layered/model/qwen_layered_transformer/qwen_layered_rope.py
@@ -0,0 +1,204 @@
+from __future__ import annotations
+
+import mlx.core as mx
+import numpy as np
+from mlx import nn
+
+
+class QwenLayeredRoPE(nn.Module):
+ """
+ Layer3D RoPE for Qwen-Image-Layered.
+
+ Extends the base RoPE to include a layer dimension:
+ - Position = (layer_index, height, width)
+ - Input condition image: layer_index = -1
+ - Output layers: layer_index = 0, 1, ..., N-1
+
+ axes_dim = [16, 56, 56] -> [layer, height, width]
+ """
+
+ def __init__(self, theta: int = 10000, axes_dim: list[int] = None, scale_rope: bool = True):
+ super().__init__()
+ if axes_dim is None:
+ axes_dim = [16, 56, 56]
+ self.theta = theta
+ self.axes_dim = axes_dim # [layer_dim, height_dim, width_dim]
+ self.scale_rope = scale_rope
+
+ # Pre-compute frequency tables for positive and negative indices
+ pos_index = np.arange(4096, dtype=np.int32)
+ neg_index = (np.arange(4096, dtype=np.int32)[::-1] * -1) - 1
+
+ # Frequency tables for each dimension
+ self.pos_freqs = np.concatenate(
+ [
+ self._rope_params(pos_index, self.axes_dim[0], self.theta), # Layer
+ self._rope_params(pos_index, self.axes_dim[1], self.theta), # Height
+ self._rope_params(pos_index, self.axes_dim[2], self.theta), # Width
+ ],
+ axis=1,
+ )
+ self.neg_freqs = np.concatenate(
+ [
+ self._rope_params(neg_index, self.axes_dim[0], self.theta),
+ self._rope_params(neg_index, self.axes_dim[1], self.theta),
+ self._rope_params(neg_index, self.axes_dim[2], self.theta),
+ ],
+ axis=1,
+ )
+
+ def _rope_params(self, index: np.ndarray, dim: int, theta: int) -> np.ndarray:
+ assert dim % 2 == 0
+ scales = np.arange(0, dim, 2, dtype=np.float32) / dim
+ omega = 1.0 / (theta**scales)
+ freqs = np.outer(index.astype(np.float32), omega)
+ cos_freqs = np.cos(freqs)
+ sin_freqs = np.sin(freqs)
+ return np.stack([cos_freqs, sin_freqs], axis=-1)
+
+ def _compute_layer_freqs(
+ self,
+ num_layers: int,
+ height: int,
+ width: int,
+ include_cond_image: bool = True,
+ ) -> tuple[np.ndarray, np.ndarray]:
+ """
+ Compute 3D positional frequencies for layered output.
+
+ Args:
+ num_layers: Number of output layers (N)
+ height: Latent height
+ width: Latent width
+ include_cond_image: Whether to include condition image at layer=-1
+ """
+ axes_splits = [x // 2 for x in self.axes_dim]
+ freqs_pos = np.split(self.pos_freqs, np.cumsum(axes_splits)[:-1], axis=1)
+ freqs_neg = np.split(self.neg_freqs, np.cumsum(axes_splits)[:-1], axis=1)
+
+ all_cos = []
+ all_sin = []
+
+ # Condition image at layer=-1 (if included)
+ if include_cond_image:
+ cond_layer_cos, cond_layer_sin = self._compute_single_layer_freqs(
+ layer_idx=-1,
+ height=height,
+ width=width,
+ freqs_pos=freqs_pos,
+ freqs_neg=freqs_neg,
+ )
+ all_cos.append(cond_layer_cos)
+ all_sin.append(cond_layer_sin)
+
+ # Output layers at layer=0..N-1
+ for layer_idx in range(num_layers):
+ layer_cos, layer_sin = self._compute_single_layer_freqs(
+ layer_idx=layer_idx,
+ height=height,
+ width=width,
+ freqs_pos=freqs_pos,
+ freqs_neg=freqs_neg,
+ )
+ all_cos.append(layer_cos)
+ all_sin.append(layer_sin)
+
+ # Concatenate all layers
+ img_cos = np.concatenate(all_cos, axis=0)
+ img_sin = np.concatenate(all_sin, axis=0)
+
+ return img_cos, img_sin
+
+ def _compute_single_layer_freqs(
+ self,
+ layer_idx: int,
+ height: int,
+ width: int,
+ freqs_pos: list[np.ndarray],
+ freqs_neg: list[np.ndarray],
+ ) -> tuple[np.ndarray, np.ndarray]:
+ """Compute frequencies for a single layer."""
+ seq_len = height * width
+
+ # Layer dimension - single index
+ if layer_idx >= 0:
+ freqs_layer = freqs_pos[0][layer_idx:layer_idx + 1]
+ else:
+ # Negative index for condition image
+ freqs_layer = freqs_neg[0][abs(layer_idx) - 1:abs(layer_idx)]
+ freqs_layer = np.broadcast_to(freqs_layer, (seq_len, freqs_layer.shape[-2], 2))
+
+ # Height dimension with optional scaling
+ if self.scale_rope:
+ freqs_height = np.concatenate(
+ [freqs_neg[1][-(height - height // 2):], freqs_pos[1][:height // 2]], axis=0
+ )
+ else:
+ freqs_height = freqs_pos[1][:height]
+ freqs_height = freqs_height.reshape(height, 1, -1, 2)
+ freqs_height = np.broadcast_to(freqs_height, (height, width, freqs_height.shape[-2], 2))
+ freqs_height = freqs_height.reshape(seq_len, -1, 2)
+
+ # Width dimension with optional scaling
+ if self.scale_rope:
+ freqs_width = np.concatenate(
+ [freqs_neg[2][-(width - width // 2):], freqs_pos[2][:width // 2]], axis=0
+ )
+ else:
+ freqs_width = freqs_pos[2][:width]
+ freqs_width = freqs_width.reshape(1, width, -1, 2)
+ freqs_width = np.broadcast_to(freqs_width, (height, width, freqs_width.shape[-2], 2))
+ freqs_width = freqs_width.reshape(seq_len, -1, 2)
+
+ # Concatenate all dimensions
+ freqs = np.concatenate([freqs_layer, freqs_height, freqs_width], axis=-2)
+
+ cos_freqs = freqs[..., 0]
+ sin_freqs = freqs[..., 1]
+
+ return cos_freqs, sin_freqs
+
+ def __call__(
+ self,
+ num_layers: int,
+ height: int,
+ width: int,
+ txt_seq_lens: list[int],
+ include_cond_image: bool = True,
+ ) -> tuple[tuple[mx.array, mx.array], tuple[mx.array, mx.array]]:
+ """
+ Compute rotary embeddings for layered output.
+
+ Args:
+ num_layers: Number of output layers
+ height: Latent height (H/16)
+ width: Latent width (W/16)
+ txt_seq_lens: List of text sequence lengths per batch
+ include_cond_image: Whether to include condition image embeddings
+
+ Returns:
+ Tuple of (image_rotary_emb, text_rotary_emb)
+ Each contains (cos, sin) arrays
+ """
+ # Compute image frequencies for all layers
+ img_cos, img_sin = self._compute_layer_freqs(
+ num_layers=num_layers,
+ height=height,
+ width=width,
+ include_cond_image=include_cond_image,
+ )
+
+ # Compute text frequencies
+ if self.scale_rope:
+ max_vid_index = max(height // 2, width // 2)
+ else:
+ max_vid_index = max(height, width)
+
+ max_len = max(txt_seq_lens)
+ txt_cos = self.pos_freqs[max_vid_index:max_vid_index + max_len, :, 0]
+ txt_sin = self.pos_freqs[max_vid_index:max_vid_index + max_len, :, 1]
+
+ return (
+ (mx.array(img_cos.astype(np.float32)), mx.array(img_sin.astype(np.float32))),
+ (mx.array(txt_cos.astype(np.float32)), mx.array(txt_sin.astype(np.float32))),
+ )
diff --git a/src/mflux/models/qwen_layered/model/qwen_layered_transformer/qwen_layered_transformer.py b/src/mflux/models/qwen_layered/model/qwen_layered_transformer/qwen_layered_transformer.py
new file mode 100644
index 0000000..b635297
--- /dev/null
+++ b/src/mflux/models/qwen_layered/model/qwen_layered_transformer/qwen_layered_transformer.py
@@ -0,0 +1,165 @@
+from __future__ import annotations
+
+import mlx.core as mx
+import numpy as np
+from mlx import nn
+
+from mflux.models.common.config.config import Config
+from mflux.models.flux.model.flux_transformer.ada_layer_norm_continuous import AdaLayerNormContinuous
+from mflux.models.qwen.model.qwen_transformer.qwen_time_text_embed import QwenTimeTextEmbed
+from mflux.models.qwen.model.qwen_transformer.qwen_transformer_block import QwenTransformerBlock
+from mflux.models.qwen.model.qwen_transformer.qwen_transformer_rms_norm import QwenTransformerRMSNorm
+from mflux.models.qwen_layered.model.qwen_layered_transformer.qwen_layered_rope import QwenLayeredRoPE
+
+
+class QwenLayeredTransformer(nn.Module):
+ """
+ VLD-MMDiT (Variable Layers Decomposition MMDiT) for Qwen-Image-Layered.
+
+ Key differences from base transformer:
+ 1. Uses Layer3D RoPE instead of standard 2D RoPE
+ 2. Accepts condition image latent (z_I) alongside noisy layers (x_t)
+ 3. Concatenates z_I and x_t along sequence dimension for joint attention
+ 4. Handles variable number of output layers N
+ """
+
+ def __init__(
+ self,
+ in_channels: int = 64,
+ out_channels: int = 16,
+ num_layers: int = 60,
+ attention_head_dim: int = 128,
+ num_attention_heads: int = 24,
+ joint_attention_dim: int = 3584,
+ patch_size: int = 2,
+ ) -> None:
+ super().__init__()
+ self.inner_dim = num_attention_heads * attention_head_dim
+ self.patch_size = patch_size
+ self.out_channels = out_channels
+
+ # Input projections
+ self.img_in = nn.Linear(in_channels, self.inner_dim)
+ self.cond_img_in = nn.Linear(in_channels, self.inner_dim) # For condition image
+
+ # Text processing
+ self.txt_norm = QwenTransformerRMSNorm(joint_attention_dim, eps=1e-6)
+ self.txt_in = nn.Linear(joint_attention_dim, self.inner_dim)
+
+ # Time embedding
+ self.time_text_embed = QwenTimeTextEmbed(timestep_proj_dim=256, inner_dim=self.inner_dim)
+
+ # Layer3D RoPE instead of standard RoPE
+ self.pos_embed = QwenLayeredRoPE(theta=10000, axes_dim=[16, 56, 56], scale_rope=True)
+
+ # Transformer blocks (same as base)
+ self.transformer_blocks = [
+ QwenTransformerBlock(dim=self.inner_dim, num_heads=num_attention_heads, head_dim=attention_head_dim)
+ for _ in range(num_layers)
+ ]
+
+ # Output
+ self.norm_out = AdaLayerNormContinuous(self.inner_dim, self.inner_dim)
+ self.proj_out = nn.Linear(self.inner_dim, patch_size * patch_size * out_channels)
+
+ def __call__(
+ self,
+ t: int,
+ config: Config,
+ hidden_states: mx.array,
+ cond_image_hidden_states: mx.array,
+ encoder_hidden_states: mx.array,
+ encoder_hidden_states_mask: mx.array,
+ num_output_layers: int,
+ ) -> mx.array:
+ """
+ Forward pass for layered decomposition.
+
+ Args:
+ t: Current timestep
+ config: Generation config
+ hidden_states: Noisy layer latents [B, N*seq_len, C]
+ cond_image_hidden_states: Condition image latents [B, seq_len, C]
+ encoder_hidden_states: Text embeddings [B, txt_len, C]
+ encoder_hidden_states_mask: Text attention mask
+ num_output_layers: Number of output layers N
+
+ Returns:
+ Predicted noise [B, N*seq_len, out_C]
+ """
+ batch_size = hidden_states.shape[0]
+
+ # Project inputs
+ hidden_states = self.img_in(hidden_states)
+ cond_image_hidden_states = self.cond_img_in(cond_image_hidden_states)
+
+ # Concatenate condition image and noisy layers along sequence dimension
+ # [B, seq + N*seq, C]
+ combined_hidden_states = mx.concatenate([cond_image_hidden_states, hidden_states], axis=1)
+
+ # Compute timestep
+ timestep = self._compute_timestep(t, config)
+ timestep = mx.broadcast_to(timestep, (batch_size,)).astype(combined_hidden_states.dtype)
+
+ # Process text
+ encoder_hidden_states = self.txt_norm(encoder_hidden_states)
+ encoder_hidden_states = self.txt_in(encoder_hidden_states)
+
+ # Time embedding
+ text_embeddings = self.time_text_embed(timestep, combined_hidden_states)
+
+ # Compute Layer3D RoPE
+ latent_height = config.height // 16
+ latent_width = config.width // 16
+ txt_seq_lens = [int(mx.sum(encoder_hidden_states_mask[i]).item()) for i in range(encoder_hidden_states_mask.shape[0])]
+
+ image_rotary_emb, txt_rotary_emb = self.pos_embed(
+ num_layers=num_output_layers,
+ height=latent_height,
+ width=latent_width,
+ txt_seq_lens=txt_seq_lens,
+ include_cond_image=True,
+ )
+
+ # Apply transformer blocks
+ for idx, block in enumerate(self.transformer_blocks):
+ encoder_hidden_states, combined_hidden_states = block(
+ hidden_states=combined_hidden_states,
+ encoder_hidden_states=encoder_hidden_states,
+ encoder_hidden_states_mask=encoder_hidden_states_mask,
+ text_embeddings=text_embeddings,
+ image_rotary_emb=(image_rotary_emb, txt_rotary_emb),
+ block_idx=idx,
+ )
+
+ # Split out the noisy layers (remove condition image)
+ cond_seq_len = cond_image_hidden_states.shape[1]
+ hidden_states = combined_hidden_states[:, cond_seq_len:, :]
+
+ # Output projection
+ hidden_states = self.norm_out(hidden_states, text_embeddings)
+ hidden_states = self.proj_out(hidden_states)
+
+ return hidden_states
+
+ @staticmethod
+ def _compute_timestep(t: int | float, config: Config) -> mx.array:
+ """Compute timestep value from step index."""
+ if isinstance(t, int):
+ if t < len(config.scheduler.sigmas):
+ time_step = config.scheduler.sigmas[t]
+ else:
+ timestep_idx = None
+ for idx, ts in enumerate(config.scheduler.timesteps):
+ if abs(int(ts.item()) - t) < 1:
+ timestep_idx = idx
+ break
+ if timestep_idx is None:
+ time_step = t / 1000.0
+ else:
+ time_step = config.scheduler.sigmas[timestep_idx]
+ else:
+ time_step = t
+
+ timestep = mx.array(np.full((1,), time_step, dtype=np.float32))
+ return timestep
diff --git a/src/mflux/models/qwen_layered/model/qwen_layered_vae/__init__.py b/src/mflux/models/qwen_layered/model/qwen_layered_vae/__init__.py
new file mode 100644
index 0000000..621e72e
--- /dev/null
+++ b/src/mflux/models/qwen_layered/model/qwen_layered_vae/__init__.py
@@ -0,0 +1 @@
+# Qwen-Layered VAE components
diff --git a/src/mflux/models/qwen_layered/model/qwen_layered_vae/qwen_layered_decoder_3d.py b/src/mflux/models/qwen_layered/model/qwen_layered_vae/qwen_layered_decoder_3d.py
new file mode 100644
index 0000000..8ca995c
--- /dev/null
+++ b/src/mflux/models/qwen_layered/model/qwen_layered_vae/qwen_layered_decoder_3d.py
@@ -0,0 +1,39 @@
+import mlx.core as mx
+from mlx import nn
+
+from mflux.models.qwen.model.qwen_vae.qwen_image_causal_conv_3d import QwenImageCausalConv3D
+from mflux.models.qwen.model.qwen_vae.qwen_image_mid_block_3d import QwenImageMidBlock3D
+from mflux.models.qwen.model.qwen_vae.qwen_image_rms_norm import QwenImageRMSNorm
+from mflux.models.qwen.model.qwen_vae.qwen_image_up_block_3d import QwenImageUpBlock3D
+
+
+class QwenLayeredDecoder3D(nn.Module):
+ """
+ RGBA-VAE Decoder with 4-channel output for Qwen-Image-Layered.
+ Decodes latents into RGBA images with alpha channel.
+ """
+
+ def __init__(self, output_channels: int = 4):
+ super().__init__()
+ self.output_channels = output_channels
+ self.conv_in = QwenImageCausalConv3D(16, 384, 3, 1, 1)
+ self.mid_block = QwenImageMidBlock3D(384, num_layers=1)
+ self.up_block0 = QwenImageUpBlock3D(384, 384, num_res_blocks=2, upsample_mode="upsample3d")
+ self.up_block1 = QwenImageUpBlock3D(192, 384, num_res_blocks=2, upsample_mode="upsample3d")
+ self.up_block2 = QwenImageUpBlock3D(192, 192, num_res_blocks=2, upsample_mode="upsample2d")
+ self.up_block3 = QwenImageUpBlock3D(96, 96, num_res_blocks=2, upsample_mode=None)
+ self.norm_out = QwenImageRMSNorm(96, images=False)
+ # 4-channel output for RGBA
+ self.conv_out = QwenImageCausalConv3D(96, output_channels, 3, 1, 1)
+
+ def __call__(self, x: mx.array) -> mx.array:
+ x = self.conv_in(x)
+ x = self.mid_block(x)
+ x = self.up_block0(x)
+ x = self.up_block1(x)
+ x = self.up_block2(x)
+ x = self.up_block3(x)
+ x = self.norm_out(x)
+ x = nn.silu(x)
+ x = self.conv_out(x)
+ return x
diff --git a/src/mflux/models/qwen_layered/model/qwen_layered_vae/qwen_layered_encoder_3d.py b/src/mflux/models/qwen_layered/model/qwen_layered_vae/qwen_layered_encoder_3d.py
new file mode 100644
index 0000000..85e2898
--- /dev/null
+++ b/src/mflux/models/qwen_layered/model/qwen_layered_vae/qwen_layered_encoder_3d.py
@@ -0,0 +1,70 @@
+import mlx.core as mx
+from mlx import nn
+
+from mflux.models.qwen.model.qwen_vae.qwen_image_causal_conv_3d import QwenImageCausalConv3D
+from mflux.models.qwen.model.qwen_vae.qwen_image_down_block_3d import QwenImageDownBlock3D
+from mflux.models.qwen.model.qwen_vae.qwen_image_mid_block_3d import QwenImageMidBlock3D
+from mflux.models.qwen.model.qwen_vae.qwen_image_rms_norm import QwenImageRMSNorm
+
+
+class QwenLayeredEncoder3D(nn.Module):
+ """
+ RGBA-VAE Encoder with 4-channel input for Qwen-Image-Layered.
+ Encodes RGBA images (or RGB with alpha=1) into latent space.
+ """
+
+ def __init__(self, input_channels: int = 4):
+ super().__init__()
+ self.input_channels = input_channels
+ self.dim = 96
+ self.z_dim = 32
+ self.dim_mult = [1, 2, 4, 4]
+ self.num_res_blocks = [2, 2, 2, 2]
+ self.attn_scales = []
+ self.temporal_downsample = [False, False, True, True]
+ self.dropout = 0.0
+
+ dims = [self.dim * u for u in [1] + self.dim_mult]
+ # 4-channel input for RGBA
+ self.conv_in = QwenImageCausalConv3D(input_channels, dims[0], 3, 1, 1)
+
+ down_blocks = []
+ for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])):
+ downsample_mode = "downsample3d" if self.temporal_downsample[i] else "downsample2d"
+ if i == len(dims) - 2:
+ downsample_mode = None
+ stage_res_blocks = self.num_res_blocks[i] if isinstance(self.num_res_blocks, list) else self.num_res_blocks
+ down_block = QwenImageDownBlock3D(
+ in_dim, out_dim, num_res_blocks=stage_res_blocks, downsample_mode=downsample_mode
+ )
+ down_blocks.append(down_block)
+ self.down_blocks = down_blocks
+
+ self.mid_block = QwenImageMidBlock3D(dims[-1], num_layers=1)
+ self.norm_out = QwenImageRMSNorm(dims[-1], images=False)
+ self.conv_out = QwenImageCausalConv3D(dims[-1], 32, 3, 1, 1)
+
+ def __call__(self, x: mx.array) -> mx.array:
+ x = self.conv_in(x)
+ for stage_idx, down_block in enumerate(self.down_blocks):
+ if stage_idx == 3:
+ for resnet in down_block.resnets:
+ residual = x
+ n1 = resnet.norm1(x)
+ a1 = nn.silu(n1)
+ c1 = resnet.conv1(a1)
+ n2 = resnet.norm2(c1)
+ a2 = nn.silu(n2)
+ c2 = resnet.conv2(a2)
+ x = c2 + residual
+ if down_block.downsamplers is not None:
+ x = down_block.downsamplers[0](x)
+ else:
+ x = down_block(x)
+
+ x = self.mid_block(x)
+ norm_in = x
+ x = self.norm_out(norm_in)
+ x = nn.silu(x)
+ encoded = self.conv_out(x)
+ return encoded
diff --git a/src/mflux/models/qwen_layered/model/qwen_layered_vae/qwen_layered_vae.py b/src/mflux/models/qwen_layered/model/qwen_layered_vae/qwen_layered_vae.py
new file mode 100644
index 0000000..61e590b
--- /dev/null
+++ b/src/mflux/models/qwen_layered/model/qwen_layered_vae/qwen_layered_vae.py
@@ -0,0 +1,113 @@
+import mlx.core as mx
+from mlx import nn
+
+from mflux.models.qwen.model.qwen_vae.qwen_image_causal_conv_3d import QwenImageCausalConv3D
+from mflux.models.qwen_layered.model.qwen_layered_vae.qwen_layered_decoder_3d import QwenLayeredDecoder3D
+from mflux.models.qwen_layered.model.qwen_layered_vae.qwen_layered_encoder_3d import QwenLayeredEncoder3D
+
+
+class QwenLayeredVAE(nn.Module):
+ """
+ RGBA-VAE for Qwen-Image-Layered.
+
+ Handles 4-channel RGBA images and supports multi-layer encoding/decoding
+ where the temporal dimension is used as the layer dimension.
+
+ - Encoder: RGBA [B, 4, N, H, W] -> Latent [B, 16, N, H/8, W/8]
+ - Decoder: Latent [B, 16, N, H/8, W/8] -> RGBA [B, 4, N, H, W]
+ """
+
+ # Same latent normalization as base Qwen-Image
+ LATENTS_MEAN = mx.array([-0.7571, -0.7089, -0.9113, 0.1075, -0.1745, 0.9653, -0.1517, 1.5508, 0.4134, -0.0715, 0.5517, -0.3632, -0.1922, -0.9497, 0.2503, -0.2921]).reshape(1, 16, 1, 1, 1) # fmt: off
+ LATENTS_STD = mx.array([2.8184, 1.4541, 2.3275, 2.6558, 1.2196, 1.7708, 2.6052, 2.0743, 3.2687, 2.1526, 2.8652, 1.5579, 1.6382, 1.1253, 2.8251, 1.916]).reshape(1, 16, 1, 1, 1) # fmt: off
+
+ def __init__(self, input_channels: int = 4, output_channels: int = 4):
+ super().__init__()
+ self.input_channels = input_channels
+ self.output_channels = output_channels
+ self.decoder = QwenLayeredDecoder3D(output_channels=output_channels)
+ self.encoder = QwenLayeredEncoder3D(input_channels=input_channels)
+ self.post_quant_conv = QwenImageCausalConv3D(16, 16, 1, 1, 0)
+ self.quant_conv = QwenImageCausalConv3D(32, 32, 1, 1, 0)
+
+ def decode(self, latents: mx.array, num_layers: int = 1) -> mx.array:
+ """
+ Decode latents to RGBA images.
+
+ Args:
+ latents: [B, 16, H, W] for single layer or [B, N, 16, H, W] for multi-layer
+ num_layers: Number of output layers
+
+ Returns:
+ RGBA images [B, 4, H, W] for single layer or [B, N, 4, H, W] for multi-layer
+ """
+ # Handle both single-layer and multi-layer cases
+ if latents.ndim == 4:
+ # Single layer: [B, 16, H, W] -> [B, 16, 1, H, W]
+ latents = latents.reshape(latents.shape[0], latents.shape[1], 1, latents.shape[2], latents.shape[3])
+ elif latents.ndim == 5 and latents.shape[2] != num_layers:
+ # Multi-layer packed: [B, N, 16, H, W] -> [B, 16, N, H, W]
+ latents = mx.transpose(latents, (0, 2, 1, 3, 4))
+
+ # Denormalize
+ latents = latents * QwenLayeredVAE.LATENTS_STD + QwenLayeredVAE.LATENTS_MEAN
+ latents = self.post_quant_conv(latents)
+ decoded = self.decoder(latents)
+
+ if num_layers == 1:
+ # Single layer output: [B, 4, 1, H, W] -> [B, 4, H, W]
+ return decoded[:, :, 0, :, :]
+ else:
+ # Multi-layer output: [B, 4, N, H, W] -> [B, N, 4, H, W]
+ return mx.transpose(decoded, (0, 2, 1, 3, 4))
+
+ def encode(self, images: mx.array) -> mx.array:
+ """
+ Encode RGBA images to latents.
+
+ Args:
+ images: RGBA images [B, 4, H, W] or [B, N, 4, H, W] for multi-layer
+
+ Returns:
+ Latents [B, 16, H/8, W/8] or [B, N, 16, H/8, W/8] for multi-layer
+ """
+ is_multi_layer = images.ndim == 5
+
+ if images.ndim == 4:
+ # Single image: [B, 4, H, W] -> [B, 4, 1, H, W]
+ images = images.reshape(images.shape[0], images.shape[1], 1, images.shape[2], images.shape[3])
+ else:
+ # Multi-layer: [B, N, 4, H, W] -> [B, 4, N, H, W]
+ images = mx.transpose(images, (0, 2, 1, 3, 4))
+
+ latents = self.encoder(images)
+ latents = self.quant_conv(latents)
+ latents = latents[:, :16, :, :, :]
+
+ # Normalize
+ latents = (latents - QwenLayeredVAE.LATENTS_MEAN) / QwenLayeredVAE.LATENTS_STD
+
+ if not is_multi_layer:
+ # Single layer: [B, 16, 1, H, W] -> [B, 16, H, W]
+ return latents[:, :, 0, :, :]
+ else:
+ # Multi-layer: [B, 16, N, H, W] -> [B, N, 16, H, W]
+ return mx.transpose(latents, (0, 2, 1, 3, 4))
+
+ def encode_condition_image(self, image: mx.array) -> mx.array:
+ """
+ Encode the input RGB condition image (converted to RGBA with alpha=1).
+
+ Args:
+ image: RGB image [B, 3, H, W]
+
+ Returns:
+ Latent [B, 16, H/8, W/8]
+ """
+ # Convert RGB to RGBA by adding alpha=1 channel
+ batch_size = image.shape[0]
+ height, width = image.shape[2], image.shape[3]
+ alpha_channel = mx.ones((batch_size, 1, height, width), dtype=image.dtype)
+ rgba_image = mx.concatenate([image, alpha_channel], axis=1)
+
+ return self.encode(rgba_image)
diff --git a/src/mflux/models/qwen_layered/qwen_layered_initializer.py b/src/mflux/models/qwen_layered/qwen_layered_initializer.py
new file mode 100644
index 0000000..9a899f5
--- /dev/null
+++ b/src/mflux/models/qwen_layered/qwen_layered_initializer.py
@@ -0,0 +1,82 @@
+import mlx.core as mx
+
+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.qwen.model.qwen_text_encoder.qwen_text_encoder import QwenTextEncoder
+from mflux.models.qwen.model.qwen_transformer.qwen_transformer import QwenTransformer # Use base transformer
+from mflux.models.qwen_layered.model.qwen_layered_vae.qwen_layered_vae import QwenLayeredVAE
+from mflux.models.qwen_layered.weights.qwen_layered_weight_definition import QwenLayeredWeightDefinition
+
+
+class QwenLayeredInitializer:
+ """Initializer for Qwen-Image-Layered model."""
+
+ @staticmethod
+ def init(
+ model,
+ model_config: ModelConfig,
+ quantize: int | None,
+ model_path: str | None = None,
+ lora_paths: list[str] | None = None,
+ lora_scales: list[float] | None = None,
+ ) -> None:
+ """Initialize the Qwen-Image-Layered model."""
+ path = model_path if model_path else model_config.model_name
+
+ QwenLayeredInitializer._init_config(model, model_config)
+ weights = QwenLayeredInitializer._load_weights(path)
+ QwenLayeredInitializer._init_tokenizers(model, path)
+ QwenLayeredInitializer._init_models(model)
+ QwenLayeredInitializer._apply_weights(model, weights, quantize)
+ # Note: LoRA not supported yet for layered model
+
+ @staticmethod
+ def _init_config(model, model_config: ModelConfig) -> None:
+ model.model_config = model_config
+ model.lora_paths = None
+ model.lora_scales = None
+ model.prompt_cache = {}
+ model.callbacks = CallbackRegistry()
+ model.bits = None
+
+ @staticmethod
+ def _load_weights(path: str) -> LoadedWeights:
+ return WeightLoader.load(
+ weight_definition=QwenLayeredWeightDefinition,
+ model_path=path,
+ )
+
+ @staticmethod
+ def _init_tokenizers(model, path: str) -> None:
+ model.tokenizers = TokenizerLoader.load_all(
+ definitions=QwenLayeredWeightDefinition.get_tokenizers(),
+ model_path=path,
+ )
+
+ @staticmethod
+ def _init_models(model) -> None:
+ """Initialize model components."""
+ model.vae = QwenLayeredVAE(input_channels=4, output_channels=4)
+ model.transformer = QwenTransformer() # Use base transformer
+ model.text_encoder = QwenTextEncoder()
+
+ @staticmethod
+ def _apply_weights(model, weights: LoadedWeights, quantize: int | None) -> None:
+ """Apply weights and optionally quantize."""
+ model.bits = WeightApplier.apply_and_quantize(
+ weights=weights,
+ quantize_arg=quantize,
+ weight_definition=QwenLayeredWeightDefinition,
+ models={
+ "vae": model.vae,
+ "transformer": model.transformer,
+ "text_encoder": model.text_encoder,
+ },
+ )
+
+ # Evaluate to load weights into memory
+ mx.eval(model.parameters())
diff --git a/src/mflux/models/qwen_layered/variants/__init__.py b/src/mflux/models/qwen_layered/variants/__init__.py
new file mode 100644
index 0000000..cb64ec5
--- /dev/null
+++ b/src/mflux/models/qwen_layered/variants/__init__.py
@@ -0,0 +1 @@
+# Variants module
diff --git a/src/mflux/models/qwen_layered/variants/i2l/__init__.py b/src/mflux/models/qwen_layered/variants/i2l/__init__.py
new file mode 100644
index 0000000..64ef279
--- /dev/null
+++ b/src/mflux/models/qwen_layered/variants/i2l/__init__.py
@@ -0,0 +1 @@
+# I2L (Image to Layers) module
diff --git a/src/mflux/models/qwen_layered/variants/i2l/qwen_image_layered.py b/src/mflux/models/qwen_layered/variants/i2l/qwen_image_layered.py
new file mode 100644
index 0000000..7cb39a4
--- /dev/null
+++ b/src/mflux/models/qwen_layered/variants/i2l/qwen_image_layered.py
@@ -0,0 +1,358 @@
+from pathlib import Path
+from typing import List
+
+import mlx.core as mx
+import numpy as np
+from mlx import nn
+from PIL import Image
+
+from mflux.models.common.config import ModelConfig
+from mflux.models.common.config.config import Config
+from mflux.models.common.weights.saving.model_saver import ModelSaver
+from mflux.models.qwen.model.qwen_text_encoder.qwen_prompt_encoder import QwenPromptEncoder
+from mflux.models.qwen.model.qwen_text_encoder.qwen_text_encoder import QwenTextEncoder
+from mflux.models.qwen.model.qwen_transformer.qwen_transformer import QwenTransformer # Use base transformer!
+from mflux.models.qwen_layered.model.qwen_layered_vae.qwen_layered_vae import QwenLayeredVAE
+from mflux.models.qwen_layered.qwen_layered_initializer import QwenLayeredInitializer
+from mflux.models.qwen_layered.weights.qwen_layered_weight_definition import QwenLayeredWeightDefinition
+from mflux.utils.exceptions import StopImageGenerationException
+
+
+class QwenImageLayered(nn.Module):
+ """
+ Qwen-Image-Layered model for image decomposition into RGBA layers.
+
+ Takes an input RGB image and decomposes it into N semantically disentangled
+ RGBA layers that can be independently edited and composited back together.
+ """
+
+ vae: QwenLayeredVAE
+ transformer: QwenTransformer # Use base Qwen transformer!
+ text_encoder: QwenTextEncoder
+
+ def __init__(
+ self,
+ quantize: int | None = None,
+ model_path: str | None = None,
+ lora_paths: list[str] | None = None,
+ lora_scales: list[float] | None = None,
+ model_config: ModelConfig = None,
+ ):
+ super().__init__()
+ if model_config is None:
+ model_config = ModelConfig.from_name("qwen-image-layered")
+
+ QwenLayeredInitializer.init(
+ model=self,
+ quantize=quantize,
+ model_path=model_path,
+ lora_paths=lora_paths,
+ lora_scales=lora_scales,
+ model_config=model_config,
+ )
+
+ def decompose(
+ self,
+ seed: int,
+ image_path: Path | str,
+ num_layers: int = 4,
+ num_inference_steps: int = 50,
+ guidance: float = 4.0,
+ resolution: int = 640,
+ prompt: str | None = None,
+ negative_prompt: str | None = None,
+ scheduler: str = "linear",
+ cfg_normalize: bool = True,
+ ) -> List[Image.Image]:
+ """
+ Decompose an input image into N RGBA layers.
+ """
+ # Load and preprocess input image
+ input_image = Image.open(image_path).convert("RGBA")
+
+ # Resize to target resolution while maintaining aspect ratio
+ width, height = self._compute_resolution(input_image.size, resolution)
+ input_image = input_image.resize((width, height), Image.Resampling.LANCZOS)
+
+ # Create config
+ config = Config(
+ width=width,
+ height=height,
+ guidance=guidance,
+ scheduler=scheduler,
+ model_config=self.model_config,
+ num_inference_steps=num_inference_steps,
+ )
+
+ # Encode input image to latent (RGB only)
+ input_tensor = self._image_to_tensor(input_image)
+ # Encode only RGB (first 3 channels) -> [B, 16, H/8, W/8]
+ cond_latent = self.vae.encode_condition_image(input_tensor[:, :3, :, :])
+
+ # Add layer dimension: [B, C, H, W] -> [B, C, 1, H, W] -> [B, 1, C, H, W]
+ cond_latent = mx.expand_dims(cond_latent, axis=2) # [B, C, 1, H, W]
+ cond_latent = cond_latent.transpose(0, 2, 1, 3, 4) # [B, 1, C, H, W]
+
+ # Pack condition image latent
+ image_latents = self._pack_latents(
+ cond_latent,
+ batch_size=1,
+ num_layers=1, # Single condition image
+ height=height,
+ width=width,
+ )
+
+ # Create initial noise for output layers (layers+1 to include combined)
+ mx.random.seed(seed)
+ noise = self._create_noise(
+ seed=seed,
+ num_layers=num_layers + 1, # +1 for combined output
+ height=height,
+ width=width,
+ )
+ latents = self._pack_latents(
+ noise,
+ batch_size=1,
+ num_layers=num_layers + 1,
+ height=height,
+ width=width,
+ )
+
+ # Encode prompt
+ if prompt is None or prompt == "":
+ prompt = "an image" # Placeholder - ideally auto-caption
+ if negative_prompt is None:
+ negative_prompt = "blurry, bad quality"
+
+ prompt_embeds, prompt_mask, neg_embeds, neg_mask = QwenPromptEncoder.encode_prompt(
+ prompt=prompt,
+ negative_prompt=negative_prompt,
+ prompt_cache=self.prompt_cache,
+ qwen_tokenizer=self.tokenizers["qwen"],
+ qwen_text_encoder=self.text_encoder,
+ )
+
+ # Calculate latent dimensions and img_shapes for RoPE
+ latent_height = height // 16 # VAE compression / patch
+ latent_width = width // 16
+
+ # Build cond_image_grid: (num_layers+1 output) + 1 condition
+ # Each layer is: (1 frame, latent_height, latent_width)
+ # For layered model, we need to generate RoPE for ALL layers + condition
+ # The base transformer will compute shapes for: [(1, H, W)] + cond_image_grid
+ # So we pass (num_layers+1 - 1) additional grids for the noisy layers,
+ # plus 1 for condition = num_layers + 1 additional grids
+ cond_image_grid = [(1, latent_height, latent_width) for _ in range(num_layers + 1)]
+
+ # Denoising loop
+ print(f"Decomposing into {num_layers} layers...")
+ try:
+ for step_idx, t in enumerate(config.time_steps):
+ # Scale model input
+ latents = config.scheduler.scale_model_input(latents, t)
+
+ # KEY DIFFERENCE: Concatenate noisy latents with condition image
+ # Like Diffusers: latent_model_input = torch.cat([latents, image_latents], dim=1)
+ latent_model_input = mx.concatenate([latents, image_latents], axis=1)
+
+ # Predict noise with positive prompt using BASE transformer
+ noise_pred = self.transformer(
+ t=t,
+ config=config,
+ hidden_states=latent_model_input,
+ encoder_hidden_states=prompt_embeds,
+ encoder_hidden_states_mask=prompt_mask,
+ cond_image_grid=cond_image_grid,
+ )
+ # Only take the first part (excludes condition image)
+ noise_pred = noise_pred[:, : latents.shape[1], :]
+
+ # Predict noise with negative prompt
+ noise_pred_neg = self.transformer(
+ t=t,
+ config=config,
+ hidden_states=mx.concatenate([latents, image_latents], axis=1),
+ encoder_hidden_states=neg_embeds,
+ encoder_hidden_states_mask=neg_mask,
+ cond_image_grid=cond_image_grid,
+ )
+ noise_pred_neg = noise_pred_neg[:, : latents.shape[1], :]
+
+ # Apply CFG
+ guided_noise = self._compute_guided_noise(noise_pred, noise_pred_neg, guidance, cfg_normalize)
+
+ # Scheduler step
+ latents = config.scheduler.step(noise=guided_noise, timestep=t, latents=latents)
+
+ # Evaluate for progress
+ mx.eval(latents)
+
+ if (step_idx + 1) % 10 == 0 or step_idx == 0:
+ print(f" Step {step_idx + 1}/{num_inference_steps}")
+ except KeyboardInterrupt:
+ raise StopImageGenerationException(
+ f"Stopping decomposition at step {step_idx + 1}/{num_inference_steps}"
+ ) from None
+
+ # Unpack latents
+ latents = self._unpack_latents(
+ latents,
+ num_layers=num_layers + 1, # +1 for combined
+ height=height,
+ width=width,
+ )
+
+ # Skip first frame (combined image) - like Diffusers line 886
+ # latents[:, :, 1:] - skip first layer
+ latents = latents[:, 1:, :, :, :] # Shape: [B, layers, C, H, W]
+
+ # Decode each layer
+ output_images = []
+ print(" Decoding layers...")
+ for layer_idx in range(num_layers):
+ layer_latent = latents[:, layer_idx : layer_idx + 1, :, :, :] # [B, 1, C, H, W]
+ layer_latent = layer_latent[:, 0, :, :, :] # [B, C, H, W]
+ decoded = self.vae.decode(layer_latent, num_layers=1)
+ rgba_image = self._tensor_to_image(decoded)
+ output_images.append(rgba_image)
+
+ print(f"Decomposition complete: {num_layers} layers")
+ return output_images
+
+ def _compute_resolution(self, original_size: tuple, target_bucket: int) -> tuple:
+ """Compute target resolution maintaining aspect ratio."""
+ w, h = original_size
+ aspect = w / h
+
+ if aspect >= 1.0:
+ # Landscape or square
+ new_w = target_bucket
+ new_h = int(target_bucket / aspect)
+ else:
+ # Portrait
+ new_h = target_bucket
+ new_w = int(target_bucket * aspect)
+
+ # Round to nearest multiple of 16 for VAE compatibility
+ new_w = (new_w // 16) * 16
+ new_h = (new_h // 16) * 16
+
+ return max(new_w, 16), max(new_h, 16)
+
+ def _create_noise(self, seed: int, num_layers: int, height: int, width: int) -> mx.array:
+ """Create initial noise for output layers."""
+ latent_height = height // 8 # VAE compression
+ latent_width = width // 8
+ num_channels = 16 # Qwen VAE latent channels
+
+ mx.random.seed(seed)
+ noise = mx.random.normal(shape=(1, num_layers, num_channels, latent_height, latent_width))
+ return noise.astype(mx.bfloat16)
+
+ def _pack_latents(
+ self,
+ latents: mx.array,
+ batch_size: int,
+ num_layers: int,
+ height: int,
+ width: int,
+ ) -> mx.array:
+ """
+ Pack latents for transformer input.
+
+ From Diffusers:
+ latents = latents.view(batch_size, layers, num_channels_latents, height // 2, 2, width // 2, 2)
+ latents = latents.permute(0, 1, 3, 5, 2, 4, 6)
+ latents = latents.reshape(batch_size, layers * (height // 2) * (width // 2), num_channels_latents * 4)
+ """
+ # latents: [B, layers, C, H, W]
+ latent_height = height // 8 // 2 # VAE + patch
+ latent_width = width // 8 // 2
+ num_channels = latents.shape[2]
+
+ # Reshape: [B, layers, C, H/2, 2, W/2, 2]
+ latents = latents.reshape(batch_size, num_layers, num_channels, latent_height, 2, latent_width, 2)
+
+ # Permute: [B, layers, H/2, W/2, C, 2, 2]
+ latents = latents.transpose(0, 1, 3, 5, 2, 4, 6)
+
+ # Reshape: [B, layers * H/2 * W/2, C * 4]
+ latents = latents.reshape(batch_size, num_layers * latent_height * latent_width, num_channels * 4)
+
+ return latents
+
+ def _unpack_latents(
+ self,
+ latents: mx.array,
+ num_layers: int,
+ height: int,
+ width: int,
+ ) -> mx.array:
+ """
+ Unpack latents after transformer.
+
+ From Diffusers:
+ latents = latents.view(batch_size, layers + 1, height // 2, width // 2, channels // 4, 2, 2)
+ latents = latents.permute(0, 1, 4, 2, 5, 3, 6)
+ latents = latents.reshape(batch_size, layers + 1, channels // (2 * 2), height, width)
+ """
+ batch_size = latents.shape[0]
+ channels = latents.shape[2]
+
+ latent_height = height // 8 // 2 # VAE + patch
+ latent_width = width // 8 // 2
+
+ # Reshape: [B, layers, H/2, W/2, C/4, 2, 2]
+ latents = latents.reshape(batch_size, num_layers, latent_height, latent_width, channels // 4, 2, 2)
+
+ # Permute: [B, layers, C/4, H/2, 2, W/2, 2]
+ latents = latents.transpose(0, 1, 4, 2, 5, 3, 6)
+
+ # Reshape: [B, layers, C, H, W]
+ full_height = latent_height * 2
+ full_width = latent_width * 2
+ latents = latents.reshape(batch_size, num_layers, channels // 4, full_height, full_width)
+
+ return latents
+
+ def _image_to_tensor(self, image: Image.Image) -> mx.array:
+ """Convert PIL RGBA image to tensor [1, 4, H, W] in [-1, 1]."""
+ arr = np.array(image).astype(np.float32) / 255.0
+ arr = arr * 2.0 - 1.0 # [0, 1] -> [-1, 1]
+ arr = np.transpose(arr, (2, 0, 1)) # [H, W, C] -> [C, H, W]
+ arr = np.expand_dims(arr, 0) # [1, C, H, W]
+ return mx.array(arr)
+
+ def _tensor_to_image(self, tensor: mx.array) -> Image.Image:
+ """Convert tensor [1, 4, H, W] in [-1, 1] to PIL RGBA image."""
+ arr = np.array(tensor[0]) # [4, H, W]
+ arr = np.transpose(arr, (1, 2, 0)) # [H, W, 4]
+ arr = (arr + 1.0) / 2.0 # [-1, 1] -> [0, 1]
+ arr = np.clip(arr * 255.0, 0, 255).astype(np.uint8)
+ return Image.fromarray(arr, mode="RGBA")
+
+ @staticmethod
+ def _compute_guided_noise(
+ noise: mx.array,
+ noise_neg: mx.array,
+ guidance: float,
+ normalize: bool = True,
+ ) -> mx.array:
+ """Apply classifier-free guidance with optional normalization."""
+ combined = noise_neg + guidance * (noise - noise_neg)
+
+ if normalize:
+ cond_norm = mx.sqrt(mx.sum(noise * noise, axis=-1, keepdims=True) + 1e-12)
+ combined_norm = mx.sqrt(mx.sum(combined * combined, axis=-1, keepdims=True) + 1e-12)
+ combined = combined * (cond_norm / combined_norm)
+
+ return combined
+
+ def save_model(self, base_path: str) -> None:
+ """Save the model with current quantization."""
+ ModelSaver.save_model(
+ model=self,
+ bits=self.bits,
+ base_path=base_path,
+ weight_definition=QwenLayeredWeightDefinition,
+ )
diff --git a/src/mflux/models/qwen_layered/weights/__init__.py b/src/mflux/models/qwen_layered/weights/__init__.py
new file mode 100644
index 0000000..2a4148a
--- /dev/null
+++ b/src/mflux/models/qwen_layered/weights/__init__.py
@@ -0,0 +1 @@
+# Weights module
diff --git a/src/mflux/models/qwen_layered/weights/qwen_layered_weight_definition.py b/src/mflux/models/qwen_layered/weights/qwen_layered_weight_definition.py
new file mode 100644
index 0000000..0f4f228
--- /dev/null
+++ b/src/mflux/models/qwen_layered/weights/qwen_layered_weight_definition.py
@@ -0,0 +1,78 @@
+from typing import List
+
+import mlx.core as mx
+
+from mflux.models.common.tokenizer import LanguageTokenizer
+from mflux.models.common.weights.loading.weight_definition import ComponentDefinition, TokenizerDefinition
+from mflux.models.qwen_layered.weights.qwen_layered_weight_mapping import QwenLayeredWeightMapping
+
+
+class QwenLayeredWeightDefinition:
+ """
+ Weight definition for Qwen-Image-Layered model.
+
+ Components:
+ - VAE: RGBA-VAE with 4-channel I/O
+ - Transformer: VLD-MMDiT with Layer3D RoPE
+ - Text Encoder: Qwen2.5-VL (same as base, skip quantization)
+ """
+
+ @staticmethod
+ def get_components() -> List[ComponentDefinition]:
+ return [
+ ComponentDefinition(
+ name="vae",
+ hf_subdir="vae",
+ loading_mode="single",
+ mapping_getter=QwenLayeredWeightMapping.get_vae_mapping,
+ ),
+ ComponentDefinition(
+ name="transformer",
+ hf_subdir="transformer",
+ loading_mode="multi_glob",
+ mapping_getter=QwenLayeredWeightMapping.get_transformer_mapping,
+ ),
+ ComponentDefinition(
+ name="text_encoder",
+ hf_subdir="text_encoder",
+ loading_mode="multi_json",
+ precision=mx.bfloat16,
+ skip_quantization=True, # Quantization causes significant semantic degradation
+ mapping_getter=QwenLayeredWeightMapping.get_text_encoder_mapping,
+ ),
+ ]
+
+ @staticmethod
+ def get_tokenizers() -> List[TokenizerDefinition]:
+ return [
+ TokenizerDefinition(
+ name="qwen",
+ hf_subdir="tokenizer",
+ tokenizer_class="Qwen2Tokenizer",
+ encoder_class=LanguageTokenizer,
+ max_length=1024,
+ template="<|im_start|>system\nDescribe the image by detailing the color, shape, size, texture, quantity, text, spatial relationships of the objects and background:<|im_end|>\n<|im_start|>user\n{}<|im_end|>\n<|im_start|>assistant\n",
+ download_patterns=["tokenizer/**", "added_tokens.json", "chat_template.jinja"],
+ ),
+ ]
+
+ @staticmethod
+ def get_download_patterns() -> List[str]:
+ return [
+ "vae/*.safetensors",
+ "vae/*.json",
+ "transformer/*.safetensors",
+ "transformer/*.json",
+ "text_encoder/*.safetensors",
+ "text_encoder/*.json",
+ ]
+
+ @staticmethod
+ def quantization_predicate(path: str, module) -> bool:
+ """
+ Determine if a module should be quantized.
+
+ Quantizes all modules with to_quantized method (nn.Linear, etc.)
+ except those in skip_quantization components (text_encoder).
+ """
+ return hasattr(module, "to_quantized")
diff --git a/src/mflux/models/qwen_layered/weights/qwen_layered_weight_mapping.py b/src/mflux/models/qwen_layered/weights/qwen_layered_weight_mapping.py
new file mode 100644
index 0000000..33a4cd3
--- /dev/null
+++ b/src/mflux/models/qwen_layered/weights/qwen_layered_weight_mapping.py
@@ -0,0 +1,57 @@
+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
+from mflux.models.qwen.weights.qwen_weight_mapping import QwenWeightMapping
+
+
+class QwenLayeredWeightMapping(WeightMapping):
+ """
+ Weight mapping for Qwen-Image-Layered model.
+
+ Extends the base Qwen mapping with:
+ - RGBA-VAE encoder/decoder (4-channel I/O)
+ - Additional timestep embedding for layered model
+ """
+
+ @staticmethod
+ def get_transformer_mapping() -> List[WeightTarget]:
+ """
+ Get transformer weight mapping.
+
+ The layered model has an additional timestep embedding weight.
+ Otherwise identical to base Qwen transformer.
+ """
+ # Start with base Qwen mapping
+ mappings = QwenWeightMapping.get_transformer_mapping()
+
+ # Add layered-specific mappings
+ mappings.extend([
+ # Additional timestep embedding for layers
+ WeightTarget(
+ to_pattern="time_text_embed.addition_t_embedding.weight",
+ from_pattern=["time_text_embed.addition_t_embedding.weight"],
+ required=False,
+ ),
+ ])
+
+ return mappings
+
+ @staticmethod
+ def get_vae_mapping() -> List[WeightTarget]:
+ """
+ Get VAE weight mapping for RGBA-VAE.
+
+ Same structure as base Qwen VAE, but encoder/decoder handle 4 channels.
+ The weight shapes for conv_in/conv_out will have different channel counts.
+ """
+ return QwenWeightMapping.get_vae_mapping()
+
+ @staticmethod
+ def get_text_encoder_mapping() -> List[WeightTarget]:
+ """
+ Get text encoder mapping.
+
+ Identical to base Qwen - same Qwen2.5-VL encoder.
+ """
+ return QwenWeightMapping.get_text_encoder_mapping()
diff --git a/uv.lock b/uv.lock
index 892272d..8fbb528 100644
--- a/uv.lock
+++ b/uv.lock
@@ -1,5 +1,5 @@
version = 1
-revision = 2
+revision = 3
requires-python = ">=3.10"
resolution-markers = [
"python_full_version >= '3.13' and sys_platform == 'darwin'",
@@ -792,7 +792,7 @@ wheels = [
[[package]]
name = "mflux"
-version = "0.13.2"
+version = "0.13.3"
source = { editable = "." }
dependencies = [
{ name = "huggingface-hub" },