diff --git a/README.md b/README.md index d3f0ecc4..900108b9 100644 --- a/README.md +++ b/README.md @@ -257,6 +257,7 @@ Documentation sources are also available in [`docs/`](docs/). | Stable Diffusion 3 (medium) | [Scaling Rectified Flow Transformers for High-Resolution Image Synthesis](https://arxiv.org/abs/2403.03206) | `diffusers` | | Stable Diffusion 3.5 (large, large-turbo, medium) | [Scaling Rectified Flow Transformers for High-Resolution Image Synthesis](https://arxiv.org/abs/2403.03206) | `diffusers` | | Qwen-Image | [Qwen-Image Technical Report](https://arxiv.org/abs/2508.02324) | `diffusers` | + | Qwen-Image-2.1 | [Qwen/Qwen-Image-2.1](https://huggingface.co/Qwen/Qwen-Image-2.1) (single-stream DiT + Qwen3-VL) | `diffusers` |
@@ -283,7 +284,7 @@ Documentation sources are also available in [`docs/`](docs/). ## 📜 License -This project leverages [timm](https://github.com/huggingface/pytorch-image-models#licenses), [transformers](https://github.com/huggingface/transformers#license) and [diffusers](https://github.com/huggingface/diffusers#license) for converting pretrained weights from PyTorch to Keras. For licensing details, please refer to the respective repositories. Converted weights keep their upstream license (for example, the Stable Diffusion checkpoints are CreativeML OpenRAIL-M / OpenRAIL++-M, SDXL-Turbo is non-commercial under the Stability AI Community License, and Qwen-Image is Apache-2.0). +This project leverages [timm](https://github.com/huggingface/pytorch-image-models#licenses), [transformers](https://github.com/huggingface/transformers#license) and [diffusers](https://github.com/huggingface/diffusers#license) for converting pretrained weights from PyTorch to Keras. For licensing details, please refer to the respective repositories. Converted weights keep their upstream license (for example, the Stable Diffusion checkpoints are CreativeML OpenRAIL-M / OpenRAIL++-M, SDXL-Turbo is non-commercial under the Stability AI Community License, Qwen-Image is Apache-2.0, and Qwen-Image-2.1 is under the Qwen Research License). - 🔖 **zeromodels Code**: This repository is licensed under the [Apache 2.0 License](https://www.apache.org/licenses/LICENSE-2.0). diff --git a/docs/getting_started.md b/docs/getting_started.md index 044b99b0..65a7e29e 100644 --- a/docs/getting_started.md +++ b/docs/getting_started.md @@ -55,7 +55,8 @@ results = processor.post_process_object_detection( [Qwen3-VL](qwen3_vl.md) ┬╖ [InternVL](internvl.md) ┬╖ [Kimi K2.5](kimi_k25.md) ┬╖ [LocateAnything](locateanything.md) ┬╖ - [Stable Diffusion](stable_diffusion.md) ┬╖ [Qwen-Image](qwen_image.md) + [Stable Diffusion](stable_diffusion.md) ┬╖ [Qwen-Image](qwen_image.md) ┬╖ + [Qwen-Image-2.1](qwen_image_21.md) - **Speech** diff --git a/docs/index.md b/docs/index.md index c149f5d0..f5af4bdc 100644 --- a/docs/index.md +++ b/docs/index.md @@ -327,7 +327,8 @@ Vision-language generation, grounding, and text-to-image diffusion. [Qwen3-VL](qwen3_vl.md) · [InternVL](internvl.md) · [Kimi K2.5](kimi_k25.md) · [LocateAnything](locateanything.md) · -[Stable Diffusion](stable_diffusion.md) · [Qwen-Image](qwen_image.md) +[Stable Diffusion](stable_diffusion.md) · [Qwen-Image](qwen_image.md) · +[Qwen-Image-2.1](qwen_image_21.md) diff --git a/docs/loading_weights.md b/docs/loading_weights.md index 58cf6521..35d7d498 100644 --- a/docs/loading_weights.md +++ b/docs/loading_weights.md @@ -17,7 +17,7 @@ load, and how long it takes. | # | Way | What happens | Used by | |---|---|---|---| -| 1 | [**HuggingFace Hub**](#1-hub-keras-weights) | `zeromodels/`. The repo's `zm_config.json` rebuilds the model and `model.weights.h5` (or a sharded `.weights.json`) loads with no conversion. | Vision, detection, segmentation, depth, speech, text encoders, CLIP-family, classification backbones, diffusion (Stable Diffusion, Qwen-Image; hosted only: no way 2 / 3) | +| 1 | [**HuggingFace Hub**](#1-hub-keras-weights) | `zeromodels/`. The repo's `zm_config.json` rebuilds the model and `model.weights.h5` (or a sharded `.weights.json`) loads with no conversion. | Vision, detection, segmentation, depth, speech, text encoders, CLIP-family, classification backbones, diffusion (Stable Diffusion, Qwen-Image, Qwen-Image-2.1; hosted only: no way 2 / 3) | | 2 | [**On the fly**](#2-on-the-fly-conversion) | A bare variant whose entry carries an `hf_id`. Upstream safetensors are downloaded and converted in process. | The LLMs and VLMs: Qwen, Llama, Gemma, DeepSeek, GLM, Mistral, ... | | 3 | [**`hf:` prefix**](#3-the-hf-prefix) | Any Hub repo, named explicitly. Same conversion machinery as way 2, but you pick the repo. | Fine-tunes and community weights, for any architecture | diff --git a/docs/models.md b/docs/models.md index 3b623407..342ff0f5 100644 --- a/docs/models.md +++ b/docs/models.md @@ -111,3 +111,4 @@ Vision-language encoders, generative VLMs, grounding across detection, OCR, poin - [Stable Diffusion 3](stable_diffusion_3.md) - [Stable Diffusion 3.5](stable_diffusion_3_5.md) - [Qwen-Image](qwen_image.md) +- [Qwen-Image-2.1](qwen_image_21.md) diff --git a/docs/qwen_image.md b/docs/qwen_image.md index 155e9d9e..9e9722d1 100644 --- a/docs/qwen_image.md +++ b/docs/qwen_image.md @@ -45,6 +45,7 @@ Links: - License: [Apache-2.0](https://huggingface.co/Qwen/Qwen-Image/blob/main/LICENSE) See also [qwen2_5_vl.md](qwen2_5_vl.md) (the text tower), +[qwen_image_21.md](qwen_image_21.md) (2.1 single-stream / unpatched), [stable_diffusion_3.md](stable_diffusion_3.md) (flow-match / MMDiT-style diffusion in ZeroModels). diff --git a/docs/qwen_image_21.md b/docs/qwen_image_21.md new file mode 100644 index 00000000..47657360 --- /dev/null +++ b/docs/qwen_image_21.md @@ -0,0 +1,298 @@ +# Qwen-Image-2.1 + +
+Weights: pretrained Keras weights live on Hugging Face under +zeromodels/<variant> +(each repo carries zm_config.json + sharded +*.weights.json / *.weights.h5 + +tokenizer.json). Load with from_weights("zeromodels/<variant>"). +
+ +Qwen-Image-2.1, ported to pure Keras 3: latent text-to-image flow-matching with a +32-layer **single-stream** block-causal DiT, a residual 64-channel KL autoencoder +(16× spatial, RGBA), and the Qwen3-VL text tower. The whole model is **one +container**, `QwenImage21Model`. `QwenImage21TextToImage` adds `generate`. + +The weights are converted once, offline, and hosted: on-the-fly `hf:` conversion +is deliberately **not supported** for diffusion models. + +Key facts of the port: + +- **Unpatched latents**: the denoiser sees `(B, H · W, 64)` tokens (VAE scale 16; + no 2×2 packing). At 1024px that is a `(B, 4096, 64)` sequence. +- **Block-causal attention**: text is causal; the target image block is + bidirectional and can attend to all preceding text. +- **`causal_condition`**: text tokens modulate from `t = 0` (timestep-independent), + matching Diffusers' KV-cache-ready conditioning. +- **True CFG optional**: Diffusers defaults to `true_cfg_scale=1.0` (no guidance). + Pass `guidance_scale > 1` with a negative prompt to enable dual forwards. +- **Pre-norm text features**: the text tower returns decoder outputs *before* the + final RMSNorm, matching Diffusers' forward hook on the language-model norm. +- **Schedulers match Diffusers**: `FlowMatchEulerDiscreteScheduler` with dynamic + resolution shifting (`mu` from image sequence length) and `shift_terminal`. + +Links: + +- Source: [`Qwen/Qwen-Image-2.1`](https://huggingface.co/Qwen/Qwen-Image-2.1) +- Reference: [diffusers `QwenImage21Pipeline`](https://huggingface.co/docs/diffusers/api/pipelines/qwenimage) +- License: [Qwen Research License](https://huggingface.co/Qwen/Qwen-Image-2.1) +- See also [qwen_image.md](qwen_image.md) (1.0 double-stream / packed), + [qwen3_vl.md](qwen3_vl.md) (text tower) + +## Variants + +Preconverted, bfloat16 weights are hosted under `zeromodels/`. Load with +`from_weights("zeromodels/")`. Each repo is one container: DiT + VAE + +Qwen3-VL text tower, ~28 GiB at 16-bit. + +| Variant | Hub | Source | +|---|---|---| +| `qwen-image-2.1` | [`zeromodels/qwen-image-2.1`](https://huggingface.co/zeromodels/qwen-image-2.1) | [`Qwen/Qwen-Image-2.1`](https://huggingface.co/Qwen/Qwen-Image-2.1) | + +Default `generate_args`: 40 flow-match steps, `guidance_scale=1.0`, 1024×1024. + +## API + +Configs are typed: `QwenImage21Config` (composite, `model_type` `"qwen_image_21"`) +over `QwenImage21TransformerConfig`, `QwenImage21VAEConfig` and +`QwenImage21TextConfig`, plus the checkpoint's `scheduler_config` and Qwen special +token ids. Each repo's `zm_config.json` parses through it; the constructor stays +flat, with the sub-config fields prefixed `transformer_` / `vae_` / `text_`. + +### `QwenImage21TextToImage` + +The text-to-image task: the `QwenImage21Model` container plus `BaseDiffusion`'s +`generate`. It supplies `encode_prompt` (ChatML template drop after the text +encoder), `predict_noise` on the transformer, `decode_latents` on the VAE with +mean/std un-normalization, and a `FlowMatchEulerDiscreteScheduler` from +`scheduler_config`. + +```python +generate( + input_ids, + attention_mask=None, + negative_input_ids=None, + negative_attention_mask=None, + num_inference_steps=None, + guidance_scale=None, + seed=None, + latents=None, + height=None, + width=None, + output_type="image", +) +``` + +| Arg | Default | Meaning | +|---|---|---| +| `input_ids` | required | ChatML-templated token ids, `**tokenizer(prompts)` | +| `attention_mask` | `None` | padding mask from the tokenizer | +| `negative_input_ids` | `None` | tokenized negative prompt; CFG only when set **and** `guidance_scale > 1` | +| `negative_attention_mask` | `None` | mask for the negative ids | +| `num_inference_steps` | `None` | scheduler steps; `generate_args` (40) when unset | +| `guidance_scale` | `None` | true CFG strength; `generate_args` (1.0) when unset | +| `seed` | `None` | seed for the initial latent | +| `latents` | `None` | explicit unpatched initial latent `(B, H·W, 64)` | +| `height` / `width` | `None` | output pixel size; `default_sample_size * 16` (1024) when unset | +| `output_type` | `"image"` | `"image"` for uint8 RGB, `"latent"` for sequence latents | + +Returns `(batch, height, width, 3)` uint8 numpy images (VAE RGBA is cropped to RGB). + +| Constructor arg | Default | Meaning | +|---|---|---| +| `scheduler` | `None` | a `BaseScheduler`; built from `scheduler_config` when unset | +| `transformer_sample_size` | `64` | latent side the DiT graph is built for (image / 16) | +| `vae_sample_size` | `1024` | image size the VAE graphs are built for | +| `transformer_*` / `vae_*` / `text_*` | Qwen-Image-2.1 | flat sub-config fields (see [Configuration](configuration.md)) | +| `bos_token_id` / `eos_token_id` / `pad_token_id` | `151643` / `151645` / `151643` | Qwen special tokens | +| `scheduler_config` | `None` | Diffusers scheduler dict (dynamic shifting, `max_shift=0.9`, ...) | +| `prompt_template_encode_start_idx` | `14` | ChatML system-prefix tokens dropped after the text encoder | +| `max_sequence_length` | `512` | prompt embed length after the template drop | + +### `QwenImage21Model` + +The container: one functional model whose graph is three disconnected paths. +Inputs cover the transformer (`sample`, `timestep`, `encoder_hidden_states`, +`encoder_hidden_states_mask`), the VAE (`image`, `latent`) and the text tower +(`token_ids`, `padding_mask`). Components are exposed as `.transformer`, `.vae` +and `.text_encoder`. + +### `QwenImage21Transformer2DModel` + +The denoiser (Diffusers' `QwenImage21Transformer2DModel`): a 32-layer +single-stream DiT over **unpatched** latents and Qwen3-VL text features, with +3-axis RoPE (`axes_dims_rope=(16, 56, 56)`). Inputs +`{"sample": (B, seq, 64), "timestep": (B,), "encoder_hidden_states": (B, text_seq, 4096)}` +(plus mask), output `{"sample": (B, seq, 64)}` target velocity. + +| Arg | Default | Meaning | +|---|---|---| +| `patch_size` | `1` | unpatched (asserted) | +| `in_channels` / `out_channels` | `64` / `64` | latent token width | +| `num_layers` | `32` | single-stream blocks | +| `attention_head_dim` / `num_attention_heads` | `128` / `32` | head geometry (inner dim 4096) | +| `context_in_dim` | `4096` | text feature width from Qwen3-VL | +| `mlp_ratio` | `3` | SwiGLU expansion | +| `causal_condition` | `True` | text modulates from `t=0` | +| `sample_size` | `64` | latent spatial side the graph is built for | +| `text_seq_len` | `512` | static text length after the template drop | + +### `AutoencoderKLQwenImage21` + +The VAE (Diffusers' `AutoencoderKLQwenImage21`): residual encoder/decoder, +`encode(image)` → `(B, H/16, W/16, 64)` latents, `decode(latent)` → +`(B, H, W, 4)` RGBA in `[-1, 1]`. Latent normalisation uses `latents_mean` / +`latents_std` around the denoiser. + +| Arg | Default | Meaning | +|---|---|---| +| `base_dim` / `decoder_base_dim` | `96` / `144` | channel widths | +| `z_dim` | `64` | latent channels | +| `dim_mult` | `(1, 2, 4, 8, 8)` | width multipliers per level | +| `scale_factor_spatial` | `16` | pixel → latent downscale | +| `input_channels` / `out_channels` | `4` / `4` | RGBA | +| `is_residual` | `True` | residual blocks | +| `sample_size` | `1024` | image size the graphs are built for | +| `latents_mean` / `latents_std` | Qwen-Image-2.1 | per-channel latent normalisation | + +### `QwenImage21TextEncoderModel` + +Qwen3-VL text tower only (no vision / LM head). Returns +`{"last_hidden_state"}` as **pre-final-norm** hidden states. + +### Schedulers + +| Scheduler | Notes | +|---|---| +| `FlowMatchEulerDiscreteScheduler` | `use_dynamic_shifting=True`, exponential time shift, `shift_terminal=0.02`, `max_shift=0.9` | + +At generate time the task sets timesteps with resolution-dependent `mu` and a +linspace sigma schedule, matching Diffusers. + +## Preprocessing + +### `QwenImage21Tokenizer` + +Qwen3 BPE with Diffusers' T2I ChatML template (`Comprehend and analyze the +provided prompt.`). Returns `{"input_ids", "attention_mask"}`. `encode_prompt` +drops `prompt_template_encode_start_idx` (14) tokens and keeps at most +`max_sequence_length` (512). + +```python +QwenImage21Tokenizer(hf_id=None, tokenizer_file=None, max_seq_len=1024) +``` + +## End-to-end example + +### Single prompt + +```python +import os + +os.environ["KERAS_BACKEND"] = "torch" + +from PIL import Image +from zeromodels.models.qwen_image_21 import ( + QwenImage21TextToImage, + QwenImage21Tokenizer, +) + +model = QwenImage21TextToImage.from_weights("zeromodels/qwen-image-2.1") +tokenizer = QwenImage21Tokenizer.from_weights("zeromodels/qwen-image-2.1") + +inputs = tokenizer("a photo of a capybara wearing a wizard hat, soft window light") +images = model.generate( + **inputs, + height=1024, + width=1024, + num_inference_steps=40, + guidance_scale=1.0, + seed=0, +) + +Image.fromarray(images[0]).save("capybara.png") # (1024, 1024, 3) uint8 +``` + +### Optional true CFG + +```python +inputs = tokenizer("a watercolor lighthouse at sunset") +negative = tokenizer("blurry, low quality")["input_ids"] + +images = model.generate( + **inputs, + negative_input_ids=negative, + guidance_scale=4.0, + height=1024, + width=1024, +) +``` + +### Reproducible latents + +At 1024px the unpatched noise is `(batch, 64, 64, 64)` → `(batch, 4096, 64)`: + +```python +import numpy as np +from zeromodels.models.qwen_image_21.qwen_image_21_model import pack_latents + +h, w, channels = 64, 64, 64 +noise = np.random.default_rng(0).standard_normal((1, h, w, channels)).astype("float32") +latents = pack_latents(noise, h, w) +images = model.generate(**tokenizer("a bowl of ramen"), latents=latents) +``` + +### Other resolutions + +Graphs are built for a fixed size; weights are not. Rebuild with constructor +overrides (multiples of 32px; checkpoint targets 1024): + +```python +model = QwenImage21TextToImage.from_weights( + "zeromodels/qwen-image-2.1", + transformer_sample_size=32, + vae_sample_size=512, +) +images = model.generate(**tokenizer("a mountain lake at dawn"), height=512, width=512) +``` + +### Offline conversion + +```bash +pip install zeromodels[conversion] +# KERAS_BACKEND=torch; prefer CPU for the full ~28 GiB bf16 build +python -m zeromodels.models.qwen_image_21.convert_qwen_image_21_diffusers_to_keras +``` + +`transfer_qwen_image_21(repo)` streams Diffusers shards into a Keras container and +writes sharded `*.weights.json`. Building the full text tower on GPU can OOM +during Functional shape tracing; convert on CPU, then load for CUDA inference. + +## Data Format + +**Channels-last for the VAE.** Transformer works on sequence tokens. + +| | Shape | +|---|---| +| `latents` passed to `generate` (unpatched) | `(batch, (H/16)·(W/16), 64)` | +| VAE encode input | `(batch, H, W, 4)` RGBA in `[-1, 1]` | +| VAE decode output | `(batch, H, W, 4)`; `generate` returns RGB uint8 | +| VAE latent | `(batch, H/16, W/16, 64)` | +| Transformer `sample` | `(batch, seq, 64)` | + +## Memory and speed + +The bf16 container is about 28 GiB (DiT ~14 GiB + text ~14 GiB + VAE ~1.4 GiB). +Building the full Functional graph at 1024px on GPU can OOM during attention shape +tracing (layers use `compute_output_spec` to avoid materializing seq² temps). Prefer +CPU convert / build, then CUDA generate. A 1024px run needs substantial activation +headroom on an 80 GB card; prefer fused attention +(`keras.ops.dot_product_attention` / torch SDPA). + +## Loading Fine-tuned Weights + +The hosted checkpoint is the supported weight; any repo laid out like it +(`zm_config.json` declaring `QwenImage21Model`, sharded `*.weights.json` / +`*.weights.h5`, `tokenizer.json`) loads with `from_weights("/")`. The +`hf:` prefix raises for diffusion models: convert Diffusers format once with +`convert_qwen_image_21_diffusers_to_keras.py` and host the result. diff --git a/tests/base/model_test_registry.py b/tests/base/model_test_registry.py index c558b204..66cf90d6 100644 --- a/tests/base/model_test_registry.py +++ b/tests/base/model_test_registry.py @@ -4846,6 +4846,63 @@ "expected_output_shape": dict(_qwen_image_outputs), } +# Qwen-Image-2.1: unpatched single-stream DiT + residual RGBA VAE + Qwen3-VL text. +# sample_size=4 → seq=16; VAE 64px → latent 4 with scale 16. +_qwen_image_21_tiny = { + "transformer_sample_size": 4, + "transformer_patch_size": 1, + "transformer_in_channels": 16, + "transformer_out_channels": 16, + "transformer_num_layers": 1, + "transformer_attention_head_dim": 8, + "transformer_num_attention_heads": 2, + "transformer_context_in_dim": 32, + "transformer_mlp_ratio": 2, + "transformer_axes_dims_rope": (2, 2, 4), + "max_sequence_length": 8, + "vae_sample_size": 64, + "vae_base_dim": 16, + "vae_decoder_base_dim": 16, + "vae_z_dim": 8, + "vae_dim_mult": (1, 2, 2), + "vae_num_res_blocks": 1, + "vae_input_channels": 4, + "vae_out_channels": 4, + "vae_is_residual": True, + "vae_scale_factor_spatial": 4, + "text_embed_dim": 32, + "text_mlp_dim": 64, + "text_num_layers": 1, + "text_num_heads": 2, + "text_num_kv_heads": 1, + "max_seq_len": 16, + "vocab_size": 256, + "default_sample_size": 4, +} + +_qwen_image_21_outputs = { + "noise_pred": (2, 16, 16), + "moments": (2, 16, 16, 16), + "image": (2, 64, 64, 4), + "prompt_embeds": (2, 16, 32), +} +MODEL_TEST_CONFIGS["QwenImage21Model"] = { + "module": "zeromodels.models.qwen_image_21", + "model_cls": "QwenImage21Model", + "model_type": "diffusion", + "init_kwargs": dict(_qwen_image_21_tiny), + "input_factory": "qwen_image_21_input", + "expected_output_shape": dict(_qwen_image_21_outputs), +} +MODEL_TEST_CONFIGS["QwenImage21TextToImage"] = { + "module": "zeromodels.models.qwen_image_21", + "model_cls": "QwenImage21TextToImage", + "model_type": "diffusion", + "init_kwargs": dict(_qwen_image_21_tiny), + "input_factory": "qwen_image_21_input", + "expected_output_shape": dict(_qwen_image_21_outputs), +} + def get_all_model_ids(): return list(MODEL_TEST_CONFIGS.keys()) diff --git a/tests/fixtures/cross_backend_parity.json b/tests/fixtures/cross_backend_parity.json index ac35e568..9cd433b3 100644 --- a/tests/fixtures/cross_backend_parity.json +++ b/tests/fixtures/cross_backend_parity.json @@ -20615,6 +20615,598 @@ ] } ], + "QwenImage21Model": [ + { + "shape": [ + 2, + 64, + 64, + 4 + ], + "sample": [ + 0.000958, + -0.002757, + -0.002676, + -0.002686, + -0.002682, + -0.002686, + -0.002666, + -0.002711, + -0.00273, + -0.0208, + -0.020851, + -0.02082, + -0.020843, + -0.02082, + -0.020862, + -0.020826, + -0.020863, + -0.020826, + 0.015441, + 0.015452, + 0.015469, + 0.015462, + 0.015427, + 0.015388, + 0.015356, + 0.015404, + 0.015399, + 0.033911, + 0.033944, + 0.03398, + 0.033954, + 0.03374, + 0.034017, + 0.033818, + 0.033958, + 0.033937, + -0.002685, + -0.002678, + -0.002668, + -0.002742, + -0.002671, + -0.002652, + -0.00272, + -0.002738, + -0.002786, + -0.020785, + -0.020731, + -0.020771, + -0.020838, + -0.020827, + -0.020796, + -0.020871, + -0.020885, + -0.020835, + 0.015476, + 0.015424, + 0.015454, + 0.015447, + 0.015401, + 0.015431, + 0.015359, + 0.015402, + 0.015371, + 0.032569 + ] + }, + { + "shape": [ + 2, + 16, + 16, + 16 + ], + "sample": [ + 0.043742, + 0.006526, + -0.00716, + 0.002203, + 0.007001, + -0.003265, + 0.041306, + 0.048222, + 0.044051, + 0.006727, + -0.007272, + 0.002304, + 0.007318, + -0.003362, + 0.041373, + 0.048279, + 0.043725, + 0.006633, + -0.007264, + 0.00234, + 0.007361, + -0.003407, + 0.041377, + 0.048261, + 0.043969, + 0.006626, + -0.007279, + 0.002379, + 0.007225, + -0.003334, + 0.041278, + 0.048426, + 0.043925, + 0.006591, + -0.007236, + 0.002311, + 0.007166, + -0.003305, + 0.04118, + 0.048263, + 0.043971, + 0.006603, + -0.007174, + 0.002288, + 0.007199, + -0.003511, + 0.041368, + 0.048239, + 0.04406, + 0.006657, + -0.007153, + 0.002314, + 0.007224, + -0.003533, + 0.041315, + 0.048327, + 0.043871, + 0.006593, + -0.007152, + 0.002406, + 0.007264, + -0.003716, + 0.041199, + -0.03484 + ] + }, + { + "shape": [ + 2, + 16, + 16 + ], + "sample": [ + 0.042214, + 0.105714, + 0.056851, + 0.076787, + 0.029758, + 0.00917, + 0.095684, + 0.106859, + 0.032329, + -0.067297, + 0.057183, + -0.051838, + -0.013893, + 0.035384, + -0.029907, + 0.154545, + -0.02961, + 0.08136, + 0.013182, + -0.054818, + -0.108304, + -0.122532, + -0.010875, + -0.090636, + -0.027775, + 0.009631, + -0.066816, + -0.096933, + -0.096951, + 0.185838, + -0.088463, + -0.070053, + -0.012736, + -0.047557, + 0.10501, + 0.094757, + 0.087138, + -0.094607, + -0.200976, + 0.063318, + 0.071323, + -0.137504, + 0.045576, + -0.086316, + -0.025822, + -0.020573, + 0.134116, + 0.039678, + -0.002558, + -0.075023, + 0.19758, + 0.001701, + 0.030064, + -0.005096, + -0.038978, + -0.093928, + -0.026556, + -0.030843, + -0.01305, + -0.017115, + 0.036709, + 0.005811, + 0.139922, + 0.009673 + ] + }, + { + "shape": [ + 2, + 16, + 32 + ], + "sample": [ + 0.002967, + -0.002949, + 0.002967, + -0.002949, + 0.002967, + 0.026536, + -0.02298, + 0.026536, + -0.02298, + 0.023317, + 0.011212, + 0.023317, + 0.011212, + 0.000229, + 0.004552, + 0.000229, + 0.004552, + 0.012545, + -0.008191, + 0.012545, + -0.008191, + 0.043657, + -0.018314, + 0.043657, + -0.018314, + 0.043657, + 0.022612, + 0.015819, + 0.022612, + 0.015819, + -0.001567, + -0.012379, + -0.001567, + -0.012379, + 0.008989, + 0.015044, + 0.008989, + 0.015044, + 0.005536, + -0.028382, + 0.005536, + -0.028382, + 0.01766, + 0.026689, + 0.01766, + 0.026689, + 0.01766, + -0.036786, + 0.017117, + -0.036786, + 0.017117, + -0.033625, + -0.005554, + -0.033625, + -0.005554, + -0.017666, + -0.012415, + -0.017666, + -0.012415, + 0.012932, + -0.040328, + 0.012932, + -0.040328, + -0.023516 + ] + } + ], + "QwenImage21TextToImage": [ + { + "shape": [ + 2, + 64, + 64, + 4 + ], + "sample": [ + 0.000958, + -0.002757, + -0.002676, + -0.002686, + -0.002682, + -0.002686, + -0.002666, + -0.002711, + -0.00273, + -0.0208, + -0.020851, + -0.02082, + -0.020843, + -0.02082, + -0.020862, + -0.020826, + -0.020863, + -0.020826, + 0.015441, + 0.015452, + 0.015469, + 0.015462, + 0.015427, + 0.015388, + 0.015356, + 0.015404, + 0.015399, + 0.033911, + 0.033944, + 0.03398, + 0.033954, + 0.03374, + 0.034017, + 0.033818, + 0.033958, + 0.033937, + -0.002685, + -0.002678, + -0.002668, + -0.002742, + -0.002671, + -0.002652, + -0.00272, + -0.002738, + -0.002786, + -0.020785, + -0.020731, + -0.020771, + -0.020838, + -0.020827, + -0.020796, + -0.020871, + -0.020885, + -0.020835, + 0.015476, + 0.015424, + 0.015454, + 0.015447, + 0.015401, + 0.015431, + 0.015359, + 0.015402, + 0.015371, + 0.032569 + ] + }, + { + "shape": [ + 2, + 16, + 16, + 16 + ], + "sample": [ + 0.043742, + 0.006526, + -0.00716, + 0.002203, + 0.007001, + -0.003265, + 0.041306, + 0.048222, + 0.044051, + 0.006727, + -0.007272, + 0.002304, + 0.007318, + -0.003362, + 0.041373, + 0.048279, + 0.043725, + 0.006633, + -0.007264, + 0.00234, + 0.007361, + -0.003407, + 0.041377, + 0.048261, + 0.043969, + 0.006626, + -0.007279, + 0.002379, + 0.007225, + -0.003334, + 0.041278, + 0.048426, + 0.043925, + 0.006591, + -0.007236, + 0.002311, + 0.007166, + -0.003305, + 0.04118, + 0.048263, + 0.043971, + 0.006603, + -0.007174, + 0.002288, + 0.007199, + -0.003511, + 0.041368, + 0.048239, + 0.04406, + 0.006657, + -0.007153, + 0.002314, + 0.007224, + -0.003533, + 0.041315, + 0.048327, + 0.043871, + 0.006593, + -0.007152, + 0.002406, + 0.007264, + -0.003716, + 0.041199, + -0.03484 + ] + }, + { + "shape": [ + 2, + 16, + 16 + ], + "sample": [ + 0.042214, + 0.105714, + 0.056851, + 0.076787, + 0.029758, + 0.00917, + 0.095684, + 0.106859, + 0.032329, + -0.067297, + 0.057183, + -0.051838, + -0.013893, + 0.035384, + -0.029907, + 0.154545, + -0.02961, + 0.08136, + 0.013182, + -0.054818, + -0.108304, + -0.122532, + -0.010875, + -0.090636, + -0.027775, + 0.009631, + -0.066816, + -0.096933, + -0.096951, + 0.185838, + -0.088463, + -0.070053, + -0.012736, + -0.047557, + 0.10501, + 0.094757, + 0.087138, + -0.094607, + -0.200976, + 0.063318, + 0.071323, + -0.137504, + 0.045576, + -0.086316, + -0.025822, + -0.020573, + 0.134116, + 0.039678, + -0.002558, + -0.075023, + 0.19758, + 0.001701, + 0.030064, + -0.005096, + -0.038978, + -0.093928, + -0.026556, + -0.030843, + -0.01305, + -0.017115, + 0.036709, + 0.005811, + 0.139922, + 0.009673 + ] + }, + { + "shape": [ + 2, + 16, + 32 + ], + "sample": [ + 0.002967, + -0.002949, + 0.002967, + -0.002949, + 0.002967, + 0.026536, + -0.02298, + 0.026536, + -0.02298, + 0.023317, + 0.011212, + 0.023317, + 0.011212, + 0.000229, + 0.004552, + 0.000229, + 0.004552, + 0.012545, + -0.008191, + 0.012545, + -0.008191, + 0.043657, + -0.018314, + 0.043657, + -0.018314, + 0.043657, + 0.022612, + 0.015819, + 0.022612, + 0.015819, + -0.001567, + -0.012379, + -0.001567, + -0.012379, + 0.008989, + 0.015044, + 0.008989, + 0.015044, + 0.005536, + -0.028382, + 0.005536, + -0.028382, + 0.01766, + 0.026689, + 0.01766, + 0.026689, + 0.01766, + -0.036786, + 0.017117, + -0.036786, + 0.017117, + -0.033625, + -0.005554, + -0.033625, + -0.005554, + -0.017666, + -0.012415, + -0.017666, + -0.012415, + 0.012932, + -0.040328, + 0.012932, + -0.040328, + -0.023516 + ] + } + ], "QwenImageModel": [ { "shape": [ diff --git a/tests/fixtures/dummy_inputs.py b/tests/fixtures/dummy_inputs.py index 2467e0f9..8a312a50 100644 --- a/tests/fixtures/dummy_inputs.py +++ b/tests/fixtures/dummy_inputs.py @@ -264,6 +264,32 @@ def qwen_image_input( } +def qwen_image_21_input( + batch_size=2, + img_seq=16, + in_channels=16, + image_size=64, + latent_size=16, + z_dim=8, + text_seq_len=8, + context_in_dim=32, + max_seq_len=16, +): + """Dummy inputs for the Qwen-Image-2.1 container graph (unpatched latents).""" + return { + "sample": ops.ones((batch_size, img_seq, in_channels)), + "timestep": ops.ones((batch_size,)), + "encoder_hidden_states": ops.ones((batch_size, text_seq_len, context_in_dim)), + "encoder_hidden_states_mask": ops.ones( + (batch_size, text_seq_len), dtype="int32" + ), + "image": ops.ones((batch_size, image_size, image_size, 4)), + "latent": ops.ones((batch_size, latent_size, latent_size, z_dim)), + "token_ids": ops.ones((batch_size, max_seq_len), dtype="int32"), + "padding_mask": ops.ones((batch_size, max_seq_len), dtype="int32"), + } + + def tips_v2_text_input(batch_size=2, max_seq_len=16): return { "token_ids": ops.ones((batch_size, max_seq_len), dtype="int32"), diff --git a/tests/integration/test_data_formats.py b/tests/integration/test_data_formats.py index 765944c0..517db8b1 100644 --- a/tests/integration/test_data_formats.py +++ b/tests/integration/test_data_formats.py @@ -40,6 +40,9 @@ # Qwen-Image's Wan-derived VAE is built channels_last (NTHWC Conv3D). "QwenImageModel", "QwenImageTextToImage", + # Qwen-Image-2.1 residual VAE is channels_last (spatial Conv2d, RGBA). + "QwenImage21Model", + "QwenImage21TextToImage", # Qwen-VL inputs are pre-patchified (no spatial axes) -> layout-agnostic. "Qwen2VLModel", "Qwen2_5VLModel", diff --git a/website/mkdocs.yml b/website/mkdocs.yml index 0a838c93..06eb039b 100644 --- a/website/mkdocs.yml +++ b/website/mkdocs.yml @@ -238,6 +238,7 @@ nav: - Stable Diffusion 3: stable_diffusion_3.md - Stable Diffusion 3.5: stable_diffusion_3_5.md - Qwen-Image: qwen_image.md + - Qwen-Image-2.1: qwen_image_21.md - TIPSv2: tipsv2.md - Loading Weights: loading_weights.md diff --git a/zeromodels/auto/auto_mapping_names.py b/zeromodels/auto/auto_mapping_names.py index 92fe9034..6b925e3b 100644 --- a/zeromodels/auto/auto_mapping_names.py +++ b/zeromodels/auto/auto_mapping_names.py @@ -282,6 +282,7 @@ "stable_diffusion_3": "StableDiffusion3Model", "stable_diffusion_3_5": "StableDiffusion3_5Model", "qwen_image": "QwenImageModel", + "qwen_image_21": "QwenImage21Model", "swin": "SwinModel", "swinv2": "SwinV2Model", "t5": "T5Model", @@ -420,6 +421,7 @@ "stable_diffusion_3": "StableDiffusion3TextToImage", "stable_diffusion_3_5": "StableDiffusion3_5TextToImage", "qwen_image": "QwenImageTextToImage", + "qwen_image_21": "QwenImage21TextToImage", }, "ImageToImage": { "stable_diffusion_xl_refiner": "StableDiffusionXLRefinerImageToImage", @@ -689,6 +691,10 @@ "stable_diffusion_3_5": "StableDiffusion3_5Config", "stable_diffusion_3_t5_encoder": "StableDiffusion3T5EncoderConfig", "qwen_image": "QwenImageConfig", + "qwen_image_21": "QwenImage21Config", + "qwen_image_21_transformer_2d": "QwenImage21TransformerConfig", + "qwen_image_21_text": "QwenImage21TextConfig", + "autoencoder_kl_qwen_image_21": "QwenImage21VAEConfig", "swin": "SwinConfig", "swinv2": "SwinV2Config", "t5": "T5Config", @@ -792,6 +798,7 @@ "stable_diffusion_3": "StableDiffusion3Tokenizer", "stable_diffusion_3_5": "StableDiffusion3_5Tokenizer", "qwen_image": "QwenImageTokenizer", + "qwen_image_21": "QwenImage21Tokenizer", "t5": "T5Tokenizer", "tipsv2": "Tipsv2Tokenizer", "whisper": "WhisperTokenizer", diff --git a/zeromodels/models/__init__.py b/zeromodels/models/__init__.py index fb9e924e..74301652 100644 --- a/zeromodels/models/__init__.py +++ b/zeromodels/models/__init__.py @@ -105,6 +105,7 @@ qwen3_vl, qwen3_vl_moe, qwen_image, + qwen_image_21, regnet, res2net, resmlp, diff --git a/zeromodels/models/qwen3_vl/qwen3_vl_layers.py b/zeromodels/models/qwen3_vl/qwen3_vl_layers.py index e0c8fd38..885091d5 100644 --- a/zeromodels/models/qwen3_vl/qwen3_vl_layers.py +++ b/zeromodels/models/qwen3_vl/qwen3_vl_layers.py @@ -43,6 +43,12 @@ def __init__(self, embed_dim, mlp_dim, **kwargs): self.up = layers.Dense(mlp_dim, use_bias=False, name="up") self.down = layers.Dense(embed_dim, use_bias=False, name="down") + def build(self, input_shape): + self.gate.build(input_shape) + self.up.build(input_shape) + self.down.build((*tuple(input_shape)[:-1], self.mlp_dim)) + self.built = True + def call(self, x): return self.down(ops.silu(self.gate(x)) * self.up(x)) @@ -85,6 +91,17 @@ def __init__( self.query_norm = Qwen3VLRMSNorm(eps=norm_eps, name="query_norm") self.key_norm = Qwen3VLRMSNorm(eps=norm_eps, name="key_norm") + def build(self, input_shape): + self.query.build(input_shape) + self.key.build(input_shape) + self.value.build(input_shape) + self.output_proj.build( + (*tuple(input_shape)[:-1], self.num_heads * self.head_dim) + ) + self.query_norm.build((self.head_dim,)) + self.key_norm.build((self.head_dim,)) + self.built = True + def call( self, hidden_states, @@ -140,6 +157,19 @@ def call( out = self.output_proj(out) return (out, new_kv) if use_cache else out + def compute_output_spec( + self, + hidden_states, + cos, + sin, + attention_mask=None, + past_key_value=None, + use_cache=False, + ): + # Skip fused_attention during Functional shape tracing (seq² temps OOM on GPU). + out = keras.KerasTensor(hidden_states.shape, dtype=self.compute_dtype) + return (out, None) if use_cache else out + def decode_step( self, hidden_states, cos, sin, cache_k, cache_v, write_pos, key_mask ): @@ -229,6 +259,13 @@ def __init__( self.mlp_norm = Qwen3VLRMSNorm(eps=norm_eps, name="mlp_norm") self.mlp = Qwen3VLMLP(embed_dim, mlp_dim, name="mlp") + def build(self, input_shape): + self.attention_norm.build(input_shape) + self.attention.build(input_shape) + self.mlp_norm.build(input_shape) + self.mlp.build(input_shape) + self.built = True + def call( self, hidden_states, @@ -257,6 +294,19 @@ def call( hidden_states = residual + self.mlp(hidden_states) return (hidden_states, new_kv) if use_cache else hidden_states + def compute_output_spec( + self, + hidden_states, + cos, + sin, + attention_mask=None, + past_key_value=None, + use_cache=False, + ): + # Residual stream keeps its shape; skip attention softmax during graph build. + out = keras.KerasTensor(hidden_states.shape, dtype=self.compute_dtype) + return (out, None) if use_cache else out + def decode_step( self, hidden_states, cos, sin, cache_k, cache_v, write_pos, key_mask ): diff --git a/zeromodels/models/qwen_image_21/__init__.py b/zeromodels/models/qwen_image_21/__init__.py new file mode 100644 index 00000000..cace1f2d --- /dev/null +++ b/zeromodels/models/qwen_image_21/__init__.py @@ -0,0 +1,27 @@ +from .qwen_image_21_config import ( + QwenImage21Config, + QwenImage21TextConfig, + QwenImage21TransformerConfig, + QwenImage21VAEConfig, +) +from .qwen_image_21_model import ( + AutoencoderKLQwenImage21, + QwenImage21Model, + QwenImage21TextEncoderModel, + QwenImage21TextToImage, + QwenImage21Transformer2DModel, +) +from .qwen_image_21_tokenizer import QwenImage21Tokenizer + +__all__ = [ + "AutoencoderKLQwenImage21", + "QwenImage21Config", + "QwenImage21Model", + "QwenImage21TextConfig", + "QwenImage21TextEncoderModel", + "QwenImage21TextToImage", + "QwenImage21Tokenizer", + "QwenImage21Transformer2DModel", + "QwenImage21TransformerConfig", + "QwenImage21VAEConfig", +] diff --git a/zeromodels/models/qwen_image_21/convert_qwen_image_21_diffusers_to_keras.py b/zeromodels/models/qwen_image_21/convert_qwen_image_21_diffusers_to_keras.py new file mode 100644 index 00000000..845135b3 --- /dev/null +++ b/zeromodels/models/qwen_image_21/convert_qwen_image_21_diffusers_to_keras.py @@ -0,0 +1,397 @@ +from __future__ import annotations + +from typing import Dict + +import numpy as np +from tqdm import tqdm + +from zeromodels.conversion.exceptions import ( + WeightMappingError, + WeightShapeMismatchError, +) +from zeromodels.conversion.weight_split_util import split_model_weights +from zeromodels.conversion.weight_transfer_util import ( + compare_keras_torch_names, + transfer_weights, + zeros_init, +) + +QWEN_IMAGE_21_SOURCES = { + "qwen-image-2.1": "Qwen/Qwen-Image-2.1", +} + +WEIGHT_NAME_MAPPING: Dict[str, str] = { + "__": ".", + "/kernel": ".weight", + "/gamma": ".weight", + "/beta": ".bias", + "/scale": ".weight", + "/": ".", + "gamma": "weight", + "beta": "bias", + "kernel": "weight", + # Qwen3-VL text tower + "token_embedding.embeddings": "model.embed_tokens.weight", + "language_model.final_norm.weight": "model.norm.weight", + "language_model.": "model.", + "decoder_layer_": "layers.", + "attention.query_norm": "self_attn.q_norm", + "attention.key_norm": "self_attn.k_norm", + "attention.query": "self_attn.q_proj", + "attention.key": "self_attn.k_proj", + "attention.value": "self_attn.v_proj", + "attention.output_proj": "self_attn.o_proj", + "attention_norm": "input_layernorm", + "mlp_norm": "post_attention_layernorm", + "mlp.gate": "mlp.gate_proj", + "mlp.up": "mlp.up_proj", + "mlp.down": "mlp.down_proj", +} + + +def config_from_diffusers(repo, token=None): + import json + + from huggingface_hub import hf_hub_download + from diffusers import FlowMatchEulerDiscreteScheduler + from transformers import AutoConfig + + from zeromodels.models.qwen_image_21.qwen_image_21_config import ( + QwenImage21Config, + QwenImage21VAEConfig, + ) + from zeromodels.models.qwen_image_21.qwen_image_21_model import ( + QwenImage21Transformer2DModel, + ) + + transformer = json.load( + open( + hf_hub_download(repo, "config.json", subfolder="transformer", token=token), + encoding="utf-8", + ) + ) + vae = json.load( + open( + hf_hub_download(repo, "config.json", subfolder="vae", token=token), + encoding="utf-8", + ) + ) + text = AutoConfig.from_pretrained( + repo, subfolder="text_encoder", token=token + ).to_dict() + text_inner = text.get("text_config") or text + scheduler = { + k: v + for k, v in FlowMatchEulerDiscreteScheduler.load_config( + repo, subfolder="scheduler", token=token + ).items() + if k == "_class_name" or not k.startswith("_") + } + temperal = tuple(vae.get("temperal_downsample", (False, True, True, True))) + rope = text_inner.get("rope_parameters") or text_inner.get("rope_scaling") or {} + return QwenImage21Config( + transformer_config={ + **QwenImage21Transformer2DModel.kwargs_from_diffusers_config(transformer), + "sample_size": 64, + "text_seq_len": 512, + }, + vae_config=QwenImage21VAEConfig( + base_dim=vae.get("base_dim", 96), + decoder_base_dim=vae.get("decoder_base_dim", 144), + z_dim=vae.get("z_dim", 64), + dim_mult=tuple(vae.get("dim_mult", (1, 2, 4, 8, 8))), + num_res_blocks=vae.get("num_res_blocks", 2), + attn_scales=tuple(vae.get("attn_scales") or ()), + temperal_downsample=temperal, + dropout=vae.get("dropout", 0.0), + input_channels=vae.get("in_channels", 4), + out_channels=vae.get("out_channels", 4), + is_residual=bool(vae.get("is_residual", True)), + scale_factor_spatial=vae.get("scale_factor_spatial", 16), + latents_mean=tuple(vae["latents_mean"]), + latents_std=tuple(vae["latents_std"]), + sample_size=1024, + ), + text_config={ + "vocab_size": text_inner.get("vocab_size", 151936), + "embed_dim": text_inner.get("hidden_size", 4096), + "mlp_dim": text_inner.get("intermediate_size", 12288), + "num_layers": text_inner.get("num_hidden_layers", 36), + "num_heads": text_inner.get("num_attention_heads", 32), + "num_kv_heads": text_inner.get("num_key_value_heads", 8), + "head_dim": text_inner.get("head_dim", 128), + "norm_eps": text_inner.get("rms_norm_eps", 1e-6), + "rope_theta": float( + rope.get("rope_theta", text_inner.get("rope_theta", 5000000.0)) + ), + "mrope_section": tuple(rope.get("mrope_section", (24, 20, 20))), + "tie_embeddings": text.get("tie_word_embeddings", False), + "max_seq_len": 1024, + }, + scheduler_config=scheduler, + bos_token_id=text_inner.get("bos_token_id", 151643), + eos_token_id=text_inner.get("eos_token_id", 151645), + pad_token_id=text_inner.get("bos_token_id", 151643), + image_token_id=text.get("image_token_id", 151655), + ) + + +def transfer_qwen_image_21( + repo, token=None, dtype="bfloat16", build_sample_size=8, config=None +): + import gc + import json + + from huggingface_hub import hf_hub_download + from safetensors import safe_open + + from zeromodels.base.base_mixin import build_dtype_scope + from zeromodels.models.qwen_image_21.qwen_image_21_model import QwenImage21Model + + config = config or config_from_diffusers(repo, token=token) + flat = config.constructor_kwargs() + flat["transformer_sample_size"] = build_sample_size + flat["vae_sample_size"] = max(build_sample_size * 16, 64) + flat["max_sequence_length"] = min(int(flat.get("max_sequence_length", 512)), 64) + flat["transformer_text_seq_len"] = flat["max_sequence_length"] + flat["max_seq_len"] = min(int(flat.get("max_seq_len", 1024)), 64) + + print(f"[1/4] Building QwenImage21Model (dtype={dtype})…", flush=True) + with build_dtype_scope(dtype), zeros_init(): + model = QwenImage21Model(**flat) + + # VAE RMSNorm checkpoints keep the name ``gamma`` (not ``weight``). + vae_mapping = { + k: v for k, v in WEIGHT_NAME_MAPPING.items() if k not in ("/gamma", "gamma") + } + for step, (component, subfolder, mapping, index_name, filename) in enumerate( + ( + ( + model.transformer, + "transformer", + WEIGHT_NAME_MAPPING, + "diffusion_pytorch_model.safetensors.index.json", + None, + ), + ( + model.vae, + "vae", + vae_mapping, + "diffusion_pytorch_model.safetensors.index.json", + "diffusion_pytorch_model.safetensors", + ), + ), + start=2, + ): + print(f"[{step}/4] Transferring {subfolder}…", flush=True) + state = {} + try: + if index_name is not None: + index_path = hf_hub_download( + repo, index_name, subfolder=subfolder, token=token + ) + with open(index_path, encoding="utf-8") as f: + weight_map = json.load(f)["weight_map"] + shard_paths = { + shard: hf_hub_download( + repo, shard, subfolder=subfolder, token=token + ) + for shard in sorted(set(weight_map.values())) + } + + class _State: + def __contains__(self, key): + return key in weight_map + + def __getitem__(self, key): + with safe_open( + shard_paths[weight_map[key]], framework="np" + ) as shard: + return shard.get_tensor(key) + + def keys(self): + return weight_map.keys() + + def __iter__(self): + return iter(weight_map) + + state = _State() + else: + raise FileNotFoundError + except Exception: + path = hf_hub_download( + repo, + filename or "diffusion_pytorch_model.safetensors", + subfolder=subfolder, + token=token, + ) + state = {} + with safe_open(path, framework="np") as shard: + for key in shard.keys(): + state[key] = shard.get_tensor(key) + + consumed = set() + trainable, non_trainable = split_model_weights(component) + prefer_leaf = subfolder in ("vae", "transformer") + for keras_weight, _ in tqdm( + trainable + non_trainable, + desc=f"Transferring {subfolder} weights to Keras", + ): + leaf = "/".join(keras_weight.path.split("/")[-2:]) + if prefer_leaf: + candidates = [leaf] + else: + parts = keras_weight.path.split("/") + if parts and parts[0] in ( + component.name, + "transformer", + "vae", + "QwenImage21Transformer2DModel", + "AutoencoderKLQwenImage21", + ): + rel = "/".join(parts[1:]) + else: + rel = "/".join(parts) + candidates = [rel, leaf] + key = None + for cand in candidates: + mapped = cand + for old, new in mapping.items(): + mapped = mapped.replace(old, new) + if mapped in state: + key = mapped + break + if key is None: + raise WeightMappingError( + keras_weight.path, + candidates[0] + if prefer_leaf + else "/".join(keras_weight.path.split("/")[-2:]), + ) + consumed.add(key) + raw = state[key] + arr = np.asarray(raw) + kshape = tuple(keras_weight.shape) + if len(kshape) == 5 and arr.ndim == 5: + arr = np.transpose(arr, (2, 3, 4, 1, 0)) + elif len(kshape) == 4 and arr.ndim == 4: + arr = np.transpose(arr, (2, 3, 1, 0)) + elif len(kshape) == 2 and arr.ndim == 4 and tuple(arr.shape[-2:]) == (1, 1): + arr = arr[:, :, 0, 0] + elif ( + arr.ndim > 1 + and len(kshape) == 1 + and int(np.prod(arr.shape)) == kshape[0] + ): + arr = arr.reshape(kshape) + if len(keras_weight.shape) in (4, 5): + if tuple(keras_weight.shape) != arr.shape: + raise WeightShapeMismatchError( + keras_weight.path, keras_weight.shape, key, np.shape(raw) + ) + keras_weight.assign(arr) + continue + if tuple(keras_weight.shape) != tuple(arr.shape): + if not compare_keras_torch_names( + keras_weight.path, keras_weight, key, raw + ): + raise WeightShapeMismatchError( + keras_weight.path, + keras_weight.shape, + key, + np.asarray(raw).shape, + ) + transfer_weights(keras_weight.path, keras_weight, arr) + del state + gc.collect() + + print("[4/4] Transferring text encoder…", flush=True) + index_path = hf_hub_download( + repo, "model.safetensors.index.json", subfolder="text_encoder", token=token + ) + with open(index_path, encoding="utf-8") as f: + weight_map = json.load(f)["weight_map"] + shard_paths = { + shard: hf_hub_download(repo, shard, subfolder="text_encoder", token=token) + for shard in sorted(set(weight_map.values())) + } + hf_keys = {} + for key in weight_map: + if key.startswith("model.visual.") or key.startswith("lm_head."): + continue + if key == "model.norm.weight": + continue + if key.startswith("model.language_model."): + rest = key[len("model.language_model.") :] + hf_keys["model." + rest] = key + hf_keys[rest] = key + elif key.startswith("model.") and not key.startswith("model.visual."): + hf_keys[key] = key + hf_keys[key[len("model.") :]] = key + hf_keys.pop("model.norm.weight", None) + hf_keys.pop("norm.weight", None) + consumed = set() + text_encoder = model.text_encoder + for weight in tqdm( + text_encoder.weights, desc="Transferring text_encoder weights to Keras" + ): + name = weight.path.removeprefix(f"{text_encoder.name}/") + if name.startswith("language_model/"): + name = name[len("language_model/") :] + for old, new in WEIGHT_NAME_MAPPING.items(): + name = name.replace(old, new) + if name not in hf_keys and name.startswith("model."): + alt = name[len("model.") :] + if alt in hf_keys: + name = alt + if name not in hf_keys: + raise WeightMappingError(weight.path, name) + consumed.add(name) + shard_key = hf_keys[name] + with safe_open(shard_paths[weight_map[shard_key]], framework="np") as shard: + transfer_weights(weight.path, weight, shard.get_tensor(shard_key)) + del weight_map, shard_paths, hf_keys + gc.collect() + return model, config + + +if __name__ == "__main__": + import gc + import os + + import keras + + OUT_DIR = os.environ.get( + "ZM_OUT_DIR", "C:/Users/gites/Desktop/code/qwen_image_21_weights" + ) + os.makedirs(OUT_DIR, exist_ok=True) + MAX_SHARD_GB = 5.0 + token = os.environ.get("HF_TOKEN") + dtype = os.environ.get("ZM_DTYPE", "bfloat16") + device = os.environ.get("ZM_DEVICE", "cpu") + selected = [v for v in os.environ.get("ZM_VARIANTS", "").split(",") if v] + sources = { + variant: source + for variant, source in QWEN_IMAGE_21_SOURCES.items() + if not selected or variant in selected + } + + for variant, source in sources.items(): + print(f"\n{'=' * 60}\nConverting: {variant} <- {source}\n{'=' * 60}") + with keras.device(device): + model, config = transfer_qwen_image_21(source, token=token, dtype=dtype) + + itemsize = 2 if "16" in dtype else 4 + n_bytes = sum(int(np.prod(w.shape)) * itemsize for w in model.weights) + stem = os.path.join(OUT_DIR, variant.replace("-", "_").replace(".", "_")) + if n_bytes > MAX_SHARD_GB * 1024**3: + out = f"{stem}.weights.json" + model.save_weights(out, max_shard_size=MAX_SHARD_GB) + else: + out = f"{stem}.weights.h5" + model.save_weights(out) + print(f" saved -> {out} ({n_bytes / 1024**3:.2f} GB {dtype})") + + del model + keras.backend.clear_session() + gc.collect() diff --git a/zeromodels/models/qwen_image_21/qwen_image_21_config.py b/zeromodels/models/qwen_image_21/qwen_image_21_config.py new file mode 100644 index 00000000..4d2e289b --- /dev/null +++ b/zeromodels/models/qwen_image_21/qwen_image_21_config.py @@ -0,0 +1,244 @@ +from zeromodels.base import BaseConfig +from zeromodels.models.qwen3_vl.qwen3_vl_config import Qwen3VLTextConfig + +DEFAULT_LATENTS_MEAN = ( + 0.5126, + 0.7721, + -0.0631, + 1.3506, + -0.7855, + -2.1025, + -0.3458, + 1.3722, + 1.8873, + -1.7177, + -0.6510, + 0.2732, + 0.7562, + -0.6163, + -1.0277, + 3.8363, + 2.0210, + 0.0472, + 0.9320, + 2.0087, + 2.4954, + -0.1391, + -1.4249, + 1.8464, + -0.5236, + 1.2826, + 3.7046, + -1.3035, + 2.7286, + -1.4518, + -1.9036, + -1.9955, + -0.0342, + -1.0265, + -0.7636, + 3.0555, + 0.0746, + -3.0751, + -0.1076, + 1.7376, + -1.0914, + -1.9435, + -0.2784, + -1.3680, + 0.4809, + -0.4433, + 0.3764, + 0.5729, + -2.0595, + 1.0960, + -1.3260, + -2.0211, + -5.0179, + 0.5275, + 4.0162, + 1.8505, + 0.3026, + 1.9373, + 1.4937, + 0.2632, + 0.5547, + -1.7121, + -0.1562, + 0.0304, +) +DEFAULT_LATENTS_STD = ( + 3.2001, + 3.2936, + 3.4321, + 3.0091, + 3.1061, + 4.0379, + 4.0705, + 3.7910, + 3.0785, + 3.6500, + 3.9308, + 3.0904, + 2.8778, + 3.7675, + 3.7320, + 5.0756, + 3.2864, + 4.0397, + 3.1317, + 4.0443, + 2.9249, + 3.9454, + 3.0988, + 4.2489, + 3.4896, + 3.8513, + 3.9323, + 3.4719, + 3.7498, + 4.2830, + 3.5694, + 4.2467, + 3.9037, + 3.2947, + 5.0770, + 3.5075, + 3.2700, + 3.4767, + 2.8063, + 5.1125, + 3.5327, + 4.7833, + 3.1286, + 4.1819, + 3.8527, + 3.8312, + 3.5605, + 4.3875, + 3.9624, + 4.0168, + 3.5643, + 4.0550, + 5.5614, + 4.2963, + 4.4080, + 3.4959, + 3.8747, + 3.7608, + 3.5735, + 3.1490, + 3.7662, + 3.6746, + 3.4563, + 3.8161, +) + + +class QwenImage21VAEConfig(BaseConfig): + """Configuration for :class:`AutoencoderKLQwenImage21`. + + Defaults match ``Qwen/Qwen-Image-2.1`` ``vae/config.json`` (64-channel latent, + 16× spatial compression, residual encoder/decoder, 4-channel RGBA IO). + """ + + model_type = "autoencoder_kl_qwen_image_21" + + base_dim: int = 96 + decoder_base_dim: int = 144 + z_dim: int = 64 + dim_mult: tuple = (1, 2, 4, 8, 8) + num_res_blocks: int = 2 + attn_scales: tuple = () + temperal_downsample: tuple = (False, True, True, True) + dropout: float = 0.0 + input_channels: int = 4 + out_channels: int = 4 + is_residual: bool = True + scale_factor_spatial: int = 16 + scale_factor_temporal: int = 8 + latents_mean: tuple = DEFAULT_LATENTS_MEAN + latents_std: tuple = DEFAULT_LATENTS_STD + sample_size: int = 1024 + + +class QwenImage21TransformerConfig(BaseConfig): + """Configuration for :class:`QwenImage21Transformer2DModel`. + + Defaults match ``Qwen/Qwen-Image-2.1`` ``transformer/config.json`` (32-layer + single-stream DiT, 32 heads × 128, context width 4096, unpatched latents). + """ + + model_type = "qwen_image_21_transformer_2d" + + patch_size: int = 1 + in_channels: int = 64 + out_channels: int = 64 + num_layers: int = 32 + attention_head_dim: int = 128 + num_attention_heads: int = 32 + context_in_dim: int = 4096 + mlp_ratio: int = 3 + axes_dims_rope: tuple = (16, 56, 56) + eps: float = 1e-6 + causal_condition: bool = True + sample_size: int = 64 + text_seq_len: int = 512 + + +class QwenImage21TextConfig(Qwen3VLTextConfig): + """Qwen3-VL text tower used as the Qwen-Image-2.1 prompt encoder. + + Defaults match ``Qwen/Qwen-Image-2.1`` ``text_encoder/config.json`` text half + (36 layers, 4096-d, 32 heads / 8 KV). + """ + + model_type = "qwen_image_21_text" + + vocab_size: int = 151936 + embed_dim: int = 4096 + mlp_dim: int = 12288 + num_layers: int = 36 + num_heads: int = 32 + num_kv_heads: int = 8 + head_dim: int = 128 + norm_eps: float = 1e-6 + rope_theta: float = 5000000.0 + mrope_section: tuple = (24, 20, 20) + tie_embeddings: bool = False + max_seq_len: int = 1024 + + +class QwenImage21Config(BaseConfig): + """Configuration for :class:`QwenImage21Model` / :class:`QwenImage21TextToImage`. + + One hosted container: single-stream DiT + Qwen-Image-2.1 VAE + Qwen3-VL text + encoder. Nested serialize; flat constructor with ``transformer_`` / ``vae_`` / + ``text_`` prefixes. + """ + + model_type = "qwen_image_21" + + sub_configs = { + "transformer_config": QwenImage21TransformerConfig, + "vae_config": QwenImage21VAEConfig, + "text_config": QwenImage21TextConfig, + } + sub_config_prefixes = { + "transformer_config": "transformer_", + "vae_config": "vae_", + "text_config": "text_", + } + group_extras = {"text_config": ("vocab_size", "max_seq_len")} + + transformer_config: QwenImage21TransformerConfig | dict | None = None + vae_config: QwenImage21VAEConfig | dict | None = None + text_config: QwenImage21TextConfig | dict | None = None + scheduler_config: dict | None = None + prompt_template_encode_start_idx: int = 14 + max_sequence_length: int = 512 + default_sample_size: int = 64 + bos_token_id: int = 151643 + eos_token_id: int = 151645 + pad_token_id: int = 151643 + image_token_id: int = 151655 diff --git a/zeromodels/models/qwen_image_21/qwen_image_21_layers.py b/zeromodels/models/qwen_image_21/qwen_image_21_layers.py new file mode 100644 index 00000000..25cbad0c --- /dev/null +++ b/zeromodels/models/qwen_image_21/qwen_image_21_layers.py @@ -0,0 +1,585 @@ +from __future__ import annotations + +import math + +import keras +import numpy as np +from keras import layers, ops + +from zeromodels.base.base_attention import fused_attention +from zeromodels.models.qwen_image.qwen_image_layers import ( + QwenImageRMSNorm, + apply_rotary_emb_qwen, +) +from zeromodels.models.stable_diffusion.stable_diffusion_layers import ( + safe_name, + timestep_embedding, +) + +NORM_EPS = 1e-6 +MASK_NEG = -1e4 +IMG_TOKENS_PER_SLOT = 4 + + +def rope_params(index, dim, theta): + freqs = np.outer( + np.asarray(index, dtype=np.float64), + 1.0 / np.power(float(theta), np.arange(0, dim, 2, dtype=np.float64) / dim), + ) + return freqs.astype(np.float32) + + +def build_qwenimage21_rope_angles(img_shapes, image_pad_mask, axes_dim, theta=10000): + """Return RoPE angles ``[S, sum(axes_dim)//2]`` for a joint text/image layout. + + Mirrors Diffusers ``QwenImage21Rope`` (real-valued angles, not complex). + ``image_pad_mask`` is a 1-D bool/int array of length ``S`` with ``True`` at + image-token positions; ``img_shapes`` is a list of ``(frame, height, width)``. + """ + image_pad_mask = np.asarray(image_pad_mask, dtype=bool) + total_len = int(image_pad_mask.shape[0]) + axes_dim = list(axes_dim) + + pos_index = np.arange(8192) + neg_index = np.arange(1024)[::-1] * -1 - 1 + freqs = [ + np.concatenate( + [rope_params(pos_index, dim, theta), rope_params(neg_index, dim, theta)], + axis=0, + ) + for dim in axes_dim + ] + + frame_index, image_height_index, image_width_index = [], [], [] + cursor, position = 0, 0 + is_image = image_pad_mask.tolist() + + for _, height, width in img_shapes: + block_start = is_image.index(True, cursor) + text_len = block_start - cursor + frame_index.extend(range(position, position + text_len)) + position += text_len + + cursor = block_start + height * width + frame_index.extend([position] * (height * width)) + position += max(height, width) + + image_height_index.extend( + [ + h + for h in range(-(height - height // 2), height // 2) + for _ in range(width) + ] + ) + image_width_index.extend( + [w for _ in range(height) for w in range(-(width - width // 2), width // 2)] + ) + + if cursor < total_len: + frame_index.extend(range(position, position + total_len - cursor)) + + frame_index = np.asarray(frame_index, dtype=np.int64) + height_index = frame_index.copy() + width_index = frame_index.copy() + height_index[image_pad_mask] = np.asarray(image_height_index, dtype=np.int64) + width_index[image_pad_mask] = np.asarray(image_width_index, dtype=np.int64) + + return np.concatenate( + [freqs[0][frame_index], freqs[1][height_index], freqs[2][width_index]], + axis=-1, + ) + + +def build_token_metadata(image_pad_mask, img_shapes): + """Label joint-sequence tokens with image-block ids (Diffusers helper).""" + image_pad_mask = np.asarray(image_pad_mask, dtype=bool) + image_positions = np.flatnonzero(image_pad_mask) + block_lengths = [int(math.prod(shape)) for shape in img_shapes] + if sum(block_lengths) != image_positions.size: + raise ValueError( + f"img_shapes accounts for {sum(block_lengths)} image tokens but " + f"image_pad_mask marks {image_positions.size}." + ) + image_ids = np.full(image_pad_mask.shape, -1, dtype=np.int32) + block_ids = np.repeat(np.arange(len(block_lengths), dtype=np.int32), block_lengths) + image_ids[image_positions] = block_ids + target_token_mask = np.zeros_like(image_pad_mask, dtype=bool) + target_token_mask[image_positions[-block_lengths[-1] :]] = True + return image_ids, target_token_mask + + +def build_block_causal_additive_mask(image_ids, key_valid=None): + """Dense additive attention mask for block-causal attention. + + Allowed when ``(q >= kv) or same_image_block``, and ``key_valid[kv]``. + Returns ``(1, 1, S, S)`` float32 mask with ``0`` / ``MASK_NEG``. + """ + image_ids = np.asarray(image_ids, dtype=np.int32) + seq = image_ids.shape[0] + q = np.arange(seq)[:, None] + kv = np.arange(seq)[None, :] + same = (image_ids[:, None] == image_ids[None, :]) & (image_ids[:, None] >= 0) + allowed = (q >= kv) | same + if key_valid is not None: + key_valid = np.asarray(key_valid, dtype=bool) + allowed = allowed & key_valid[None, :] + mask = np.where(allowed, 0.0, MASK_NEG).astype(np.float32) + return mask[None, None] + + +@keras.saving.register_keras_serializable(package="zeromodels") +class QwenImage21ZeroCenterRMSNorm(layers.Layer): + """RMSNorm with zero-centered weight (effective scale ``weight + 1``).""" + + def __init__(self, eps=NORM_EPS, module_path=None, **kwargs): + if module_path is not None: + kwargs.setdefault("name", safe_name(module_path)) + super().__init__(**kwargs) + self.eps = eps + self.module_path = module_path + + def build(self, input_shape): + self.weight = self.add_weight( + name="weight", + shape=(int(input_shape[-1]),), + initializer="zeros", + trainable=True, + ) + self.built = True + + def call(self, x): + dtype = x.dtype + x32 = ops.cast(x, "float32") + rrms = ops.rsqrt(ops.mean(ops.square(x32), axis=-1, keepdims=True) + self.eps) + out = x32 * rrms * (ops.cast(self.weight, "float32") + 1.0) + return ops.cast(out, dtype) + + def compute_output_shape(self, input_shape): + return tuple(input_shape) + + def get_config(self): + config = super().get_config() + config.update({"eps": self.eps, "module_path": self.module_path}) + return config + + +@keras.saving.register_keras_serializable(package="zeromodels") +class QwenImage21TimestepProjEmbeddings(layers.Layer): + """Diffusers ``QwenImage21TimestepProjEmbeddings`` (bias-free MLP).""" + + def __init__( + self, + embedding_dim, + module_path="time_text_embed", + time_freq_dim=256, + scale=1000.0, + **kwargs, + ): + kwargs.setdefault("name", safe_name(module_path)) + super().__init__(**kwargs) + self.embedding_dim = embedding_dim + self.module_path = module_path + self.time_freq_dim = time_freq_dim + self.scale = scale + embedder = f"{module_path}.timestep_embedder" + self.linear_1 = layers.Dense( + embedding_dim, use_bias=False, name=safe_name(f"{embedder}.linear_1") + ) + self.linear_2 = layers.Dense( + embedding_dim, use_bias=False, name=safe_name(f"{embedder}.linear_2") + ) + + def build(self, timestep_shape): + freq_shape = (timestep_shape[0], self.time_freq_dim) + self.linear_1.build(freq_shape) + self.linear_2.build((timestep_shape[0], self.embedding_dim)) + self.built = True + + def call(self, timestep): + t = ops.cast(timestep, "float32") * self.scale + emb = timestep_embedding( + t, + self.time_freq_dim, + flip_sin_to_cos=True, + downscale_freq_shift=0, + ) + emb = self.linear_1(emb) + emb = ops.silu(emb) + return self.linear_2(emb) + + def compute_output_shape(self, timestep_shape): + return (timestep_shape[0], self.embedding_dim) + + def get_config(self): + config = super().get_config() + config.update( + { + "embedding_dim": self.embedding_dim, + "module_path": self.module_path, + "time_freq_dim": self.time_freq_dim, + "scale": self.scale, + } + ) + return config + + +@keras.saving.register_keras_serializable(package="zeromodels") +class QwenImage21TextProjection(layers.Layer): + """Diffusers ``QwenImage21TextProjection``.""" + + def __init__( + self, context_in_dim, hidden_size, eps=NORM_EPS, module_path="txt_in", **kwargs + ): + kwargs.setdefault("name", safe_name(module_path)) + super().__init__(**kwargs) + self.context_in_dim = context_in_dim + self.hidden_size = hidden_size + self.eps = eps + self.module_path = module_path + self.text_norm = QwenImage21ZeroCenterRMSNorm( + eps=eps, module_path=f"{module_path}.text_norm" + ) + self.in_layer = layers.Dense( + hidden_size, use_bias=False, name=safe_name(f"{module_path}.in_layer") + ) + self.out_layer = layers.Dense( + hidden_size, use_bias=False, name=safe_name(f"{module_path}.out_layer") + ) + + def build(self, input_shape): + self.text_norm.build(input_shape) + self.in_layer.build(input_shape) + self.out_layer.build((*input_shape[:-1], self.hidden_size)) + self.built = True + + def call(self, x): + x = self.text_norm(x) + x = self.in_layer(x) + x = keras.activations.gelu(x, approximate=True) + return self.out_layer(x) + + def compute_output_shape(self, input_shape): + return (*tuple(input_shape)[:-1], self.hidden_size) + + def get_config(self): + config = super().get_config() + config.update( + { + "context_in_dim": self.context_in_dim, + "hidden_size": self.hidden_size, + "eps": self.eps, + "module_path": self.module_path, + } + ) + return config + + +@keras.saving.register_keras_serializable(package="zeromodels") +class QwenImage21SwiGLUFeedForward(layers.Layer): + """Diffusers ``QwenImage21SwiGLUFeedForward``.""" + + def __init__(self, hidden_size, mlp_hidden_size, module_path, **kwargs): + kwargs.setdefault("name", safe_name(module_path)) + super().__init__(**kwargs) + self.hidden_size = hidden_size + self.mlp_hidden_size = mlp_hidden_size + self.module_path = module_path + self.proj = layers.Dense( + mlp_hidden_size, use_bias=False, name=safe_name(f"{module_path}.proj") + ) + self.gate_layer = layers.Dense( + mlp_hidden_size, use_bias=False, name=safe_name(f"{module_path}.gate_layer") + ) + self.out = layers.Dense( + hidden_size, use_bias=False, name=safe_name(f"{module_path}.out") + ) + + def build(self, input_shape): + self.proj.build(input_shape) + self.gate_layer.build(input_shape) + self.out.build((*input_shape[:-1], self.mlp_hidden_size)) + self.built = True + + def call(self, x): + return self.out(ops.silu(self.gate_layer(x)) * self.proj(x)) + + def compute_output_shape(self, input_shape): + return (*tuple(input_shape)[:-1], self.hidden_size) + + def get_config(self): + config = super().get_config() + config.update( + { + "hidden_size": self.hidden_size, + "mlp_hidden_size": self.mlp_hidden_size, + "module_path": self.module_path, + } + ) + return config + + +@keras.saving.register_keras_serializable(package="zeromodels") +class QwenImage21AdaLayerNormContinuous(layers.Layer): + """Final adaptive LayerNorm (scale only, no shift).""" + + def __init__( + self, + embedding_dim, + conditioning_embedding_dim, + eps=NORM_EPS, + module_path="norm_out", + **kwargs, + ): + kwargs.setdefault("name", safe_name(module_path)) + super().__init__(**kwargs) + self.embedding_dim = embedding_dim + self.conditioning_embedding_dim = conditioning_embedding_dim + self.eps = eps + self.module_path = module_path + self.linear = layers.Dense( + embedding_dim, use_bias=False, name=safe_name(f"{module_path}.linear") + ) + self.norm = layers.LayerNormalization( + axis=-1, + epsilon=eps, + center=False, + scale=False, + name=safe_name(f"{module_path}.norm"), + ) + + def build(self, input_shapes): + hidden_shape = input_shapes[0] + cond_shape = input_shapes[1] + self.norm.build(hidden_shape) + self.linear.build(cond_shape) + self.built = True + + def call(self, inputs): + hidden_states, conditioning, target_token_mask = inputs + scale = self.linear(ops.silu(ops.cast(conditioning, hidden_states.dtype))) + scale = select_modulation_rows(scale, target_token_mask) + return self.norm(hidden_states) * (1.0 + scale) + + def compute_output_shape(self, input_shapes): + return tuple(input_shapes[0]) + + def get_config(self): + config = super().get_config() + config.update( + { + "embedding_dim": self.embedding_dim, + "conditioning_embedding_dim": self.conditioning_embedding_dim, + "eps": self.eps, + "module_path": self.module_path, + } + ) + return config + + +def select_modulation_rows(params, target_token_mask): + """Broadcast per-sample modulation over tokens (causal_condition aware).""" + if target_token_mask is None: + return ops.expand_dims(params, 1) + real = ops.expand_dims(params[:-1], 1) + zero = ops.expand_dims(params[-1:], 0) + mask = ops.reshape(ops.cast(target_token_mask, "bool"), (1, -1, 1)) + return ops.where(mask, real, zero) + + +@keras.saving.register_keras_serializable(package="zeromodels") +class QwenImage21Attention(layers.Layer): + """Single-stream attention (bias-free QKV / out, per-head QK RMSNorm).""" + + def __init__(self, dim, heads, dim_head, eps=NORM_EPS, module_path=None, **kwargs): + if module_path is not None: + kwargs.setdefault("name", safe_name(module_path)) + super().__init__(**kwargs) + self.dim = dim + self.heads = heads + self.dim_head = dim_head + self.eps = eps + self.module_path = module_path + inner = heads * dim_head + path = module_path or "attn" + self.to_q = layers.Dense(inner, use_bias=False, name=safe_name(f"{path}.to_q")) + self.to_k = layers.Dense(inner, use_bias=False, name=safe_name(f"{path}.to_k")) + self.to_v = layers.Dense(inner, use_bias=False, name=safe_name(f"{path}.to_v")) + self.to_out = layers.Dense( + dim, use_bias=False, name=safe_name(f"{path}.to_out.0") + ) + self.norm_q = QwenImageRMSNorm(eps=eps, module_path=f"{path}.norm_q") + self.norm_k = QwenImageRMSNorm(eps=eps, module_path=f"{path}.norm_k") + + def build(self, input_shape): + self.to_q.build(input_shape) + self.to_k.build(input_shape) + self.to_v.build(input_shape) + head_shape = (*input_shape[:-1], self.heads, self.dim_head) + self.norm_q.build(head_shape) + self.norm_k.build(head_shape) + self.to_out.build((*input_shape[:-1], self.heads * self.dim_head)) + self.built = True + + def call(self, hidden_states, rotary_emb=None, attention_mask=None): + query = self.to_q(hidden_states) + key = self.to_k(hidden_states) + value = self.to_v(hidden_states) + + def unflatten(x): + shape = ops.shape(x) + return ops.reshape(x, (shape[0], shape[1], self.heads, self.dim_head)) + + query = self.norm_q(unflatten(query)) + key = self.norm_k(unflatten(key)) + value = unflatten(value) + + if rotary_emb is not None: + query = apply_rotary_emb_qwen(query, rotary_emb, use_real=True) + key = apply_rotary_emb_qwen(key, rotary_emb, use_real=True) + + query = ops.transpose(query, (0, 2, 1, 3)) + key = ops.transpose(key, (0, 2, 1, 3)) + value = ops.transpose(value, (0, 2, 1, 3)) + scale = self.dim_head**-0.5 + out = fused_attention(query, key, value, scale, attention_mask=attention_mask) + out = ops.transpose(out, (0, 2, 1, 3)) + shape = ops.shape(out) + out = ops.reshape(out, (shape[0], shape[1], self.heads * self.dim_head)) + return self.to_out(out) + + def compute_output_shape(self, input_shape): + return tuple(input_shape) + + def compute_output_spec(self, hidden_states, rotary_emb=None, attention_mask=None): + return keras.KerasTensor(hidden_states.shape, dtype=self.compute_dtype) + + def get_config(self): + config = super().get_config() + config.update( + { + "dim": self.dim, + "heads": self.heads, + "dim_head": self.dim_head, + "eps": self.eps, + "module_path": self.module_path, + } + ) + return config + + +@keras.saving.register_keras_serializable(package="zeromodels") +class QwenImage21TransformerBlock(layers.Layer): + """Single-stream block with shared external modulation.""" + + def __init__( + self, + dim, + num_attention_heads, + attention_head_dim, + mlp_ratio=3, + eps=NORM_EPS, + module_path=None, + **kwargs, + ): + if module_path is not None: + kwargs.setdefault("name", safe_name(module_path)) + super().__init__(**kwargs) + self.dim = dim + self.num_attention_heads = num_attention_heads + self.attention_head_dim = attention_head_dim + self.mlp_ratio = mlp_ratio + self.eps = eps + self.module_path = module_path + path = module_path or "block" + self.img_norm1 = layers.LayerNormalization( + axis=-1, + epsilon=eps, + center=False, + scale=False, + name=safe_name(f"{path}.img_norm1"), + ) + self.attn = QwenImage21Attention( + dim, + num_attention_heads, + attention_head_dim, + eps=eps, + module_path=f"{path}.attn", + ) + self.img_norm2 = layers.LayerNormalization( + axis=-1, + epsilon=eps, + center=False, + scale=False, + name=safe_name(f"{path}.img_norm2"), + ) + self.img_mlp = QwenImage21SwiGLUFeedForward( + dim, dim * mlp_ratio, module_path=f"{path}.img_mlp" + ) + + def build(self, input_shape): + self.img_norm1.build(input_shape) + self.attn.build(input_shape) + self.img_norm2.build(input_shape) + self.img_mlp.build(input_shape) + self.built = True + + def modulate(self, hidden_states, mod_params, target_token_mask): + scale, gate = ops.split(mod_params, 2, axis=-1) + scale = select_modulation_rows(scale, target_token_mask) + gate = select_modulation_rows(gate, target_token_mask) + return hidden_states * (1.0 + scale), gate + + def call( + self, + hidden_states, + modulation, + rotary_emb=None, + attention_mask=None, + target_token_mask=None, + ): + mod1, mod2 = ops.split(modulation, 2, axis=-1) + img_modulated, img_gate1 = self.modulate( + self.img_norm1(hidden_states), mod1, target_token_mask + ) + attn_output = self.attn( + img_modulated, rotary_emb=rotary_emb, attention_mask=attention_mask + ) + hidden_states = hidden_states + ops.tanh(img_gate1) * attn_output + + img_modulated2, img_gate2 = self.modulate( + self.img_norm2(hidden_states), mod2, target_token_mask + ) + hidden_states = hidden_states + ops.tanh(img_gate2) * self.img_mlp( + img_modulated2 + ) + + if keras.backend.standardize_dtype(hidden_states.dtype) == "float16": + hidden_states = ops.clip(hidden_states, -65504.0, 65504.0) + return hidden_states + + def compute_output_shape(self, input_shape): + return tuple(input_shape) + + def compute_output_spec( + self, + hidden_states, + modulation, + rotary_emb=None, + attention_mask=None, + target_token_mask=None, + ): + return keras.KerasTensor(hidden_states.shape, dtype=self.compute_dtype) + + def get_config(self): + config = super().get_config() + config.update( + { + "dim": self.dim, + "num_attention_heads": self.num_attention_heads, + "attention_head_dim": self.attention_head_dim, + "mlp_ratio": self.mlp_ratio, + "eps": self.eps, + "module_path": self.module_path, + } + ) + return config diff --git a/zeromodels/models/qwen_image_21/qwen_image_21_model.py b/zeromodels/models/qwen_image_21/qwen_image_21_model.py new file mode 100644 index 00000000..d9e97445 --- /dev/null +++ b/zeromodels/models/qwen_image_21/qwen_image_21_model.py @@ -0,0 +1,857 @@ +from __future__ import annotations + +import keras +import numpy as np +from keras import layers, ops + +from zeromodels.base import BaseDiffusion, BaseModel, CausalMask +from zeromodels.base.base_mixin import inference_scope +from zeromodels.base.base_scheduler import ( + FlowMatchEulerDiscreteScheduler, + get_scheduler, +) +from zeromodels.models.qwen3_vl.qwen3_vl_model import ( + Qwen3VLTextModel, + qwen3_text_cos_sin, +) +from zeromodels.models.qwen_image_21.qwen_image_21_config import ( + DEFAULT_LATENTS_MEAN, + DEFAULT_LATENTS_STD, + QwenImage21Config, + QwenImage21TextConfig, + QwenImage21TransformerConfig, + QwenImage21VAEConfig, +) +from zeromodels.models.qwen_image_21.qwen_image_21_layers import ( + IMG_TOKENS_PER_SLOT, + MASK_NEG, + QwenImage21AdaLayerNormContinuous, + QwenImage21TextProjection, + QwenImage21TimestepProjEmbeddings, + QwenImage21TransformerBlock, + build_block_causal_additive_mask, + build_qwenimage21_rope_angles, + build_token_metadata, +) +from zeromodels.models.qwen_image_21.qwen_image_21_vae import ( + QwenImage21CausalConv, + QwenImage21Decoder3d, + QwenImage21Encoder3d, +) +from zeromodels.models.stable_diffusion.stable_diffusion_layers import safe_name + +QWEN_IMAGE_21_HUB_SIBLINGS = frozenset({"QwenImage21Model", "QwenImage21TextToImage"}) + + +def pack_latents(latents, height, width): + """Flatten ``(B, H, W, C)`` → ``(B, H*W, C)`` (2.1 consumes latents unpatched).""" + batch = ops.shape(latents)[0] + channels = ops.shape(latents)[-1] + return ops.reshape(latents, (batch, height * width, channels)) + + +def unpack_latents(latents, height, width, channels): + """Unflatten ``(B, H*W, C)`` → ``(B, H, W, C)``.""" + batch = ops.shape(latents)[0] + return ops.reshape(latents, (batch, height, width, channels)) + + +@keras.saving.register_keras_serializable(package="zeromodels") +class AutoencoderKLQwenImage21(BaseModel): + """Qwen-Image-2.1 VAE (Diffusers ``AutoencoderKLQwenImage21``), channels-last. + + 64-channel latents, 16× spatial compression, residual encoder/decoder, RGBA + (4-channel) IO. Graph built for ``sample_size``; weights are resolution- + independent. + """ + + config_class = QwenImage21VAEConfig + HF_MODEL_TYPE = None + + def __init__( + self, + base_dim=96, + decoder_base_dim=144, + z_dim=64, + dim_mult=(1, 2, 4, 8, 8), + num_res_blocks=2, + attn_scales=(), + temperal_downsample=(False, True, True, True), + dropout=0.0, + input_channels=4, + out_channels=4, + is_residual=True, + scale_factor_spatial=16, + latents_mean=DEFAULT_LATENTS_MEAN, + latents_std=DEFAULT_LATENTS_STD, + sample_size=1024, + name="AutoencoderKLQwenImage21", + **kwargs, + ): + dim_mult = tuple(dim_mult) + temperal_downsample = tuple(temperal_downsample) + attn_scales = tuple(attn_scales) + latents_mean = tuple(latents_mean) + latents_std = tuple(latents_std) + temperal_upsample = tuple(reversed(temperal_downsample)) + decoder_base_dim = ( + int(decoder_base_dim) if decoder_base_dim is not None else int(base_dim) + ) + + h_img, w_img = ( + sample_size + if isinstance(sample_size, (tuple, list)) + else (sample_size, sample_size) + ) + spatial_compression_ratio = int(scale_factor_spatial) + h_lat, w_lat = ( + h_img // spatial_compression_ratio, + w_img // spatial_compression_ratio, + ) + + encoder = QwenImage21Encoder3d( + dim=base_dim, + z_dim=z_dim * 2, + dim_mult=dim_mult, + num_res_blocks=num_res_blocks, + attn_scales=attn_scales, + temperal_downsample=temperal_downsample, + dropout=dropout, + input_channels=input_channels, + is_residual=is_residual, + module_path="encoder", + ) + decoder = QwenImage21Decoder3d( + dim=decoder_base_dim, + z_dim=z_dim, + dim_mult=dim_mult, + num_res_blocks=num_res_blocks, + attn_scales=attn_scales, + temperal_upsample=temperal_upsample, + dropout=dropout, + out_channels=out_channels, + is_residual=is_residual, + module_path="decoder", + ) + quant_conv = QwenImage21CausalConv( + z_dim * 2, kernel_size=1, padding=0, module_path="quant_conv" + ) + post_quant_conv = QwenImage21CausalConv( + z_dim, kernel_size=1, padding=0, module_path="post_quant_conv" + ) + + image_in = layers.Input(shape=(h_img, w_img, input_channels), name="image") + latent_in = layers.Input(shape=(h_lat, w_lat, z_dim), name="latent") + + image_5d = ops.expand_dims(image_in, axis=1) + moments_5d = quant_conv(encoder(image_5d)) + moments = ops.squeeze(moments_5d, axis=1) + + latent_5d = ops.expand_dims(latent_in, axis=1) + decoded_5d = decoder(post_quant_conv(latent_5d)) + decoded = ops.squeeze(decoded_5d, axis=1) + decoded = ops.clip(decoded, -1.0, 1.0) + + super().__init__( + inputs={"image": image_in, "latent": latent_in}, + outputs={"moments": moments, "sample": decoded}, + name=name, + **kwargs, + ) + + self.base_dim = base_dim + self.decoder_base_dim = decoder_base_dim + self.z_dim = z_dim + self.dim_mult = dim_mult + self.num_res_blocks = num_res_blocks + self.attn_scales = attn_scales + self.temperal_downsample = temperal_downsample + self.temperal_upsample = temperal_upsample + self.dropout = dropout + self.input_channels = input_channels + self.out_channels = out_channels + self.is_residual = is_residual + self.latents_mean = latents_mean + self.latents_std = latents_std + self.sample_size = sample_size + self.spatial_compression_ratio = spatial_compression_ratio + self.vae_scale_factor = spatial_compression_ratio + self.encoder = encoder + self.decoder = decoder + self.quant_conv = quant_conv + self.post_quant_conv = post_quant_conv + + def to_ndhwc(self, x, is_latent=False): + static_ndim = len(x.shape) + if static_ndim == 4: + return ops.expand_dims(x, axis=1) + if static_ndim == 5: + c1 = int(x.shape[1]) if x.shape[1] is not None else None + c_last = int(x.shape[-1]) if x.shape[-1] is not None else None + expect = self.z_dim if is_latent else self.input_channels + if c1 == expect and c_last != expect: + return ops.transpose(x, (0, 2, 3, 4, 1)) + return x + raise ValueError(f"Expected 4D or 5D tensor, got shape {x.shape}") + + def encode(self, x, sample=False, seed=None): + was_4d = len(x.shape) == 4 + x5 = self.to_ndhwc(x, is_latent=False) + moments = self.quant_conv(self.encoder(x5)) + mean, logvar = ops.split(moments, 2, axis=-1) + if sample: + logvar = ops.clip(logvar, -30.0, 20.0) + std = ops.exp(0.5 * logvar) + noise = keras.random.normal(ops.shape(mean), dtype=mean.dtype, seed=seed) + z = mean + std * noise + else: + z = mean + if was_4d: + z = ops.squeeze(z, axis=1) + return z + + def decode(self, z): + was_4d = len(z.shape) == 4 + z5 = self.to_ndhwc(z, is_latent=True) + out = self.decoder(self.post_quant_conv(z5)) + out = ops.clip(out, -1.0, 1.0) + if was_4d: + out = ops.squeeze(out, axis=1) + return out + + @classmethod + def kwargs_from_diffusers_config(cls, cfg): + return { + "base_dim": cfg.get("base_dim", 96), + "decoder_base_dim": cfg.get("decoder_base_dim", 144), + "z_dim": cfg.get("z_dim", 64), + "dim_mult": tuple(cfg.get("dim_mult", (1, 2, 4, 8, 8))), + "num_res_blocks": cfg.get("num_res_blocks", 2), + "attn_scales": tuple(cfg.get("attn_scales") or ()), + "temperal_downsample": tuple( + cfg.get("temperal_downsample", (False, True, True, True)) + ), + "dropout": cfg.get("dropout", 0.0), + "input_channels": cfg.get("in_channels", 4), + "out_channels": cfg.get("out_channels", 4), + "is_residual": bool(cfg.get("is_residual", True)), + "scale_factor_spatial": cfg.get("scale_factor_spatial", 16), + "latents_mean": tuple(cfg.get("latents_mean", DEFAULT_LATENTS_MEAN)), + "latents_std": tuple(cfg.get("latents_std", DEFAULT_LATENTS_STD)), + } + + +@keras.saving.register_keras_serializable(package="zeromodels") +class QwenImage21Transformer2DModel(BaseModel): + """Qwen-Image-2.1 denoiser (Diffusers ``QwenImage21Transformer2DModel``). + + Single-stream block-causal DiT over unpatched latents and Qwen3-VL text + features. Built for a fixed T2I layout (``sample_size`` latent side + + ``text_seq_len``). Output ``{"sample": target velocity}``. + """ + + HF_MODEL_TYPE = None + config_class = QwenImage21TransformerConfig + + def __init__( + self, + patch_size=1, + in_channels=64, + out_channels=64, + num_layers=32, + attention_head_dim=128, + num_attention_heads=32, + context_in_dim=4096, + mlp_ratio=3, + axes_dims_rope=(16, 56, 56), + eps=1e-6, + causal_condition=True, + sample_size=64, + text_seq_len=512, + name="QwenImage21Transformer2DModel", + **kwargs, + ): + keras_kwargs = {k: kwargs.pop(k) for k in ("trainable", "dtype") if k in kwargs} + axes_dims_rope = tuple(axes_dims_rope) + inner_dim = num_attention_heads * attention_head_dim + sample_h = ( + sample_size[0] if isinstance(sample_size, (tuple, list)) else sample_size + ) + sample_w = ( + sample_size[1] if isinstance(sample_size, (tuple, list)) else sample_size + ) + assert patch_size == 1, "Qwen-Image-2.1 consumes latents unpatched" + img_seq = sample_h * sample_w + img_shapes = [(1, sample_h, sample_w)] + + target_slots = (sample_h * sample_w) // IMG_TOKENS_PER_SLOT + slot_mask = np.concatenate( + [ + np.zeros(text_seq_len, dtype=bool), + np.ones(target_slots, dtype=bool), + ] + ) + image_pad_mask = np.repeat( + slot_mask, np.where(slot_mask, IMG_TOKENS_PER_SLOT, 1) + ) + joint_seq = int(image_pad_mask.shape[0]) + image_ids, target_token_mask = build_token_metadata(image_pad_mask, img_shapes) + rope_angles = build_qwenimage21_rope_angles( + img_shapes, image_pad_mask, axes_dims_rope + ) + attn_mask_np = build_block_causal_additive_mask(image_ids) + + time_text_embed = QwenImage21TimestepProjEmbeddings( + embedding_dim=inner_dim, module_path="time_text_embed" + ) + txt_in = QwenImage21TextProjection( + context_in_dim, inner_dim, eps=eps, module_path="txt_in" + ) + img_in = layers.Dense(inner_dim, use_bias=False, name=safe_name("img_in")) + modulation = layers.Dense( + 4 * inner_dim, use_bias=False, name=safe_name("modulation.1") + ) + blocks = [ + QwenImage21TransformerBlock( + dim=inner_dim, + num_attention_heads=num_attention_heads, + attention_head_dim=attention_head_dim, + mlp_ratio=mlp_ratio, + eps=eps, + module_path=f"transformer_blocks.{i}", + ) + for i in range(num_layers) + ] + norm_out = QwenImage21AdaLayerNormContinuous( + inner_dim, inner_dim, eps=eps, module_path="norm_out" + ) + proj_out = layers.Dense( + patch_size * patch_size * (out_channels or in_channels), + use_bias=False, + name=safe_name("proj_out"), + ) + + sample_in = layers.Input(shape=(img_seq, in_channels), name="sample") + timestep_in = layers.Input(shape=(), name="timestep") + enc_in = layers.Input( + shape=(text_seq_len, context_in_dim), name="encoder_hidden_states" + ) + enc_mask_in = layers.Input( + shape=(text_seq_len,), dtype="int32", name="encoder_hidden_states_mask" + ) + + hidden = img_in(sample_in) + encoder = txt_in(enc_in) + + joint = ops.concatenate([encoder, hidden], axis=1) + + rotary = ops.convert_to_tensor(rope_angles) + attn_mask = ops.convert_to_tensor(attn_mask_np) + target_mask_t = ops.convert_to_tensor(target_token_mask) + text_pos = np.flatnonzero(~image_pad_mask).astype(np.int32) + + text_scatter = np.zeros((joint_seq, text_seq_len), dtype=np.float32) + for j, pos in enumerate(text_pos): + text_scatter[pos, j] = 1.0 + text_scatter_t = ops.convert_to_tensor(text_scatter) + text_valid_on_joint = ops.einsum( + "bt,jt->bj", ops.cast(enc_mask_in, "float32"), text_scatter_t + ) + is_text = ops.convert_to_tensor((~image_pad_mask).astype(np.float32)) + key_ok = text_valid_on_joint * is_text[None, :] + (1.0 - is_text[None, :]) + attn_mask = attn_mask + (1.0 - key_ok)[:, None, None, :] * MASK_NEG + + if causal_condition: + t_all = ops.concatenate( + [timestep_in, ops.zeros((1,), dtype=timestep_in.dtype)], axis=0 + ) + temb = time_text_embed(t_all) + mod_mask = target_mask_t + else: + temb = time_text_embed(timestep_in) + mod_mask = None + + mod = modulation(ops.silu(temb)) + for block in blocks: + joint = block( + joint, + mod, + rotary_emb=rotary, + attention_mask=attn_mask, + target_token_mask=mod_mask, + ) + joint = norm_out([joint, temb, mod_mask]) + output_full = proj_out(joint) + output = output_full[:, -img_seq:, :] + super().__init__( + inputs={ + "sample": sample_in, + "timestep": timestep_in, + "encoder_hidden_states": enc_in, + "encoder_hidden_states_mask": enc_mask_in, + }, + outputs={"sample": output}, + name=name, + **keras_kwargs, + ) + + self.patch_size = patch_size + self.in_channels = in_channels + self.out_channels = out_channels or in_channels + self.num_layers = num_layers + self.attention_head_dim = attention_head_dim + self.num_attention_heads = num_attention_heads + self.context_in_dim = context_in_dim + self.mlp_ratio = mlp_ratio + self.axes_dims_rope = axes_dims_rope + self.eps = eps + self.causal_condition = causal_condition + self.sample_size = sample_size + self.text_seq_len = text_seq_len + self.inner_dim = inner_dim + self.img_seq = img_seq + self.joint_seq = joint_seq + self.time_text_embed = time_text_embed + self.txt_in = txt_in + self.img_in = img_in + self.modulation = modulation + self.transformer_blocks = blocks + self.norm_out = norm_out + self.proj_out = proj_out + + @classmethod + def kwargs_from_diffusers_config(cls, cfg): + return { + "patch_size": cfg.get("patch_size", 1), + "in_channels": cfg.get("in_channels", 64), + "out_channels": cfg.get("out_channels", 64), + "num_layers": cfg.get("num_layers", 32), + "attention_head_dim": cfg.get("attention_head_dim", 128), + "num_attention_heads": cfg.get("num_attention_heads", 32), + "context_in_dim": cfg.get("context_in_dim", 4096), + "mlp_ratio": cfg.get("mlp_ratio", 3), + "axes_dims_rope": tuple(cfg.get("axes_dims_rope", (16, 56, 56))), + "eps": cfg.get("eps", 1e-6), + "causal_condition": bool(cfg.get("causal_condition", True)), + } + + +@keras.saving.register_keras_serializable(package="zeromodels") +class QwenImage21TextEncoderModel(BaseModel): + """Qwen-Image-2.1 prompt encoder: Qwen3-VL text tower, no vision / LM head. + + Returns **pre-final-norm** hidden states (Diffusers hooks the last RMSNorm so + the transformer sees the decoder output before it). + """ + + HF_MODEL_TYPE = None + config_class = QwenImage21TextConfig + output_logits = False + + def __init__(self, name="text_encoder", **kwargs): + keras_kwargs = {k: kwargs.pop(k) for k in ("trainable", "dtype") if k in kwargs} + if "max_seq_len" in kwargs and not any( + k in kwargs for k in ("embed_dim", "num_layers", "vocab_size") + ): + pass + config = self.config_class.from_dict(kwargs) if kwargs else self.config_class() + max_seq_len = int(getattr(config, "max_seq_len", 1024)) + + language_model = Qwen3VLTextModel( + vocab_size=config.vocab_size, + embed_dim=config.embed_dim, + mlp_dim=config.mlp_dim, + num_layers=config.num_layers, + num_heads=config.num_heads, + num_kv_heads=config.num_kv_heads, + head_dim=config.head_dim, + norm_eps=config.norm_eps, + name="language_model", + ) + causal_mask = CausalMask(name="causal_mask") + + input_ids = layers.Input(shape=(None,), dtype="int32", name="input_ids") + attention_mask = layers.Input( + shape=(None,), dtype="int32", name="attention_mask" + ) + hidden = language_model.token_embedding(input_ids) + pos1 = ops.cumsum(ops.ones_like(input_ids), axis=-1) - 1 + pos = ops.stack([pos1, pos1, pos1], axis=0) + cos, sin = qwen3_text_cos_sin( + pos, config.head_dim, config.rope_theta, tuple(config.mrope_section) + ) + mask = causal_mask(input_ids, attention_mask) + for layer in language_model.decoder_layers: + hidden = layer(hidden, cos, sin, attention_mask=mask) + + super().__init__( + inputs={"input_ids": input_ids, "attention_mask": attention_mask}, + outputs={"last_hidden_state": hidden}, + name=name, + **keras_kwargs, + ) + self.language_model = language_model + self.causal_mask_layer = causal_mask + self.max_seq_len = max_seq_len + self.vocab_size = config.vocab_size + self.embed_dim = config.embed_dim + self.mlp_dim = config.mlp_dim + self.num_layers = config.num_layers + self.num_heads = config.num_heads + self.num_kv_heads = config.num_kv_heads + self.head_dim = config.head_dim + self.norm_eps = config.norm_eps + self.rope_theta = config.rope_theta + self.mrope_section = tuple(config.mrope_section) + self.tie_embeddings = config.tie_embeddings + + +@keras.saving.register_keras_serializable(package="zeromodels") +class QwenImage21Model(BaseModel): + """Qwen-Image-2.1 weights container: transformer + VAE + Qwen3-VL text tower.""" + + config_class = QwenImage21Config + HF_MODEL_TYPE = None + HUB_REPO_SIBLINGS = QWEN_IMAGE_21_HUB_SIBLINGS + + def __init__(self, name="QwenImage21Model", **kwargs): + keras_kwargs = {k: kwargs.pop(k) for k in ("trainable", "dtype") if k in kwargs} + config = self.config_class.from_dict(kwargs) + components = self.build_components(config) + inputs, outputs = self.build_graph(config, components) + super().__init__(inputs=inputs, outputs=outputs, name=name, **keras_kwargs) + for attr, component in components.items(): + setattr(self, attr, component) + + def build_components(self, config): + d, v = config.transformer_config, config.vae_config + transformer = QwenImage21Transformer2DModel( + patch_size=d.patch_size, + in_channels=d.in_channels, + out_channels=d.out_channels, + num_layers=d.num_layers, + attention_head_dim=d.attention_head_dim, + num_attention_heads=d.num_attention_heads, + context_in_dim=d.context_in_dim, + mlp_ratio=d.mlp_ratio, + axes_dims_rope=d.axes_dims_rope, + eps=d.eps, + causal_condition=d.causal_condition, + sample_size=d.sample_size, + text_seq_len=config.max_sequence_length, + ) + vae = AutoencoderKLQwenImage21( + base_dim=v.base_dim, + decoder_base_dim=v.decoder_base_dim, + z_dim=v.z_dim, + dim_mult=v.dim_mult, + num_res_blocks=v.num_res_blocks, + attn_scales=v.attn_scales, + temperal_downsample=v.temperal_downsample, + dropout=v.dropout, + input_channels=v.input_channels, + out_channels=v.out_channels, + is_residual=v.is_residual, + scale_factor_spatial=v.scale_factor_spatial, + latents_mean=v.latents_mean, + latents_std=v.latents_std, + sample_size=v.sample_size, + ) + text_kw = ( + config.text_config.constructor_kwargs() + if hasattr(config.text_config, "constructor_kwargs") + else dict(config.text_config) + ) + text_encoder = QwenImage21TextEncoderModel(**text_kw) + return {"transformer": transformer, "vae": vae, "text_encoder": text_encoder} + + def build_graph(self, config, components): + d, v, t = config.transformer_config, config.vae_config, config.text_config + transformer, vae, text_encoder = ( + components["transformer"], + components["vae"], + components["text_encoder"], + ) + text_seq = config.max_sequence_length + img_h = img_w = ( + v.sample_size + if not isinstance(v.sample_size, (tuple, list)) + else v.sample_size[0] + ) + lat_h = img_h // vae.vae_scale_factor + lat_w = img_w // vae.vae_scale_factor + + inputs = { + "sample": layers.Input( + shape=(transformer.img_seq, d.in_channels), name="sample" + ), + "timestep": layers.Input(shape=(), name="timestep"), + "encoder_hidden_states": layers.Input( + shape=(text_seq, d.context_in_dim), name="encoder_hidden_states" + ), + "encoder_hidden_states_mask": layers.Input( + shape=(text_seq,), dtype="int32", name="encoder_hidden_states_mask" + ), + "image": layers.Input(shape=(img_h, img_w, v.input_channels), name="image"), + "latent": layers.Input(shape=(lat_h, lat_w, v.z_dim), name="latent"), + "token_ids": layers.Input( + shape=(t.max_seq_len,), dtype="int32", name="token_ids" + ), + "padding_mask": layers.Input( + shape=(t.max_seq_len,), dtype="int32", name="padding_mask" + ), + } + noise_pred = transformer( + { + "sample": inputs["sample"], + "timestep": inputs["timestep"], + "encoder_hidden_states": inputs["encoder_hidden_states"], + "encoder_hidden_states_mask": inputs["encoder_hidden_states_mask"], + } + )["sample"] + vae_out = vae({"image": inputs["image"], "latent": inputs["latent"]}) + text_out = text_encoder( + {"input_ids": inputs["token_ids"], "attention_mask": inputs["padding_mask"]} + ) + return inputs, { + "noise_pred": noise_pred, + "moments": vae_out["moments"], + "image": vae_out["sample"], + "prompt_embeds": text_out["last_hidden_state"], + } + + def get_config(self): + config = super().get_config() + config.update(self.config.constructor_kwargs()) + return config + + def from_hf(self, *args, **kwargs): + raise NotImplementedError( + "On-the-fly hf: conversion is not supported for Qwen-Image-2.1; " + "use convert_qwen_image_21_diffusers_to_keras.py and " + "from_weights('zeromodels/qwen-image-2.1')." + ) + + +@keras.saving.register_keras_serializable(package="zeromodels") +class QwenImage21TextToImage(QwenImage21Model, BaseDiffusion): + """Text-to-image Qwen-Image-2.1 (Diffusers ``QwenImage21Pipeline``). + + :: + + model = QwenImage21TextToImage.from_weights("zeromodels/qwen-image-2.1") + tok = QwenImage21Tokenizer.from_weights("zeromodels/qwen-image-2.1") + image = model.generate(**tok("a cat"), height=1024, width=1024) + + Defaults match the reference: 40 flow-match steps, ``guidance_scale=1.0`` + (no CFG). Pass ``guidance_scale>1`` with a negative prompt for true CFG. + """ + + config_class = QwenImage21Config + HUB_REPO_SIBLINGS = QWEN_IMAGE_21_HUB_SIBLINGS + generate_args = {"num_inference_steps": 40, "guidance_scale": 1.0} + DEFAULT_GUIDANCE_SCALE = 1.0 + + def __init__(self, scheduler=None, name="QwenImage21TextToImage", **kwargs): + scheduler_config = kwargs.get("scheduler_config") + super().__init__(name=name, **kwargs) + if scheduler is None: + scheduler = ( + get_scheduler(scheduler_config) + if scheduler_config + else self.default_scheduler() + ) + self.scheduler = scheduler + + def default_scheduler(self): + return FlowMatchEulerDiscreteScheduler( + shift=1.0, + use_dynamic_shifting=True, + base_shift=0.5, + max_shift=0.9, + base_image_seq_len=256, + max_image_seq_len=8192, + shift_terminal=0.02, + time_shift_type="exponential", + ) + + @property + def vae_scale_factor(self): + return self.vae.vae_scale_factor + + @property + def latent_shape(self): + h, w = self.latent_side() + return (h * w, self.vae.z_dim) + + def latent_side(self, height=None, width=None): + scale = self.vae_scale_factor * 2 + if height is None or width is None: + side = self.config.default_sample_size * self.vae_scale_factor + height = width = side + h = 2 * (int(height) // scale) + w = 2 * (int(width) // scale) + return h, w + + def unconditional_ids(self, batch): + length = self.config.text_config.max_seq_len + row = [self.config.pad_token_id] * length + return ops.convert_to_tensor([row] * batch, dtype="int32") + + def encode_prompt(self, input_ids, attention_mask=None, **conditioning): + del conditioning + input_ids = ops.cast(ops.convert_to_tensor(input_ids), "int32") + if attention_mask is None: + attention_mask = ops.ones_like(input_ids) + else: + attention_mask = ops.cast(ops.convert_to_tensor(attention_mask), "int32") + + out = self.text_encoder( + {"input_ids": input_ids, "attention_mask": attention_mask} + ) + hidden = out["last_hidden_state"] + drop = int(self.config.prompt_template_encode_start_idx) + max_len = int(self.config.max_sequence_length) + + hidden_np = ops.convert_to_numpy(hidden) + mask_np = ops.convert_to_numpy(attention_mask) + batch = hidden_np.shape[0] + dim = hidden_np.shape[-1] + embeds = np.zeros((batch, max_len, dim), dtype=hidden_np.dtype) + out_mask = np.zeros((batch, max_len), dtype=np.int32) + for i in range(batch): + valid = hidden_np[i][mask_np[i].astype(bool)] + valid = valid[drop:][:max_len] + n = valid.shape[0] + embeds[i, :n] = valid + out_mask[i, :n] = 1 + return { + "encoder_hidden_states": ops.convert_to_tensor(embeds), + "encoder_hidden_states_mask": ops.convert_to_tensor(out_mask), + } + + def predict_noise(self, latents, timesteps, embeddings): + return self.transformer( + { + "sample": latents, + "timestep": timesteps, + "encoder_hidden_states": embeddings["encoder_hidden_states"], + "encoder_hidden_states_mask": embeddings["encoder_hidden_states_mask"], + } + )["sample"] + + def decode_latents(self, latents, height=None, width=None): + h, w = self.latent_side(height, width) + latents = unpack_latents(latents, h, w, self.vae.z_dim) + mean = ops.convert_to_tensor(self.vae.latents_mean, dtype="float32") + std = ops.convert_to_tensor(self.vae.latents_std, dtype="float32") + mean = ops.reshape(mean, (1, 1, 1, -1)) + std = ops.reshape(std, (1, 1, 1, -1)) + latents = latents * std + mean + image = self.vae.decode(latents) + if int(image.shape[-1]) == 4: + image = image[..., :3] + return image + + def prepare_latents(self, batch, seed=None, latents=None, dtype="float32"): + h, w = self.latent_side( + getattr(self, "_gen_height", None), getattr(self, "_gen_width", None) + ) + channels = self.vae.z_dim + if latents is None: + noise = keras.random.normal((batch, h, w, channels), seed=seed, dtype=dtype) + latents = pack_latents(noise, h, w) + else: + latents = ops.cast(ops.convert_to_tensor(latents), dtype) + return latents * self.scheduler.init_noise_sigma + + def denoise( + self, latents, embeddings, num_inference_steps, guidance_scale, timesteps=None + ): + do_cfg = guidance_scale > 1.0 + scheduler = self.scheduler + if timesteps is None: + scheduler.set_timesteps(num_inference_steps) + timesteps = scheduler.timesteps + if do_cfg and isinstance(embeddings, (tuple, list)): + uncond_emb, cond_emb = embeddings + else: + cond_emb = embeddings + uncond_emb = None + + for t in timesteps: + t_batch = ops.full((ops.shape(latents)[0],), float(t) / 1000.0) + noise_pred = self.predict_noise(latents, t_batch, cond_emb) + if do_cfg and uncond_emb is not None: + neg_pred = self.predict_noise(latents, t_batch, uncond_emb) + noise_pred = neg_pred + guidance_scale * (noise_pred - neg_pred) + latents = scheduler.step(noise_pred, t, latents) + return latents + + def generate( + self, + input_ids, + attention_mask=None, + negative_input_ids=None, + negative_attention_mask=None, + num_inference_steps=None, + guidance_scale=None, + seed=None, + latents=None, + height=None, + width=None, + output_type="image", + **conditioning, + ): + num_inference_steps, guidance_scale, seed = self.resolve_generation_args( + num_inference_steps, guidance_scale, seed + ) + self._gen_height = height + self._gen_width = width + input_ids = ops.cast(ops.convert_to_tensor(input_ids), "int32") + batch = int(input_ids.shape[0]) + + with inference_scope(): + embeddings = self.encode_prompt(input_ids, attention_mask, **conditioning) + if guidance_scale > 1.0 and negative_input_ids is None: + do_cfg = False + uncond = None + elif negative_input_ids is not None and guidance_scale > 1.0: + uncond = self.encode_prompt(negative_input_ids, negative_attention_mask) + do_cfg = True + else: + uncond = None + do_cfg = False + + h, w = self.latent_side(height, width) + image_seq_len = h * w + sched_cfg = getattr(self.scheduler, "config_dict", None) or {} + base_seq_len = sched_cfg.get("base_image_seq_len", 256) + max_seq_len = sched_cfg.get("max_image_seq_len", 4096) + base_shift = sched_cfg.get("base_shift", 0.5) + max_shift = sched_cfg.get("max_shift", 0.9) + m = (max_shift - base_shift) / (max_seq_len - base_seq_len) + mu = image_seq_len * m + (base_shift - m * base_seq_len) + sigmas = np.linspace(1.0, 1.0 / num_inference_steps, num_inference_steps) + if hasattr(self.scheduler, "set_timesteps"): + try: + self.scheduler.set_timesteps( + num_inference_steps, sigmas=sigmas, mu=mu + ) + except TypeError: + self.scheduler.set_timesteps(num_inference_steps) + timesteps = self.scheduler.timesteps + + latents = self.prepare_latents(batch, seed=seed, latents=latents) + emb = (uncond, embeddings) if do_cfg else embeddings + latents = self.denoise( + latents, + emb, + num_inference_steps, + guidance_scale if do_cfg else 1.0, + timesteps=timesteps, + ) + if output_type == "latent": + return latents + image = self.decode_latents(latents, height=height, width=width) + return self.postprocess_image(image) diff --git a/zeromodels/models/qwen_image_21/qwen_image_21_tokenizer.py b/zeromodels/models/qwen_image_21/qwen_image_21_tokenizer.py new file mode 100644 index 00000000..f0975ba3 --- /dev/null +++ b/zeromodels/models/qwen_image_21/qwen_image_21_tokenizer.py @@ -0,0 +1,57 @@ +import keras + +from zeromodels.models.qwen3.qwen3_tokenizer import Qwen3Tokenizer + +SYS_PROMPT = "Comprehend and analyze the provided prompt." +PROMPT_TEMPLATE = ( + f"<|im_start|>system\n{SYS_PROMPT}<|im_end|>\n" + f"<|im_start|>user\n{{}}<|im_end|>\n" + f"<|im_start|>assistant\n" +) +PROMPT_TEMPLATE_START_IDX = 14 + + +@keras.saving.register_keras_serializable(package="zeromodels") +class QwenImage21Tokenizer(Qwen3Tokenizer): + """Tokenizer for Qwen-Image-2.1 text-to-image. + + Wraps each prompt in the Diffusers ChatML template (system prompt + ``Comprehend and analyze the provided prompt.``), then BPE-encodes. + Returns ``input_ids`` / ``attention_mask`` for + :meth:`QwenImage21TextToImage.generate`. + + Load:: + + tok = QwenImage21Tokenizer.from_weights("zeromodels/qwen-image-2.1") + model.generate(**tok("a photo of a cat")) + """ + + prompt_template = PROMPT_TEMPLATE + prompt_template_start_idx = PROMPT_TEMPLATE_START_IDX + + def __init__( + self, + hf_id=None, + tokenizer_file=None, + max_seq_len=1024, + **kwargs, + ): + self.max_seq_len = max_seq_len + self.tokenizer_max_length = max_seq_len + super().__init__(hf_id=hf_id, tokenizer_file=tokenizer_file, **kwargs) + + def format_prompt(self, text): + return self.prompt_template.format(text if text else " ") + + def call(self, inputs): + texts = self.normalize_texts(inputs) + templated = [self.format_prompt(t) for t in texts] + max_length = self.tokenizer_max_length + self.prompt_template_start_idx + encoded = [self.encode(t)[:max_length] for t in templated] + input_ids, attention_mask = self.pad_batch(encoded) + return {"input_ids": input_ids, "attention_mask": attention_mask} + + def get_config(self): + config = super().get_config() + config.update({"max_seq_len": self.max_seq_len}) + return config diff --git a/zeromodels/models/qwen_image_21/qwen_image_21_vae.py b/zeromodels/models/qwen_image_21/qwen_image_21_vae.py new file mode 100644 index 00000000..7bbecaad --- /dev/null +++ b/zeromodels/models/qwen_image_21/qwen_image_21_vae.py @@ -0,0 +1,892 @@ +from __future__ import annotations + +import keras +from keras import layers, ops + +from zeromodels.base.base_attention import fused_attention +from zeromodels.models.qwen_image.qwen_image_vae import QwenImageRMSNorm +from zeromodels.models.stable_diffusion.stable_diffusion_layers import safe_name + + +@keras.saving.register_keras_serializable(package="zeromodels") +class QwenImage21CausalConv(layers.Layer): + """Spatial Conv2d with NDHWC T-squeeze (Diffusers ``QwenImage21CausalConv3d``).""" + + def __init__( + self, + out_channels, + kernel_size=3, + stride=1, + padding=0, + module_path=None, + **kwargs, + ): + if module_path is not None: + kwargs.setdefault("name", safe_name(module_path)) + super().__init__(**kwargs) + self.out_channels = int(out_channels) + if isinstance(kernel_size, int): + self.kernel_size = (kernel_size, kernel_size) + else: + ks = tuple(kernel_size) + self.kernel_size = (ks[-2], ks[-1]) if len(ks) == 3 else (ks[0], ks[1]) + if isinstance(stride, int): + self.stride = (stride, stride) + else: + st = tuple(stride) + self.stride = (st[-2], st[-1]) if len(st) == 3 else (st[0], st[1]) + if isinstance(padding, int): + self.pad_h = self.pad_w = padding + else: + pad = tuple(padding) + self.pad_h = pad[0] if len(pad) >= 1 else 0 + self.pad_w = pad[1] if len(pad) >= 2 else self.pad_h + self.module_path = module_path + self.conv = layers.Conv2D( + self.out_channels, + self.kernel_size, + strides=self.stride, + padding="valid", + data_format="channels_last", + name=safe_name(module_path) if module_path else "conv", + ) + + def build(self, input_shape): + # (B, T, H, W, C) → build Conv2D on (B, H, W, C) + b, _, h, w, c = input_shape + h_p = None if h is None else h + 2 * self.pad_h + w_p = None if w is None else w + 2 * self.pad_w + self.conv.build((b, h_p, w_p, c)) + self.built = True + + def call(self, x): + # x: (B, T, H, W, C) with T=1 for T2I + b = ops.shape(x)[0] + t = ops.shape(x)[1] + h = ops.shape(x)[2] + w = ops.shape(x)[3] + c = ops.shape(x)[4] + x2 = ops.reshape(x, (b * t, h, w, c)) + if self.pad_h or self.pad_w: + x2 = ops.pad( + x2, + ( + (0, 0), + (self.pad_h, self.pad_h), + (self.pad_w, self.pad_w), + (0, 0), + ), + ) + x2 = self.conv(x2) + out_h = ops.shape(x2)[1] + out_w = ops.shape(x2)[2] + return ops.reshape(x2, (b, t, out_h, out_w, self.out_channels)) + + def compute_output_shape(self, input_shape): + b, t, h, w, _ = input_shape + if h is None or w is None: + return (b, t, None, None, self.out_channels) + h_p = h + 2 * self.pad_h + w_p = w + 2 * self.pad_w + out_h = (h_p - self.kernel_size[0]) // self.stride[0] + 1 + out_w = (w_p - self.kernel_size[1]) // self.stride[1] + 1 + return (b, t, out_h, out_w, self.out_channels) + + def get_config(self): + config = super().get_config() + config.update( + { + "out_channels": self.out_channels, + "kernel_size": self.kernel_size, + "stride": self.stride, + "padding": (self.pad_h, self.pad_w), + "module_path": self.module_path, + } + ) + return config + + +@keras.saving.register_keras_serializable(package="zeromodels") +class QwenImage21AvgDown3D(layers.Layer): + """Average downsample shortcut (Diffusers ``QwenImage21AvgDown3D``), NDHWC.""" + + def __init__(self, in_channels, out_channels, factor_t, factor_s=1, **kwargs): + super().__init__(**kwargs) + self.in_channels = int(in_channels) + self.out_channels = int(out_channels) + self.factor_t = int(factor_t) + self.factor_s = int(factor_s) + self.factor = self.factor_t * self.factor_s * self.factor_s + if self.in_channels * self.factor % self.out_channels != 0: + raise ValueError( + f"in_channels ({in_channels}) * factor ({self.factor}) must be " + f"divisible by out_channels ({out_channels})." + ) + self.group_size = self.in_channels * self.factor // self.out_channels + + def call(self, x): + ft, fs = self.factor_t, self.factor_s + pad_t = (ft - (ops.shape(x)[1] % ft)) % ft + x = ops.pad(x, ((0, 0), (pad_t, 0), (0, 0), (0, 0), (0, 0))) + b = ops.shape(x)[0] + t = ops.shape(x)[1] + h = ops.shape(x)[2] + w = ops.shape(x)[3] + c = self.in_channels + x = ops.reshape(x, (b, t // ft, ft, h // fs, fs, w // fs, fs, c)) + x = ops.transpose(x, (0, 1, 3, 5, 7, 2, 4, 6)) + x = ops.reshape(x, (b, t // ft, h // fs, w // fs, c * self.factor)) + x = ops.reshape( + x, + (b, t // ft, h // fs, w // fs, self.out_channels, self.group_size), + ) + return ops.mean(x, axis=-1) + + def compute_output_shape(self, input_shape): + b, t, h, w, _ = input_shape + ft, fs = self.factor_t, self.factor_s + + def down(size, factor): + if size is None: + return None + pad = (factor - size % factor) % factor + return (size + pad) // factor + + return ( + b, + down(t, ft) if t is not None else None, + down(h, fs), + down(w, fs), + self.out_channels, + ) + + def get_config(self): + config = super().get_config() + config.update( + { + "in_channels": self.in_channels, + "out_channels": self.out_channels, + "factor_t": self.factor_t, + "factor_s": self.factor_s, + } + ) + return config + + +@keras.saving.register_keras_serializable(package="zeromodels") +class QwenImage21DupUp3D(layers.Layer): + """Duplicate upsample shortcut (Diffusers ``QwenImage21DupUp3D``), NDHWC.""" + + def __init__(self, in_channels, out_channels, factor_t, factor_s=1, **kwargs): + super().__init__(**kwargs) + self.in_channels = int(in_channels) + self.out_channels = int(out_channels) + self.factor_t = int(factor_t) + self.factor_s = int(factor_s) + self.factor = self.factor_t * self.factor_s * self.factor_s + assert self.out_channels * self.factor % self.in_channels == 0 + self.repeats = self.out_channels * self.factor // self.in_channels + + def call(self, x, first_chunk=False): + x = ops.repeat(x, self.repeats, axis=-1) + b = ops.shape(x)[0] + t = ops.shape(x)[1] + h = ops.shape(x)[2] + w = ops.shape(x)[3] + ft, fs = self.factor_t, self.factor_s + x = ops.reshape(x, (b, t, h, w, self.out_channels, ft, fs, fs)) + x = ops.transpose(x, (0, 1, 5, 2, 6, 3, 7, 4)) + x = ops.reshape(x, (b, t * ft, h * fs, w * fs, self.out_channels)) + if first_chunk and ft > 1: + x = x[:, ft - 1 :, :, :, :] + return x + + def compute_output_shape(self, input_shape): + b, t, h, w, _ = input_shape + ft, fs = self.factor_t, self.factor_s + return ( + b, + None if t is None else t * ft, + None if h is None else h * fs, + None if w is None else w * fs, + self.out_channels, + ) + + def get_config(self): + config = super().get_config() + config.update( + { + "in_channels": self.in_channels, + "out_channels": self.out_channels, + "factor_t": self.factor_t, + "factor_s": self.factor_s, + } + ) + return config + + +@keras.saving.register_keras_serializable(package="zeromodels") +class QwenImage21Resample(layers.Layer): + """Spatial resample (Diffusers ``QwenImage21Resample``), T=1 / no temporal path.""" + + def __init__(self, dim, mode, module_path, upsample_out_dim=None, **kwargs): + kwargs.setdefault("name", safe_name(module_path)) + super().__init__(**kwargs) + self.dim = int(dim) + self.mode = mode + self.module_path = module_path + self.upsample_out_dim = ( + int(upsample_out_dim) if upsample_out_dim is not None else self.dim // 2 + ) + self._downsample = mode in ("downsample2d", "downsample3d") + self._upsample = mode in ("upsample2d", "upsample3d") + self.spatial_conv = None + if self._upsample: + self.spatial_conv = layers.Conv2D( + self.upsample_out_dim, + 3, + padding="same", + data_format="channels_last", + name=safe_name(f"{module_path}.resample.1"), + ) + elif self._downsample: + self.spatial_conv = layers.Conv2D( + self.dim, + 3, + strides=2, + padding="valid", + data_format="channels_last", + name=safe_name(f"{module_path}.resample.1"), + ) + elif mode != "none": + raise ValueError(f"Unknown resample mode {mode!r}") + + def build(self, input_shape): + b, t, h, w, c = input_shape + if self.spatial_conv is not None: + if self._downsample: + self.spatial_conv.build( + (b, None if h is None else h + 1, None if w is None else w + 1, c) + ) + else: + self.spatial_conv.build( + (b, None if h is None else h * 2, None if w is None else w * 2, c) + ) + self.built = True + + def call(self, x): + b = ops.shape(x)[0] + t = ops.shape(x)[1] + h = ops.shape(x)[2] + w = ops.shape(x)[3] + c = ops.shape(x)[4] + x2 = ops.reshape(x, (b * t, h, w, c)) + if self._upsample: + x2 = ops.image.resize(x2, (h * 2, w * 2), interpolation="nearest") + x2 = self.spatial_conv(x2) + out_c = self.upsample_out_dim + out_h, out_w = h * 2, w * 2 + elif self._downsample: + x2 = ops.pad(x2, ((0, 0), (0, 1), (0, 1), (0, 0))) + x2 = self.spatial_conv(x2) + out_c = self.dim + out_h = (h + 1 - 3) // 2 + 1 + out_w = (w + 1 - 3) // 2 + 1 + else: + out_c, out_h, out_w = c, h, w + return ops.reshape(x2, (b, t, out_h, out_w, out_c)) + + def compute_output_shape(self, input_shape): + b, t, h, w, c = input_shape + if self._upsample: + return ( + b, + t, + None if h is None else h * 2, + None if w is None else w * 2, + self.upsample_out_dim, + ) + if self._downsample: + return ( + b, + t, + None if h is None else (h + 1 - 3) // 2 + 1, + None if w is None else (w + 1 - 3) // 2 + 1, + self.dim, + ) + return tuple(input_shape) + + def get_config(self): + config = super().get_config() + config.update( + { + "dim": self.dim, + "mode": self.mode, + "module_path": self.module_path, + "upsample_out_dim": self.upsample_out_dim, + } + ) + return config + + +@keras.saving.register_keras_serializable(package="zeromodels") +class QwenImage21ResidualBlock(layers.Layer): + """Residual block with spatial causal convs.""" + + def __init__(self, in_dim, out_dim, module_path, dropout=0.0, **kwargs): + kwargs.setdefault("name", safe_name(module_path)) + super().__init__(**kwargs) + self.in_dim = int(in_dim) + self.out_dim = int(out_dim) + self.module_path = module_path + self.dropout_rate = float(dropout) + self.norm1 = QwenImageRMSNorm( + in_dim, images=False, module_path=f"{module_path}.norm1" + ) + self.conv1 = QwenImage21CausalConv( + out_dim, 3, padding=1, module_path=f"{module_path}.conv1" + ) + self.norm2 = QwenImageRMSNorm( + out_dim, images=False, module_path=f"{module_path}.norm2" + ) + self.conv2 = QwenImage21CausalConv( + out_dim, 3, padding=1, module_path=f"{module_path}.conv2" + ) + self.conv_shortcut = None + if in_dim != out_dim: + self.conv_shortcut = QwenImage21CausalConv( + out_dim, 1, padding=0, module_path=f"{module_path}.conv_shortcut" + ) + self.dropout = layers.Dropout(dropout) if dropout > 0 else None + + def build(self, input_shape): + self.norm1.build(input_shape) + self.conv1.build(input_shape) + mid = list(input_shape) + mid[-1] = self.out_dim + mid = tuple(mid) + self.norm2.build(mid) + self.conv2.build(mid) + if self.conv_shortcut is not None: + self.conv_shortcut.build(input_shape) + self.built = True + + def call(self, x, training=None): + h = self.conv_shortcut(x) if self.conv_shortcut is not None else x + x = self.norm1(x) + x = ops.silu(x) + x = self.conv1(x) + x = self.norm2(x) + x = ops.silu(x) + if self.dropout is not None: + x = self.dropout(x, training=training) + x = self.conv2(x) + return x + h + + def compute_output_shape(self, input_shape): + shape = list(input_shape) + shape[-1] = self.out_dim + return tuple(shape) + + def get_config(self): + config = super().get_config() + config.update( + { + "in_dim": self.in_dim, + "out_dim": self.out_dim, + "module_path": self.module_path, + "dropout": self.dropout_rate, + } + ) + return config + + +@keras.saving.register_keras_serializable(package="zeromodels") +class QwenImage21AttentionBlock(layers.Layer): + """Spatial self-attention over H*W (Diffusers ``QwenImage21AttentionBlock``).""" + + def __init__(self, dim, module_path, **kwargs): + kwargs.setdefault("name", safe_name(module_path)) + super().__init__(**kwargs) + self.dim = int(dim) + self.module_path = module_path + self.norm = QwenImageRMSNorm( + dim, images=False, module_path=f"{module_path}.norm" + ) + self.to_qkv = layers.Dense( + dim * 3, use_bias=True, name=safe_name(f"{module_path}.to_qkv") + ) + self.proj = layers.Dense( + dim, use_bias=True, name=safe_name(f"{module_path}.proj") + ) + + def build(self, input_shape): + self.norm.build(input_shape) + self.to_qkv.build((*input_shape[:-1], self.dim)) + self.proj.build((*input_shape[:-1], self.dim)) + self.built = True + + def call(self, x): + residual = x + x = self.norm(x) + b = ops.shape(x)[0] + t = ops.shape(x)[1] + h = ops.shape(x)[2] + w = ops.shape(x)[3] + c = self.dim + x = ops.reshape(x, (b * t, h * w, c)) + qkv = self.to_qkv(x) + q, k, v = ops.split(qkv, 3, axis=-1) + q = ops.expand_dims(q, 1) + k = ops.expand_dims(k, 1) + v = ops.expand_dims(v, 1) + out = fused_attention(q, k, v, self.dim**-0.5) + out = ops.squeeze(out, 1) + out = self.proj(out) + out = ops.reshape(out, (b, t, h, w, c)) + return out + residual + + def compute_output_shape(self, input_shape): + return tuple(input_shape) + + def get_config(self): + config = super().get_config() + config.update({"dim": self.dim, "module_path": self.module_path}) + return config + + +@keras.saving.register_keras_serializable(package="zeromodels") +class QwenImage21MidBlock(layers.Layer): + def __init__(self, dim, module_path, dropout=0.0, num_layers=1, **kwargs): + kwargs.setdefault("name", safe_name(module_path)) + super().__init__(**kwargs) + self.dim = int(dim) + self.module_path = module_path + self.resnets = [ + QwenImage21ResidualBlock( + dim, dim, f"{module_path}.resnets.0", dropout=dropout + ) + ] + self.attentions = [] + for i in range(num_layers): + self.attentions.append( + QwenImage21AttentionBlock(dim, f"{module_path}.attentions.{i}") + ) + self.resnets.append( + QwenImage21ResidualBlock( + dim, dim, f"{module_path}.resnets.{i + 1}", dropout=dropout + ) + ) + + def build(self, input_shape): + for layer in self.resnets: + layer.build(input_shape) + for layer in self.attentions: + layer.build(input_shape) + self.built = True + + def call(self, x, training=None): + x = self.resnets[0](x, training=training) + for attn, resnet in zip(self.attentions, self.resnets[1:]): + x = attn(x) + x = resnet(x, training=training) + return x + + def compute_output_shape(self, input_shape): + return tuple(input_shape) + + def get_config(self): + config = super().get_config() + config.update({"dim": self.dim, "module_path": self.module_path}) + return config + + +@keras.saving.register_keras_serializable(package="zeromodels") +class QwenImage21ResidualDownBlock(layers.Layer): + def __init__( + self, + in_dim, + out_dim, + module_path, + dropout=0.0, + num_res_blocks=2, + temperal_downsample=False, + down_flag=False, + **kwargs, + ): + kwargs.setdefault("name", safe_name(module_path)) + super().__init__(**kwargs) + self.in_dim = int(in_dim) + self.out_dim = int(out_dim) + self.module_path = module_path + self.num_res_blocks = int(num_res_blocks) + self.temperal_downsample = bool(temperal_downsample) + self.down_flag = bool(down_flag) + self.avg_shortcut = QwenImage21AvgDown3D( + in_dim, + out_dim, + factor_t=2 if temperal_downsample else 1, + factor_s=2 if down_flag else 1, + name=safe_name(f"{module_path}.avg_shortcut"), + ) + self.resnets = [] + current = in_dim + for i in range(num_res_blocks): + self.resnets.append( + QwenImage21ResidualBlock( + current, out_dim, f"{module_path}.resnets.{i}", dropout=dropout + ) + ) + current = out_dim + self.downsampler = None + if down_flag: + mode = "downsample3d" if temperal_downsample else "downsample2d" + self.downsampler = QwenImage21Resample( + out_dim, mode, f"{module_path}.downsampler" + ) + + def build(self, input_shape): + self.avg_shortcut.build(input_shape) + shape = input_shape + for layer in self.resnets: + layer.build(shape) + shape = layer.compute_output_shape(shape) + if self.downsampler is not None: + self.downsampler.build(shape) + self.built = True + + def call(self, x, training=None): + shortcut = self.avg_shortcut(x) + for resnet in self.resnets: + x = resnet(x, training=training) + if self.downsampler is not None: + x = self.downsampler(x) + return x + shortcut + + def compute_output_shape(self, input_shape): + shape = input_shape + for layer in self.resnets: + shape = layer.compute_output_shape(shape) + if self.downsampler is not None: + shape = self.downsampler.compute_output_shape(shape) + return shape + + def get_config(self): + config = super().get_config() + config.update( + { + "in_dim": self.in_dim, + "out_dim": self.out_dim, + "module_path": self.module_path, + "num_res_blocks": self.num_res_blocks, + "temperal_downsample": self.temperal_downsample, + "down_flag": self.down_flag, + } + ) + return config + + +@keras.saving.register_keras_serializable(package="zeromodels") +class QwenImage21ResidualUpBlock(layers.Layer): + def __init__( + self, + in_dim, + out_dim, + module_path, + dropout=0.0, + num_res_blocks=2, + temperal_upsample=False, + up_flag=False, + **kwargs, + ): + kwargs.setdefault("name", safe_name(module_path)) + super().__init__(**kwargs) + self.in_dim = int(in_dim) + self.out_dim = int(out_dim) + self.module_path = module_path + self.num_res_blocks = int(num_res_blocks) + self.temperal_upsample = bool(temperal_upsample) + self.up_flag = bool(up_flag) + self.avg_shortcut = None + if up_flag: + self.avg_shortcut = QwenImage21DupUp3D( + in_dim, + out_dim, + factor_t=2 if temperal_upsample else 1, + factor_s=2, + name=safe_name(f"{module_path}.avg_shortcut"), + ) + self.resnets = [] + current = in_dim + for i in range(num_res_blocks + 1): + self.resnets.append( + QwenImage21ResidualBlock( + current, out_dim, f"{module_path}.resnets.{i}", dropout=dropout + ) + ) + current = out_dim + self.upsampler = None + if up_flag: + mode = "upsample3d" if temperal_upsample else "upsample2d" + self.upsampler = QwenImage21Resample( + out_dim, mode, f"{module_path}.upsampler", upsample_out_dim=out_dim + ) + + def build(self, input_shape): + if self.avg_shortcut is not None: + self.avg_shortcut.build(input_shape) + shape = input_shape + for layer in self.resnets: + layer.build(shape) + shape = layer.compute_output_shape(shape) + if self.upsampler is not None: + self.upsampler.build(shape) + self.built = True + + def call(self, x, training=None, first_chunk=False): + shortcut = None + if self.avg_shortcut is not None: + shortcut = self.avg_shortcut(x, first_chunk=first_chunk) + for resnet in self.resnets: + x = resnet(x, training=training) + if self.upsampler is not None: + x = self.upsampler(x) + if shortcut is not None: + x = x + shortcut + return x + + def compute_output_shape(self, input_shape): + shape = input_shape + for layer in self.resnets: + shape = layer.compute_output_shape(shape) + if self.upsampler is not None: + shape = self.upsampler.compute_output_shape(shape) + return shape + + def get_config(self): + config = super().get_config() + config.update( + { + "in_dim": self.in_dim, + "out_dim": self.out_dim, + "module_path": self.module_path, + "num_res_blocks": self.num_res_blocks, + "temperal_upsample": self.temperal_upsample, + "up_flag": self.up_flag, + } + ) + return config + + +@keras.saving.register_keras_serializable(package="zeromodels") +class QwenImage21Encoder3d(layers.Layer): + def __init__( + self, + dim=96, + z_dim=128, + dim_mult=(1, 2, 4, 8, 8), + num_res_blocks=2, + attn_scales=(), + temperal_downsample=(False, True, True, True), + dropout=0.0, + input_channels=4, + is_residual=True, + module_path="encoder", + **kwargs, + ): + kwargs.setdefault("name", safe_name(module_path)) + super().__init__(**kwargs) + self.dim = int(dim) + self.z_dim = int(z_dim) + self.dim_mult = tuple(dim_mult) + self.num_res_blocks = int(num_res_blocks) + self.attn_scales = tuple(attn_scales) + self.temperal_downsample = tuple(temperal_downsample) + self.dropout_rate = float(dropout) + self.input_channels = int(input_channels) + self.is_residual = bool(is_residual) + self.module_path = module_path + + dims = [dim * u for u in [1] + list(dim_mult)] + self.conv_in = QwenImage21CausalConv( + dims[0], 3, padding=1, module_path=f"{module_path}.conv_in" + ) + self.down_blocks = [] + for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])): + self.down_blocks.append( + QwenImage21ResidualDownBlock( + in_dim, + out_dim, + f"{module_path}.down_blocks.{i}", + dropout=dropout, + num_res_blocks=num_res_blocks, + temperal_downsample=( + temperal_downsample[i] if i != len(dim_mult) - 1 else False + ), + down_flag=i != len(dim_mult) - 1, + ) + ) + self.mid_block = QwenImage21MidBlock( + dims[-1], f"{module_path}.mid_block", dropout=dropout, num_layers=1 + ) + self.norm_out = QwenImageRMSNorm( + dims[-1], images=False, module_path=f"{module_path}.norm_out" + ) + self.conv_out = QwenImage21CausalConv( + z_dim, 3, padding=1, module_path=f"{module_path}.conv_out" + ) + + def build(self, input_shape): + self.conv_in.build(input_shape) + shape = self.conv_in.compute_output_shape(input_shape) + for block in self.down_blocks: + block.build(shape) + shape = block.compute_output_shape(shape) + self.mid_block.build(shape) + self.norm_out.build(shape) + self.conv_out.build(shape) + self.built = True + + def call(self, x, training=None): + x = self.conv_in(x) + for block in self.down_blocks: + x = block(x, training=training) + x = self.mid_block(x, training=training) + x = self.norm_out(x) + x = ops.silu(x) + return self.conv_out(x) + + def compute_output_shape(self, input_shape): + shape = self.conv_in.compute_output_shape(input_shape) + for block in self.down_blocks: + shape = block.compute_output_shape(shape) + shape = list(shape) + shape[-1] = self.z_dim + return tuple(shape) + + def get_config(self): + config = super().get_config() + config.update( + { + "dim": self.dim, + "z_dim": self.z_dim, + "dim_mult": self.dim_mult, + "num_res_blocks": self.num_res_blocks, + "attn_scales": self.attn_scales, + "temperal_downsample": self.temperal_downsample, + "dropout": self.dropout_rate, + "input_channels": self.input_channels, + "is_residual": self.is_residual, + "module_path": self.module_path, + } + ) + return config + + +@keras.saving.register_keras_serializable(package="zeromodels") +class QwenImage21Decoder3d(layers.Layer): + def __init__( + self, + dim=144, + z_dim=64, + dim_mult=(1, 2, 4, 8, 8), + num_res_blocks=2, + attn_scales=(), + temperal_upsample=(True, True, True, False), + dropout=0.0, + out_channels=4, + is_residual=True, + module_path="decoder", + **kwargs, + ): + kwargs.setdefault("name", safe_name(module_path)) + super().__init__(**kwargs) + self.dim = int(dim) + self.z_dim = int(z_dim) + self.dim_mult = tuple(dim_mult) + self.num_res_blocks = int(num_res_blocks) + self.attn_scales = tuple(attn_scales) + self.temperal_upsample = tuple(temperal_upsample) + self.dropout_rate = float(dropout) + self.out_channels = int(out_channels) + self.is_residual = bool(is_residual) + self.module_path = module_path + + dims = [dim * u for u in [self.dim_mult[-1]] + list(self.dim_mult[::-1])] + self.conv_in = QwenImage21CausalConv( + dims[0], 3, padding=1, module_path=f"{module_path}.conv_in" + ) + self.mid_block = QwenImage21MidBlock( + dims[0], f"{module_path}.mid_block", dropout=dropout, num_layers=1 + ) + self.up_blocks = [] + for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])): + up_flag = i != len(self.dim_mult) - 1 + self.up_blocks.append( + QwenImage21ResidualUpBlock( + in_dim, + out_dim, + f"{module_path}.up_blocks.{i}", + dropout=dropout, + num_res_blocks=num_res_blocks, + temperal_upsample=self.temperal_upsample[i] if up_flag else False, + up_flag=up_flag, + ) + ) + self.norm_out = QwenImageRMSNorm( + dims[-1], images=False, module_path=f"{module_path}.norm_out" + ) + self.conv_out = QwenImage21CausalConv( + out_channels, 3, padding=1, module_path=f"{module_path}.conv_out" + ) + + def build(self, input_shape): + self.conv_in.build(input_shape) + shape = list(input_shape) + shape[-1] = self.dim * self.dim_mult[-1] + shape = tuple(shape) + self.mid_block.build(shape) + for block in self.up_blocks: + block.build(shape) + shape = block.compute_output_shape(shape) + self.norm_out.build(shape) + self.conv_out.build(shape) + self.built = True + + def call(self, x, training=None): + x = self.conv_in(x) + x = self.mid_block(x, training=training) + for block in self.up_blocks: + x = block(x, training=training, first_chunk=True) + x = self.norm_out(x) + x = ops.silu(x) + return self.conv_out(x) + + def compute_output_shape(self, input_shape): + b, t, h, w, _ = input_shape + scale = 2 ** (len(self.dim_mult) - 1) + return ( + b, + t, + None if h is None else h * scale, + None if w is None else w * scale, + self.out_channels, + ) + + def get_config(self): + config = super().get_config() + config.update( + { + "dim": self.dim, + "z_dim": self.z_dim, + "dim_mult": self.dim_mult, + "num_res_blocks": self.num_res_blocks, + "attn_scales": self.attn_scales, + "temperal_upsample": self.temperal_upsample, + "dropout": self.dropout_rate, + "out_channels": self.out_channels, + "is_residual": self.is_residual, + "module_path": self.module_path, + } + ) + return config