From dd292726f091289ecbbf26ded847c96f7c36cd4f Mon Sep 17 00:00:00 2001 From: IMvision12 Date: Tue, 22 Sep 2026 23:59:00 -0700 Subject: [PATCH 1/9] Add initial qwen image 2.1 --- README.md | 3 +- docs/getting_started.md | 3 +- docs/index.md | 3 +- docs/loading_weights.md | 2 +- docs/models.md | 1 + docs/qwen_image_21.md | 62 ++ tests/base/model_test_registry.py | 57 ++ tests/fixtures/cross_backend_parity.json | 592 ++++++++++++ tests/fixtures/dummy_inputs.py | 28 + tests/integration/test_data_formats.py | 3 + website/mkdocs.yml | 1 + zeromodels/auto/auto_mapping_names.py | 7 + zeromodels/models/__init__.py | 1 + zeromodels/models/qwen_image_21/__init__.py | 27 + ...onvert_qwen_image_21_diffusers_to_keras.py | 389 ++++++++ .../qwen_image_21/qwen_image_21_config.py | 246 +++++ .../qwen_image_21/qwen_image_21_layers.py | 532 +++++++++++ .../qwen_image_21/qwen_image_21_model.py | 880 ++++++++++++++++++ .../qwen_image_21/qwen_image_21_tokenizer.py | 60 ++ .../models/qwen_image_21/qwen_image_21_vae.py | 874 +++++++++++++++++ 20 files changed, 3767 insertions(+), 4 deletions(-) create mode 100644 docs/qwen_image_21.md create mode 100644 zeromodels/models/qwen_image_21/__init__.py create mode 100644 zeromodels/models/qwen_image_21/convert_qwen_image_21_diffusers_to_keras.py create mode 100644 zeromodels/models/qwen_image_21/qwen_image_21_config.py create mode 100644 zeromodels/models/qwen_image_21/qwen_image_21_layers.py create mode 100644 zeromodels/models/qwen_image_21/qwen_image_21_model.py create mode 100644 zeromodels/models/qwen_image_21/qwen_image_21_tokenizer.py create mode 100644 zeromodels/models/qwen_image_21/qwen_image_21_vae.py diff --git a/README.md b/README.md index d3f0ecc4..4229123e 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) | `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_21.md b/docs/qwen_image_21.md new file mode 100644 index 00000000..9df921c5 --- /dev/null +++ b/docs/qwen_image_21.md @@ -0,0 +1,62 @@ +# 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), and the Qwen3-VL text tower. The whole model is **one container**, +`QwenImage21Model`. `QwenImage21TextToImage` adds `generate`. + +Key facts of the port: + +- **Unpatched latents**: the denoiser sees `(B, H · W, 64)` tokens (VAE scale 16; + no 2×2 packing). +- **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. + +Links: + +- Reference: [diffusers `QwenImage21Pipeline`](https://huggingface.co/docs/diffusers/api/pipelines/qwenimage) +- Source: [`Qwen/Qwen-Image-2.1`](https://huggingface.co/Qwen/Qwen-Image-2.1) +- See also [qwen_image.md](qwen_image.md), [qwen3_vl.md](qwen3_vl.md) + +## Variants + +| 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 + +### `QwenImage21TextToImage` + +```python +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") +image = model.generate(**tokenizer("a capybara in a wizard hat"), height=1024, width=1024) +``` + +Offline conversion (no on-the-fly `hf:` for diffusion):: + +```bash +python -m zeromodels.models.qwen_image_21.convert_qwen_image_21_diffusers_to_keras +``` 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..a391575d 100644 --- a/tests/fixtures/dummy_inputs.py +++ b/tests/fixtures/dummy_inputs.py @@ -264,6 +264,34 @@ 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/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..f4833d11 --- /dev/null +++ b/zeromodels/models/qwen_image_21/convert_qwen_image_21_diffusers_to_keras.py @@ -0,0 +1,389 @@ +"""Offline Diffusers ``Qwen/Qwen-Image-2.1`` → ZeroModels Keras weight conversion. + +Converts the transformer, VAE, and Qwen3-VL text encoder into a hosted +``zeromodels/qwen-image-2.1`` layout (``zm_config.json`` + sharded weights). + +Usage:: + + python -m zeromodels.models.qwen_image_21.convert_qwen_image_21_diffusers_to_keras + +Env: + ZM_OUT_DIR output directory (default ``./qwen_image_21_weights``) + HF_TOKEN optional Hub token + ZM_DTYPE ``float16`` / ``bfloat16`` / ``float32`` (default ``bfloat16``) +""" + +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", + "/": ".", + # text encoder (Qwen3-VL language tower) + "token_embedding.embeddings": "model.embed_tokens.weight", + "language_model.final_norm.weight": "model.norm.weight", + "language_model.": "model.", + "decoder_layer_": "layers.", + "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.query_norm": "self_attn.q_norm", + "attention.key_norm": "self_attn.k_norm", + "attention_norm": "input_layernorm", + "mlp_norm": "post_attention_layernorm", + "mlp.gate": "mlp.gate_proj", + "mlp.up": "mlp.up_proj", + "mlp.down": "mlp.down_proj", + "gamma": "weight", + "beta": "bias", + "kernel": "weight", +} + + +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_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": 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) + + print(f"[1/4] Building QwenImage21Model (dtype={dtype})…", flush=True) + with build_dtype_scope(dtype), zeros_init(): + model = QwenImage21Model(**flat) + + vae_mapping = {k: v for k, v in WEIGHT_NAME_MAPPING.items() if "/gamma" not in k} + 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) + for keras_weight, _ in tqdm( + trainable + non_trainable, + desc=f"Transferring {subfolder} weights to Keras", + ): + key = "/".join(keras_weight.path.split("/")[-2:]) + # Prefer full path remapping from the component root. + full = keras_weight.path + # Strip leading component name + parts = full.split("/") + if parts and parts[0] in ( + component.name, + "transformer", + "vae", + "QwenImage21Transformer2DModel", + "AutoencoderKLQwenImage21", + ): + rel = "/".join(parts[1:]) + else: + rel = "/".join(parts) + key = rel + for old, new in mapping.items(): + key = key.replace(old, new) + if key not in state: + # Fall back to leaf-pair mapping used by Qwen-Image 1.0. + key = "/".join(keras_weight.path.split("/")[-2:]) + for old, new in mapping.items(): + key = key.replace(old, new) + if key not in state: + raise WeightMappingError(keras_weight.path, key) + 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 ( + 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: + # Text-only tower: keep language_model / embed_tokens; drop vision + lm_head + # + final norm (Diffusers reads pre-norm hidden states). + if key.startswith("model.visual.") or key.startswith("lm_head."): + continue + if key in ("model.norm.weight",) or key.endswith(".norm.weight") and key.count(".") <= 2: + if key == "model.norm.weight": + continue + if key.startswith("model.language_model."): + hf_keys["model." + key[len("model.language_model.") :]] = key + elif key.startswith("model.") and not key.startswith("model.visual."): + hf_keys[key] = key + # Drop final norm from the transferable set. + hf_keys.pop("model.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}/") + for old, new in WEIGHT_NAME_MAPPING.items(): + name = name.replace(old, new) + 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..f50a4298 --- /dev/null +++ b/zeromodels/models/qwen_image_21/qwen_image_21_config.py @@ -0,0 +1,246 @@ +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 + # Length of the tokenized system message for Diffusers' template drop. + prompt_template_encode_start_idx: int = 14 + max_sequence_length: int = 512 + # Latent grid side when height/width omitted (64 → 1024px at VAE scale 16). + 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..af562965 --- /dev/null +++ b/zeromodels/models/qwen_image_21/qwen_image_21_layers.py @@ -0,0 +1,532 @@ +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) + # params: (B+1, D) — last row is t=0; first B rows are the real timestep. + 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) + + # fused_attention expects (B, H, S, D) + 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 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) + + # Diffusers clips only under float16 (not bfloat16). + 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 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..8b68df41 --- /dev/null +++ b/zeromodels/models/qwen_image_21/qwen_image_21_model.py @@ -0,0 +1,880 @@ +"""Qwen-Image-2.1 models: single-stream transformer, container, and T2I task.""" + +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, + 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)) + + +def _t2i_image_pad_mask(text_seq_len, latent_h, latent_w): + """Bool mask over the VLM+target-slot sequence before 2×2 expansion.""" + target_slots = (latent_h * latent_w) // IMG_TOKENS_PER_SLOT + return np.concatenate( + [ + np.zeros(text_seq_len, dtype=bool), + np.ones(target_slots, dtype=bool), + ] + ) + + +def _expand_image_pad_mask(img_mask): + """Expand each VLM image slot to ``IMG_TOKENS_PER_SLOT`` latent tokens.""" + repeats = np.where(img_mask, IMG_TOKENS_PER_SLOT, 1) + return np.repeat(img_mask, repeats) + + +@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)] + + # Prefill T2I joint layout (text + target image, no condition images). + slot_mask = _t2i_image_pad_mask(text_seq_len, sample_h, sample_w) + image_pad_mask = _expand_image_pad_mask(slot_mask) + 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) + + # Pure T2I layout: text tokens then target-image tokens (no condition images). + # Equivalent to Diffusers' expand/scatter path when img_mask is text-False + + # target-slot-True. + 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) + + # Fold encoder padding into the block-causal key mask. + from zeromodels.models.qwen_image_21.qwen_image_21_layers import MASK_NEG + + 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: + # Extra t=0 row: text tokens modulate from it; target uses the real step. + 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) + # Return only target-image tokens (Diffusers pipeline slices the same way). + 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") + ): + # Allow ``QwenImage21TextEncoderModel(config)`` via from_dict path below. + 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) + # Decoder layers only — skip final_norm (Diffusers forward hook). + 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) + # VAE emits RGBA; T2I postprocess uses RGB. + 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: + # Diffusers only enables CFG when a negative prompt is provided. + 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..641ea833 --- /dev/null +++ b/zeromodels/models/qwen_image_21/qwen_image_21_tokenizer.py @@ -0,0 +1,60 @@ +"""Qwen-Image-2.1 tokenizer: ChatML template for prompt encoding.""" + +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" +) +# Tokenized system message length (Diffusers derives this via apply_chat_template). +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..c7d46b5f --- /dev/null +++ b/zeromodels/models/qwen_image_21/qwen_image_21_vae.py @@ -0,0 +1,874 @@ +"""Qwen-Image-2.1 VAE: residual Wan-style KL autoencoder (Diffusers ``AutoencoderKLQwenImage21``). + +Diffusers' ``QwenImage21CausalConv3d`` is an image specialization of Wan's causal 3D +conv: a spatial ``Conv2d`` that squeezes / unsqueezes the singleton temporal axis. +Temporal resample paths are inactive for single-frame T2I (cold feat-cache), matching +the existing Qwen-Image ``apply_temporal=False`` convention. +""" + +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: + # Diffusers may pass (1, 1) or a 3-tuple; keep the spatial pair. + 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): + # x: (B, T, H, W, C) + 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, 2, 4, 6, 7)) + 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: (B, T, H, W, C) + 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): + # x: (B, T, H, W, C) + 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 From fc0e9dfd55477882f3f8047aad12359f42e60711 Mon Sep 17 00:00:00 2001 From: IMvision12 Date: Wed, 23 Sep 2026 00:54:28 -0700 Subject: [PATCH 2/9] fix --- ...onvert_qwen_image_21_diffusers_to_keras.py | 127 ++++++++++-------- .../models/qwen_image_21/qwen_image_21_vae.py | 6 +- 2 files changed, 73 insertions(+), 60 deletions(-) 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 index f4833d11..369bc70d 100644 --- 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 @@ -1,18 +1,3 @@ -"""Offline Diffusers ``Qwen/Qwen-Image-2.1`` → ZeroModels Keras weight conversion. - -Converts the transformer, VAE, and Qwen3-VL text encoder into a hosted -``zeromodels/qwen-image-2.1`` layout (``zm_config.json`` + sharded weights). - -Usage:: - - python -m zeromodels.models.qwen_image_21.convert_qwen_image_21_diffusers_to_keras - -Env: - ZM_OUT_DIR output directory (default ``./qwen_image_21_weights``) - HF_TOKEN optional Hub token - ZM_DTYPE ``float16`` / ``bfloat16`` / ``float32`` (default ``bfloat16``) -""" - from __future__ import annotations from typing import Dict @@ -42,25 +27,28 @@ "/beta": ".bias", "/scale": ".weight", "/": ".", - # text encoder (Qwen3-VL language tower) + "gamma": "weight", + "beta": "bias", + "kernel": "weight", +} + +TEXT_WEIGHT_NAME_MAPPING: Dict[str, str] = { + **WEIGHT_NAME_MAPPING, "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.query_norm": "self_attn.q_norm", - "attention.key_norm": "self_attn.k_norm", "attention_norm": "input_layernorm", "mlp_norm": "post_attention_layernorm", "mlp.gate": "mlp.gate_proj", "mlp.up": "mlp.up_proj", "mlp.down": "mlp.down_proj", - "gamma": "weight", - "beta": "bias", - "kernel": "weight", } @@ -103,7 +91,7 @@ def config_from_diffusers(repo, token=None): if k == "_class_name" or not k.startswith("_") } temperal = tuple(vae.get("temperal_downsample", (False, True, True, True))) - rope = text_inner.get("rope_scaling") or {} + rope = text_inner.get("rope_parameters") or text_inner.get("rope_scaling") or {} return QwenImage21Config( transformer_config={ **QwenImage21Transformer2DModel.kwargs_from_diffusers_config(transformer), @@ -136,7 +124,9 @@ def config_from_diffusers(repo, token=None): "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": text_inner.get("rope_theta", 5000000.0), + "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, @@ -171,7 +161,12 @@ def transfer_qwen_image_21( with build_dtype_scope(dtype), zeros_init(): model = QwenImage21Model(**flat) - vae_mapping = {k: v for k, v in WEIGHT_NAME_MAPPING.items() if "/gamma" not in k} + # 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( ( ( @@ -239,49 +234,65 @@ def __iter__(self): consumed = set() trainable, non_trainable = split_model_weights(component) + # VAE / DiT nest safe_name prefixes; leaf-pair mapping matches Qwen-Image 1.x. + prefer_leaf = subfolder in ("vae", "transformer") for keras_weight, _ in tqdm( trainable + non_trainable, desc=f"Transferring {subfolder} weights to Keras", ): - key = "/".join(keras_weight.path.split("/")[-2:]) - # Prefer full path remapping from the component root. - full = keras_weight.path - # Strip leading component name - parts = full.split("/") - if parts and parts[0] in ( - component.name, - "transformer", - "vae", - "QwenImage21Transformer2DModel", - "AutoencoderKLQwenImage21", - ): - rel = "/".join(parts[1:]) + leaf = "/".join(keras_weight.path.split("/")[-2:]) + if prefer_leaf: + candidates = [leaf] else: - rel = "/".join(parts) - key = rel - for old, new in mapping.items(): - key = key.replace(old, new) - if key not in state: - # Fall back to leaf-pair mapping used by Qwen-Image 1.0. - key = "/".join(keras_weight.path.split("/")[-2:]) + 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(): - key = key.replace(old, new) - if key not in state: - raise WeightMappingError(keras_weight.path, key) + 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 ( - arr.ndim > 1 - and len(kshape) == 1 - and int(np.prod(arr.shape)) == kshape[0] - ): - arr = arr.reshape(kshape) + 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) + ): + # Diffusers mid-block attention uses 1×1 Conv2d; Keras uses Dense. + # Keep torch (out, in) layout — transfer_weights will transpose. + 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( @@ -334,7 +345,7 @@ def __iter__(self): text_encoder.weights, desc="Transferring text_encoder weights to Keras" ): name = weight.path.removeprefix(f"{text_encoder.name}/") - for old, new in WEIGHT_NAME_MAPPING.items(): + for old, new in TEXT_WEIGHT_NAME_MAPPING.items(): name = name.replace(old, new) if name not in hf_keys: raise WeightMappingError(weight.path, name) diff --git a/zeromodels/models/qwen_image_21/qwen_image_21_vae.py b/zeromodels/models/qwen_image_21/qwen_image_21_vae.py index c7d46b5f..ee36fe7d 100644 --- a/zeromodels/models/qwen_image_21/qwen_image_21_vae.py +++ b/zeromodels/models/qwen_image_21/qwen_image_21_vae.py @@ -134,7 +134,8 @@ def __init__(self, in_channels, out_channels, factor_t, factor_s=1, **kwargs): self.group_size = self.in_channels * self.factor // self.out_channels def call(self, x): - # x: (B, T, H, W, C) + # x: (B, T, H, W, C) — pack order must match Diffusers NCHW AvgDown3D: + # channels outer, then (factor_t, factor_s, factor_s). 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))) @@ -144,7 +145,8 @@ def call(self, x): 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, 2, 4, 6, 7)) + # B,T',ft,H',fs,W',fs,C → B,T',H',W',C,ft,fs,fs + 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, From 615d3ad2978472b8b0864b9c28d3eb55daadac87 Mon Sep 17 00:00:00 2001 From: IMvision12 Date: Wed, 23 Sep 2026 01:10:08 -0700 Subject: [PATCH 3/9] fix --- .../qwen_image_21/convert_qwen_image_21_diffusers_to_keras.py | 2 ++ 1 file changed, 2 insertions(+) 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 index 369bc70d..0f09ef4d 100644 --- 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 @@ -156,6 +156,8 @@ def transfer_qwen_image_21( 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(): From a5eb6dae069caedc116fe5d69c3ac32232350a75 Mon Sep 17 00:00:00 2001 From: IMvision12 Date: Wed, 23 Sep 2026 01:20:01 -0700 Subject: [PATCH 4/9] fix --- ...onvert_qwen_image_21_diffusers_to_keras.py | 22 ++++++++++++++----- 1 file changed, 17 insertions(+), 5 deletions(-) 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 index 0f09ef4d..9187879b 100644 --- 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 @@ -332,23 +332,35 @@ def __iter__(self): # + final norm (Diffusers reads pre-norm hidden states). if key.startswith("model.visual.") or key.startswith("lm_head."): continue - if key in ("model.norm.weight",) or key.endswith(".norm.weight") and key.count(".") <= 2: - if key == "model.norm.weight": - continue + if key == "model.norm.weight": + continue if key.startswith("model.language_model."): - hf_keys["model." + key[len("model.language_model.") :]] = key + 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 - # Drop final norm from the transferable set. + hf_keys[key[len("model.") :]] = key + # Drop final norm from the transferable set (pre-norm embeds for the DiT). 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}/") + # Also strip a nested language_model/ prefix if present. + if name.startswith("language_model/"): + name = name[len("language_model/") :] for old, new in TEXT_WEIGHT_NAME_MAPPING.items(): name = name.replace(old, new) + # token_embedding.embeddings → model.embed_tokens.weight; without the + # language_model remap, bare embed_tokens.weight must also resolve. + 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) From 9d3c12f441091a0f2b1d7de6f2db578a39ff706a Mon Sep 17 00:00:00 2001 From: IMvision12 Date: Wed, 23 Sep 2026 01:33:22 -0700 Subject: [PATCH 5/9] fix --- zeromodels/models/qwen3_vl/qwen3_vl_layers.py | 50 +++++++++++++++++++ .../qwen_image_21/qwen_image_21_layers.py | 14 ++++++ 2 files changed, 64 insertions(+) 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/qwen_image_21_layers.py b/zeromodels/models/qwen_image_21/qwen_image_21_layers.py index af562965..6daee871 100644 --- a/zeromodels/models/qwen_image_21/qwen_image_21_layers.py +++ b/zeromodels/models/qwen_image_21/qwen_image_21_layers.py @@ -430,6 +430,10 @@ def _unflatten(x): def compute_output_shape(self, input_shape): return tuple(input_shape) + def compute_output_spec(self, hidden_states, rotary_emb=None, attention_mask=None): + # Avoid seq² attention temps while tracing the Functional graph on GPU. + return keras.KerasTensor(hidden_states.shape, dtype=self.compute_dtype) + def get_config(self): config = super().get_config() config.update( @@ -517,6 +521,16 @@ def call(self, hidden_states, modulation, rotary_emb=None, attention_mask=None, 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( From a7eb88e74a3a8dfe0a555aa289f98cc0a38b8a3b Mon Sep 17 00:00:00 2001 From: IMvision12 Date: Wed, 23 Sep 2026 01:54:46 -0700 Subject: [PATCH 6/9] fix --- ...onvert_qwen_image_21_diffusers_to_keras.py | 36 +++++++++---------- 1 file changed, 18 insertions(+), 18 deletions(-) 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 index 9187879b..846da6a9 100644 --- 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 @@ -277,24 +277,24 @@ def __iter__(self): 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) - ): - # Diffusers mid-block attention uses 1×1 Conv2d; Keras uses Dense. - # Keep torch (out, in) layout — transfer_weights will transpose. - 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(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) + ): + # Diffusers mid-block attention uses 1×1 Conv2d; Keras uses Dense. + # Keep torch (out, in) layout — transfer_weights will transpose. + 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( From 6fc741531b4c5036eb3195bddaa401d0126f630a Mon Sep 17 00:00:00 2001 From: IMvision12 Date: Wed, 23 Sep 2026 02:14:02 -0700 Subject: [PATCH 7/9] fix --- ...onvert_qwen_image_21_diffusers_to_keras.py | 16 +---- .../qwen_image_21/qwen_image_21_config.py | 2 - .../qwen_image_21/qwen_image_21_layers.py | 30 ++++----- .../qwen_image_21/qwen_image_21_model.py | 61 ++++++------------- .../qwen_image_21/qwen_image_21_tokenizer.py | 3 - .../models/qwen_image_21/qwen_image_21_vae.py | 18 +----- 6 files changed, 36 insertions(+), 94 deletions(-) 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 index 846da6a9..c2266c0d 100644 --- 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 @@ -30,10 +30,7 @@ "gamma": "weight", "beta": "bias", "kernel": "weight", -} - -TEXT_WEIGHT_NAME_MAPPING: Dict[str, str] = { - **WEIGHT_NAME_MAPPING, + # Qwen3-VL text tower "token_embedding.embeddings": "model.embed_tokens.weight", "language_model.final_norm.weight": "model.norm.weight", "language_model.": "model.", @@ -236,7 +233,6 @@ def __iter__(self): consumed = set() trainable, non_trainable = split_model_weights(component) - # VAE / DiT nest safe_name prefixes; leaf-pair mapping matches Qwen-Image 1.x. prefer_leaf = subfolder in ("vae", "transformer") for keras_weight, _ in tqdm( trainable + non_trainable, @@ -286,8 +282,6 @@ def __iter__(self): and arr.ndim == 4 and tuple(arr.shape[-2:]) == (1, 1) ): - # Diffusers mid-block attention uses 1×1 Conv2d; Keras uses Dense. - # Keep torch (out, in) layout — transfer_weights will transpose. arr = arr[:, :, 0, 0] elif ( arr.ndim > 1 @@ -328,8 +322,6 @@ def __iter__(self): } hf_keys = {} for key in weight_map: - # Text-only tower: keep language_model / embed_tokens; drop vision + lm_head - # + final norm (Diffusers reads pre-norm hidden states). if key.startswith("model.visual.") or key.startswith("lm_head."): continue if key == "model.norm.weight": @@ -341,7 +333,6 @@ def __iter__(self): elif key.startswith("model.") and not key.startswith("model.visual."): hf_keys[key] = key hf_keys[key[len("model.") :]] = key - # Drop final norm from the transferable set (pre-norm embeds for the DiT). hf_keys.pop("model.norm.weight", None) hf_keys.pop("norm.weight", None) consumed = set() @@ -350,13 +341,10 @@ def __iter__(self): text_encoder.weights, desc="Transferring text_encoder weights to Keras" ): name = weight.path.removeprefix(f"{text_encoder.name}/") - # Also strip a nested language_model/ prefix if present. if name.startswith("language_model/"): name = name[len("language_model/") :] - for old, new in TEXT_WEIGHT_NAME_MAPPING.items(): + for old, new in WEIGHT_NAME_MAPPING.items(): name = name.replace(old, new) - # token_embedding.embeddings → model.embed_tokens.weight; without the - # language_model remap, bare embed_tokens.weight must also resolve. if name not in hf_keys and name.startswith("model."): alt = name[len("model.") :] if alt in hf_keys: diff --git a/zeromodels/models/qwen_image_21/qwen_image_21_config.py b/zeromodels/models/qwen_image_21/qwen_image_21_config.py index f50a4298..4d2e289b 100644 --- a/zeromodels/models/qwen_image_21/qwen_image_21_config.py +++ b/zeromodels/models/qwen_image_21/qwen_image_21_config.py @@ -235,10 +235,8 @@ class QwenImage21Config(BaseConfig): vae_config: QwenImage21VAEConfig | dict | None = None text_config: QwenImage21TextConfig | dict | None = None scheduler_config: dict | None = None - # Length of the tokenized system message for Diffusers' template drop. prompt_template_encode_start_idx: int = 14 max_sequence_length: int = 512 - # Latent grid side when height/width omitted (64 → 1024px at VAE scale 16). default_sample_size: int = 64 bos_token_id: int = 151643 eos_token_id: int = 151645 diff --git a/zeromodels/models/qwen_image_21/qwen_image_21_layers.py b/zeromodels/models/qwen_image_21/qwen_image_21_layers.py index 6daee871..8136275b 100644 --- a/zeromodels/models/qwen_image_21/qwen_image_21_layers.py +++ b/zeromodels/models/qwen_image_21/qwen_image_21_layers.py @@ -21,7 +21,7 @@ IMG_TOKENS_PER_SLOT = 4 -def _rope_params(index, dim, theta): +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), @@ -43,7 +43,7 @@ def build_qwenimage21_rope_angles(img_shapes, image_pad_mask, axes_dim, theta=10 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) + np.concatenate([rope_params(pos_index, dim, theta), rope_params(neg_index, dim, theta)], axis=0) for dim in axes_dim ] @@ -337,7 +337,7 @@ def build(self, input_shapes): 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) + scale = select_modulation_rows(scale, target_token_mask) return self.norm(hidden_states) * (1.0 + scale) def compute_output_shape(self, input_shapes): @@ -356,11 +356,10 @@ def get_config(self): return config -def _select_modulation_rows(params, target_token_mask): +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) - # params: (B+1, D) — last row is t=0; first B rows are the real timestep. 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)) @@ -404,19 +403,18 @@ def call(self, hidden_states, rotary_emb=None, attention_mask=None): key = self.to_k(hidden_states) value = self.to_v(hidden_states) - def _unflatten(x): + 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) + 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) - # fused_attention expects (B, H, S, D) query = ops.transpose(query, (0, 2, 1, 3)) key = ops.transpose(key, (0, 2, 1, 3)) value = ops.transpose(value, (0, 2, 1, 3)) @@ -431,7 +429,6 @@ def compute_output_shape(self, input_shape): return tuple(input_shape) def compute_output_spec(self, hidden_states, rotary_emb=None, attention_mask=None): - # Avoid seq² attention temps while tracing the Functional graph on GPU. return keras.KerasTensor(hidden_states.shape, dtype=self.compute_dtype) def get_config(self): @@ -492,15 +489,15 @@ def build(self, input_shape): self.img_mlp.build(input_shape) self.built = True - def _modulate(self, hidden_states, mod_params, target_token_mask): + 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) + 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( + img_modulated, img_gate1 = self.modulate( self.img_norm1(hidden_states), mod1, target_token_mask ) attn_output = self.attn( @@ -508,12 +505,11 @@ def call(self, hidden_states, modulation, rotary_emb=None, attention_mask=None, ) hidden_states = hidden_states + ops.tanh(img_gate1) * attn_output - img_modulated2, img_gate2 = self._modulate( + 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) - # Diffusers clips only under float16 (not bfloat16). if keras.backend.standardize_dtype(hidden_states.dtype) == "float16": hidden_states = ops.clip(hidden_states, -65504.0, 65504.0) return hidden_states diff --git a/zeromodels/models/qwen_image_21/qwen_image_21_model.py b/zeromodels/models/qwen_image_21/qwen_image_21_model.py index 8b68df41..d16e7076 100644 --- a/zeromodels/models/qwen_image_21/qwen_image_21_model.py +++ b/zeromodels/models/qwen_image_21/qwen_image_21_model.py @@ -1,5 +1,3 @@ -"""Qwen-Image-2.1 models: single-stream transformer, container, and T2I task.""" - from __future__ import annotations import keras @@ -23,6 +21,7 @@ ) from zeromodels.models.qwen_image_21.qwen_image_21_layers import ( IMG_TOKENS_PER_SLOT, + MASK_NEG, QwenImage21AdaLayerNormContinuous, QwenImage21TextProjection, QwenImage21TimestepProjEmbeddings, @@ -56,23 +55,6 @@ def unpack_latents(latents, height, width, channels): return ops.reshape(latents, (batch, height, width, channels)) -def _t2i_image_pad_mask(text_seq_len, latent_h, latent_w): - """Bool mask over the VLM+target-slot sequence before 2×2 expansion.""" - target_slots = (latent_h * latent_w) // IMG_TOKENS_PER_SLOT - return np.concatenate( - [ - np.zeros(text_seq_len, dtype=bool), - np.ones(target_slots, dtype=bool), - ] - ) - - -def _expand_image_pad_mask(img_mask): - """Expand each VLM image slot to ``IMG_TOKENS_PER_SLOT`` latent tokens.""" - repeats = np.where(img_mask, IMG_TOKENS_PER_SLOT, 1) - return np.repeat(img_mask, repeats) - - @keras.saving.register_keras_serializable(package="zeromodels") class AutoencoderKLQwenImage21(BaseModel): """Qwen-Image-2.1 VAE (Diffusers ``AutoencoderKLQwenImage21``), channels-last. @@ -195,7 +177,7 @@ def __init__( self.quant_conv = quant_conv self.post_quant_conv = post_quant_conv - def _to_ndhwc(self, x, is_latent=False): + def to_ndhwc(self, x, is_latent=False): static_ndim = len(x.shape) if static_ndim == 4: return ops.expand_dims(x, axis=1) @@ -210,7 +192,7 @@ def _to_ndhwc(self, x, is_latent=False): def encode(self, x, sample=False, seed=None): was_4d = len(x.shape) == 4 - x5 = self._to_ndhwc(x, is_latent=False) + 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: @@ -226,7 +208,7 @@ def encode(self, x, sample=False, seed=None): def decode(self, z): was_4d = len(z.shape) == 4 - z5 = self._to_ndhwc(z, is_latent=True) + 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: @@ -298,9 +280,16 @@ def __init__( img_seq = sample_h * sample_w img_shapes = [(1, sample_h, sample_w)] - # Prefill T2I joint layout (text + target image, no condition images). - slot_mask = _t2i_image_pad_mask(text_seq_len, sample_h, sample_w) - image_pad_mask = _expand_image_pad_mask(slot_mask) + 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( @@ -352,9 +341,6 @@ def __init__( hidden = img_in(sample_in) encoder = txt_in(enc_in) - # Pure T2I layout: text tokens then target-image tokens (no condition images). - # Equivalent to Diffusers' expand/scatter path when img_mask is text-False + - # target-slot-True. joint = ops.concatenate([encoder, hidden], axis=1) rotary = ops.convert_to_tensor(rope_angles) @@ -362,9 +348,6 @@ def __init__( target_mask_t = ops.convert_to_tensor(target_token_mask) text_pos = np.flatnonzero(~image_pad_mask).astype(np.int32) - # Fold encoder padding into the block-causal key mask. - from zeromodels.models.qwen_image_21.qwen_image_21_layers import MASK_NEG - 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 @@ -377,7 +360,6 @@ def __init__( attn_mask = attn_mask + (1.0 - key_ok)[:, None, None, :] * MASK_NEG if causal_condition: - # Extra t=0 row: text tokens modulate from it; target uses the real step. t_all = ops.concatenate( [timestep_in, ops.zeros((1,), dtype=timestep_in.dtype)], axis=0 ) @@ -398,7 +380,6 @@ def __init__( ) joint = norm_out([joint, temb, mod_mask]) output_full = proj_out(joint) - # Return only target-image tokens (Diffusers pipeline slices the same way). output = output_full[:, -img_seq:, :] super().__init__( inputs={ @@ -470,7 +451,6 @@ def __init__(self, name="text_encoder", **kwargs): if "max_seq_len" in kwargs and not any( k in kwargs for k in ("embed_dim", "num_layers", "vocab_size") ): - # Allow ``QwenImage21TextEncoderModel(config)`` via from_dict path below. pass config = self.config_class.from_dict(kwargs) if kwargs else self.config_class() max_seq_len = int(getattr(config, "max_seq_len", 1024)) @@ -499,7 +479,6 @@ def __init__(self, name="text_encoder", **kwargs): pos, config.head_dim, config.rope_theta, tuple(config.mrope_section) ) mask = causal_mask(input_ids, attention_mask) - # Decoder layers only — skip final_norm (Diffusers forward hook). for layer in language_model.decoder_layers: hidden = layer(hidden, cos, sin, attention_mask=mask) @@ -702,10 +681,10 @@ def vae_scale_factor(self): @property def latent_shape(self): - h, w = self._latent_side() + h, w = self.latent_side() return (h * w, self.vae.z_dim) - def _latent_side(self, height=None, width=None): + 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 @@ -762,7 +741,7 @@ def predict_noise(self, latents, timesteps, embeddings): )["sample"] def decode_latents(self, latents, height=None, width=None): - h, w = self._latent_side(height, width) + 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") @@ -770,13 +749,12 @@ def decode_latents(self, latents, height=None, width=None): std = ops.reshape(std, (1, 1, 1, -1)) latents = latents * std + mean image = self.vae.decode(latents) - # VAE emits RGBA; T2I postprocess uses RGB. 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( + h, w = self.latent_side( getattr(self, "_gen_height", None), getattr(self, "_gen_width", None) ) channels = self.vae.z_dim @@ -836,7 +814,6 @@ def generate( with inference_scope(): embeddings = self.encode_prompt(input_ids, attention_mask, **conditioning) if guidance_scale > 1.0 and negative_input_ids is None: - # Diffusers only enables CFG when a negative prompt is provided. do_cfg = False uncond = None elif negative_input_ids is not None and guidance_scale > 1.0: @@ -846,7 +823,7 @@ def generate( uncond = None do_cfg = False - h, w = self._latent_side(height, width) + 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) diff --git a/zeromodels/models/qwen_image_21/qwen_image_21_tokenizer.py b/zeromodels/models/qwen_image_21/qwen_image_21_tokenizer.py index 641ea833..f0975ba3 100644 --- a/zeromodels/models/qwen_image_21/qwen_image_21_tokenizer.py +++ b/zeromodels/models/qwen_image_21/qwen_image_21_tokenizer.py @@ -1,5 +1,3 @@ -"""Qwen-Image-2.1 tokenizer: ChatML template for prompt encoding.""" - import keras from zeromodels.models.qwen3.qwen3_tokenizer import Qwen3Tokenizer @@ -10,7 +8,6 @@ f"<|im_start|>user\n{{}}<|im_end|>\n" f"<|im_start|>assistant\n" ) -# Tokenized system message length (Diffusers derives this via apply_chat_template). PROMPT_TEMPLATE_START_IDX = 14 diff --git a/zeromodels/models/qwen_image_21/qwen_image_21_vae.py b/zeromodels/models/qwen_image_21/qwen_image_21_vae.py index ee36fe7d..3069dd15 100644 --- a/zeromodels/models/qwen_image_21/qwen_image_21_vae.py +++ b/zeromodels/models/qwen_image_21/qwen_image_21_vae.py @@ -1,11 +1,3 @@ -"""Qwen-Image-2.1 VAE: residual Wan-style KL autoencoder (Diffusers ``AutoencoderKLQwenImage21``). - -Diffusers' ``QwenImage21CausalConv3d`` is an image specialization of Wan's causal 3D -conv: a spatial ``Conv2d`` that squeezes / unsqueezes the singleton temporal axis. -Temporal resample paths are inactive for single-frame T2I (cold feat-cache), matching -the existing Qwen-Image ``apply_temporal=False`` convention. -""" - from __future__ import annotations import keras @@ -36,7 +28,6 @@ def __init__( if isinstance(kernel_size, int): self.kernel_size = (kernel_size, kernel_size) else: - # Diffusers may pass (1, 1) or a 3-tuple; keep the spatial pair. ks = tuple(kernel_size) self.kernel_size = (ks[-2], ks[-1]) if len(ks) == 3 else (ks[0], ks[1]) if isinstance(stride, int): @@ -134,8 +125,6 @@ def __init__(self, in_channels, out_channels, factor_t, factor_s=1, **kwargs): self.group_size = self.in_channels * self.factor // self.out_channels def call(self, x): - # x: (B, T, H, W, C) — pack order must match Diffusers NCHW AvgDown3D: - # channels outer, then (factor_t, factor_s, factor_s). 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))) @@ -145,7 +134,6 @@ def call(self, x): w = ops.shape(x)[3] c = self.in_channels x = ops.reshape(x, (b, t // ft, ft, h // fs, fs, w // fs, fs, c)) - # B,T',ft,H',fs,W',fs,C → B,T',H',W',C,ft,fs,fs 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( @@ -158,13 +146,13 @@ 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): + 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) + 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() @@ -194,7 +182,6 @@ def __init__(self, in_channels, out_channels, factor_t, factor_s=1, **kwargs): self.repeats = self.out_channels * self.factor // self.in_channels def call(self, x, first_chunk=False): - # x: (B, T, H, W, C) x = ops.repeat(x, self.repeats, axis=-1) b = ops.shape(x)[0] t = ops.shape(x)[1] @@ -412,7 +399,6 @@ def build(self, input_shape): self.built = True def call(self, x): - # x: (B, T, H, W, C) residual = x x = self.norm(x) b = ops.shape(x)[0] From d94d84321b775e8fe1ef61dbbc7d1e51ecb6d1c9 Mon Sep 17 00:00:00 2001 From: IMvision12 Date: Wed, 23 Sep 2026 02:15:39 -0700 Subject: [PATCH 8/9] docs --- README.md | 2 +- docs/qwen_image.md | 1 + docs/qwen_image_21.md | 256 ++++++++++++++++++++++++++++++++++++++++-- 3 files changed, 250 insertions(+), 9 deletions(-) diff --git a/README.md b/README.md index 4229123e..900108b9 100644 --- a/README.md +++ b/README.md @@ -257,7 +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) | `diffusers` | + | Qwen-Image-2.1 | [Qwen/Qwen-Image-2.1](https://huggingface.co/Qwen/Qwen-Image-2.1) (single-stream DiT + Qwen3-VL) | `diffusers` |
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 index 9df921c5..934a32c5 100644 --- a/docs/qwen_image_21.md +++ b/docs/qwen_image_21.md @@ -10,13 +10,16 @@ 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), and the Qwen3-VL text tower. The whole model is **one container**, -`QwenImage21Model`. `QwenImage21TextToImage` adds `generate`. +(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). + 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), @@ -24,16 +27,24 @@ Key facts of the port: - **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. + 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: -- Reference: [diffusers `QwenImage21Pipeline`](https://huggingface.co/docs/diffusers/api/pipelines/qwenimage) - Source: [`Qwen/Qwen-Image-2.1`](https://huggingface.co/Qwen/Qwen-Image-2.1) -- See also [qwen_image.md](qwen_image.md), [qwen3_vl.md](qwen3_vl.md) +- 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) | @@ -42,9 +53,145 @@ 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, @@ -52,11 +199,104 @@ from zeromodels.models.qwen_image_21 import ( model = QwenImage21TextToImage.from_weights("zeromodels/qwen-image-2.1") tokenizer = QwenImage21Tokenizer.from_weights("zeromodels/qwen-image-2.1") -image = model.generate(**tokenizer("a capybara in a wizard hat"), height=1024, width=1024) + +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, +) ``` -Offline conversion (no on-the-fly `hf:` for diffusion):: +### 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. From 4ace418c648795cb18aee5481194f39e8eb057af Mon Sep 17 00:00:00 2001 From: IMvision12 Date: Wed, 23 Sep 2026 02:16:30 -0700 Subject: [PATCH 9/9] format --- docs/qwen_image_21.md | 8 +-- tests/fixtures/dummy_inputs.py | 4 +- ...onvert_qwen_image_21_diffusers_to_keras.py | 13 ++-- .../qwen_image_21/qwen_image_21_layers.py | 65 +++++++++++++++---- .../qwen_image_21/qwen_image_21_model.py | 22 +++---- .../models/qwen_image_21/qwen_image_21_vae.py | 52 +++++++++++---- 6 files changed, 113 insertions(+), 51 deletions(-) diff --git a/docs/qwen_image_21.md b/docs/qwen_image_21.md index 934a32c5..47657360 100644 --- a/docs/qwen_image_21.md +++ b/docs/qwen_image_21.md @@ -200,9 +200,7 @@ from zeromodels.models.qwen_image_21 import ( 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" -) +inputs = tokenizer("a photo of a capybara wearing a wizard hat, soft window light") images = model.generate( **inputs, height=1024, @@ -255,9 +253,7 @@ model = QwenImage21TextToImage.from_weights( transformer_sample_size=32, vae_sample_size=512, ) -images = model.generate( - **tokenizer("a mountain lake at dawn"), height=512, width=512 -) +images = model.generate(**tokenizer("a mountain lake at dawn"), height=512, width=512) ``` ### Offline conversion diff --git a/tests/fixtures/dummy_inputs.py b/tests/fixtures/dummy_inputs.py index a391575d..8a312a50 100644 --- a/tests/fixtures/dummy_inputs.py +++ b/tests/fixtures/dummy_inputs.py @@ -279,9 +279,7 @@ def qwen_image_21_input( 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": ops.ones((batch_size, text_seq_len, context_in_dim)), "encoder_hidden_states_mask": ops.ones( (batch_size, text_seq_len), dtype="int32" ), 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 index c2266c0d..845135b3 100644 --- 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 @@ -162,9 +162,7 @@ def transfer_qwen_image_21( # 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") + 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( ( @@ -222,7 +220,8 @@ def __iter__(self): raise FileNotFoundError except Exception: path = hf_hub_download( - repo, filename or "diffusion_pytorch_model.safetensors", + repo, + filename or "diffusion_pytorch_model.safetensors", subfolder=subfolder, token=token, ) @@ -277,11 +276,7 @@ def __iter__(self): 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) - ): + elif len(kshape) == 2 and arr.ndim == 4 and tuple(arr.shape[-2:]) == (1, 1): arr = arr[:, :, 0, 0] elif ( arr.ndim > 1 diff --git a/zeromodels/models/qwen_image_21/qwen_image_21_layers.py b/zeromodels/models/qwen_image_21/qwen_image_21_layers.py index 8136275b..25cbad0c 100644 --- a/zeromodels/models/qwen_image_21/qwen_image_21_layers.py +++ b/zeromodels/models/qwen_image_21/qwen_image_21_layers.py @@ -43,7 +43,10 @@ def build_qwenimage21_rope_angles(img_shapes, image_pad_mask, axes_dim, theta=10 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) + np.concatenate( + [rope_params(pos_index, dim, theta), rope_params(neg_index, dim, theta)], + axis=0, + ) for dim in axes_dim ] @@ -62,7 +65,11 @@ def build_qwenimage21_rope_angles(img_shapes, image_pad_mask, axes_dim, theta=10 position += max(height, width) image_height_index.extend( - [h for h in range(-(height - height // 2), height // 2) for _ in range(width)] + [ + 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)] @@ -220,7 +227,9 @@ def get_config(self): class QwenImage21TextProjection(layers.Layer): """Diffusers ``QwenImage21TextProjection``.""" - def __init__(self, context_in_dim, hidden_size, eps=NORM_EPS, module_path="txt_in", **kwargs): + 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 @@ -313,7 +322,14 @@ def get_config(self): 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): + 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 @@ -324,7 +340,11 @@ def __init__(self, embedding_dim, conditioning_embedding_dim, eps=NORM_EPS, modu 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") + axis=-1, + epsilon=eps, + center=False, + scale=False, + name=safe_name(f"{module_path}.norm"), ) def build(self, input_shapes): @@ -384,7 +404,9 @@ def __init__(self, dim, heads, dim_head, eps=NORM_EPS, module_path=None, **kwarg 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.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") @@ -470,13 +492,25 @@ def __init__( 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") + 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" + 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") + 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" @@ -495,7 +529,14 @@ def modulate(self, hidden_states, mod_params, 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): + 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 @@ -508,7 +549,9 @@ def call(self, hidden_states, modulation, rotary_emb=None, attention_mask=None, 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) + 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) diff --git a/zeromodels/models/qwen_image_21/qwen_image_21_model.py b/zeromodels/models/qwen_image_21/qwen_image_21_model.py index d16e7076..d9e97445 100644 --- a/zeromodels/models/qwen_image_21/qwen_image_21_model.py +++ b/zeromodels/models/qwen_image_21/qwen_image_21_model.py @@ -10,7 +10,10 @@ FlowMatchEulerDiscreteScheduler, get_scheduler, ) -from zeromodels.models.qwen3_vl.qwen3_vl_model import Qwen3VLTextModel, qwen3_text_cos_sin +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, @@ -37,9 +40,7 @@ ) from zeromodels.models.stable_diffusion.stable_diffusion_layers import safe_name -QWEN_IMAGE_21_HUB_SIBLINGS = frozenset( - {"QwenImage21Model", "QwenImage21TextToImage"} -) +QWEN_IMAGE_21_HUB_SIBLINGS = frozenset({"QwenImage21Model", "QwenImage21TextToImage"}) def pack_latents(latents, height, width): @@ -103,7 +104,10 @@ def __init__( 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 + h_lat, w_lat = ( + h_img // spatial_compression_ratio, + w_img // spatial_compression_ratio, + ) encoder = QwenImage21Encoder3d( dim=base_dim, @@ -303,9 +307,7 @@ def __init__( 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") - ) + 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") ) @@ -590,9 +592,7 @@ def build_graph(self, config, components): "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" - ), + "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" diff --git a/zeromodels/models/qwen_image_21/qwen_image_21_vae.py b/zeromodels/models/qwen_image_21/qwen_image_21_vae.py index 3069dd15..7bbecaad 100644 --- a/zeromodels/models/qwen_image_21/qwen_image_21_vae.py +++ b/zeromodels/models/qwen_image_21/qwen_image_21_vae.py @@ -152,7 +152,13 @@ def down(size, factor): 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) + 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() @@ -259,7 +265,9 @@ 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)) + 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) @@ -291,7 +299,13 @@ def call(self, x): 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) + 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, @@ -326,10 +340,18 @@ def __init__(self, in_dim, out_dim, module_path, dropout=0.0, **kwargs): 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.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( @@ -388,9 +410,15 @@ def __init__(self, dim, module_path, **kwargs): 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")) + 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) @@ -435,7 +463,9 @@ def __init__(self, dim, module_path, dropout=0.0, num_layers=1, **kwargs): self.dim = int(dim) self.module_path = module_path self.resnets = [ - QwenImage21ResidualBlock(dim, dim, f"{module_path}.resnets.0", dropout=dropout) + QwenImage21ResidualBlock( + dim, dim, f"{module_path}.resnets.0", dropout=dropout + ) ] self.attentions = [] for i in range(num_layers):