Skip to content

Commit b6c4e39

Browse files
committed
Optimize Z-Image-Turbo serving pipeline and attention layouts
- Pre-split text_encoder and transformer during initialization to avoid host-side splitting CPU overhead. - Fully tensorize ZImageTransformer2DModel __call__ forward path using statically-padded 3D stacked sequences, bypassing cross-device all-reduce/all-gather collectives and loops. - Pass 4D transposed layouts directly to NNXAttentionOp to bypass unflatten/flatten transposes. - Maintain RMSNorm outputs in float32 and pass them directly into _apply_rope to enable seamless on-device loop/kernel fusion and avoid float32 <-> bfloat16 casting loops. - Inline latent generation logic inside pipeline call, removing redundant JIT overhead.
1 parent bda41e1 commit b6c4e39

21 files changed

Lines changed: 2518 additions & 226 deletions

src/maxdiffusion/checkpointing/checkpointing_utils.py

Lines changed: 11 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -36,6 +36,7 @@
3636
STABLE_DIFFUSION_XL_CHECKPOINT = "STABLE_DIFUSSION_XL_CHECKPOINT"
3737
FLUX_CHECKPOINT = "FLUX_CHECKPOINT"
3838
WAN_CHECKPOINT = "WAN_CHECKPOINT"
39+
Z_IMAGE_CHECKPOINT = "Z_IMAGE_CHECKPOINT"
3940

4041

4142
def create_orbax_checkpoint_manager(
@@ -62,6 +63,16 @@ def create_orbax_checkpoint_manager(
6263
item_handlers = None
6364
if checkpoint_type == FLUX_CHECKPOINT:
6465
item_names = ("flux_state", "flux_config", "vae_state", "vae_config", "scheduler", "scheduler_config")
66+
elif checkpoint_type == Z_IMAGE_CHECKPOINT:
67+
# Only `transformer_state` is trainable; the VAE and the Qwen3 text encoder
68+
# are frozen and kept as separate non-trainable items.
69+
item_names = ("transformer_state", "vae_state", "text_encoder_state", "z_image_config")
70+
item_handlers = {
71+
"z_image_config": ocp.JsonCheckpointHandler(),
72+
"transformer_state": ocp.StandardCheckpointHandler(),
73+
"vae_state": ocp.StandardCheckpointHandler(),
74+
"text_encoder_state": ocp.StandardCheckpointHandler(),
75+
}
6576
elif checkpoint_type == WAN_CHECKPOINT:
6677
item_names = ("low_noise_transformer_state", "high_noise_transformer_state", "wan_state", "wan_config")
6778
item_handlers = {

src/maxdiffusion/checkpointing/z_image_checkpointer.py

Lines changed: 415 additions & 0 deletions
Large diffs are not rendered by default.

src/maxdiffusion/common_types.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -53,6 +53,7 @@
5353
WAN2_2 = "wan2.2"
5454
LTX2_VIDEO = "ltx2_video"
5555
LTX2_3 = "ltx2.3"
56+
Z_IMAGE = "z_image"
5657

5758
WAN_MODEL = WAN2_1
5859

Lines changed: 157 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,157 @@
1+
# Copyright 2026 Google LLC
2+
#
3+
# Licensed under the Apache License, Version 2.0 (the "License");
4+
# you may not use this file except in compliance with the License.
5+
# You may obtain a copy of the License at
6+
#
7+
# https://www.apache.org/licenses/LICENSE-2.0
8+
#
9+
# Unless required by applicable law or agreed to in writing, software
10+
# distributed under the License is distributed on an "AS IS" BASIS,
11+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12+
# See the License for the specific language governing permissions and
13+
# limitations under the License.
14+
15+
# Official Z-Image inference defaults. Z-Image and Z-Image-Turbo share the
16+
# denoiser; only the recommended denoising-step count differs (see
17+
# base_zimage_turbo.yml).
18+
19+
run_name: 'zimage'
20+
21+
# If true save config to GCS in {base_output_directory}/{run_name}/
22+
save_config_to_gcs: False
23+
24+
pretrained_model_name_or_path: 'Tongyi-MAI/Z-Image'
25+
model_name: z_image
26+
model_type: 'T2I'
27+
28+
unet_checkpoint: ''
29+
# This will convert the weights to this dtype.
30+
# When running inference on TPUv5e, use weights_dtype: 'bfloat16'
31+
weights_dtype: 'bfloat16'
32+
# This sets the layer's dtype in the model. Ex: nn.Dense(dtype=activations_dtype)
33+
activations_dtype: 'bfloat16'
34+
35+
# Maximum sequence length for the text encoder
36+
max_sequence_length: 512
37+
# offloads text encoder after text encoding to save memory.
38+
offload_encoders: True
39+
40+
# Attention
41+
attention: 'flash' # Supported attention: dot_product, flash, tokamax_flash, cudnn_flash_te, ring, tokamax_ring, ulysses, ulysses_custom, ulysses_ring
42+
# Maxdiffusion has 2 types of attention sharding strategies:
43+
# 1. attention_sharding_uniform = True : same sequence sharding rules applied for q in both (self and cross attention)
44+
# 2. attention_sharding_uniform = False : Heads are sharded uniformly across devices for self attention while sequence is
45+
# sharded in cross attention q.
46+
attention_sharding_uniform: True
47+
# Empty dict lets the kernel pick defaults. Example v6e override:
48+
# flash_block_sizes: {
49+
# "block_q" : 1024,
50+
# "block_kv_compute" : 1024,
51+
# "block_kv" : 1024,
52+
# "block_q_dkv" : 1024,
53+
# "block_kv_dkv" : 1024,
54+
# "block_kv_dkv_compute" : 1024,
55+
# "block_q_dq" : 1024,
56+
# "block_kv_dq" : 1024,
57+
# "use_fused_bwd_kernel": False,
58+
# }
59+
flash_block_sizes: {}
60+
61+
# Output directory
62+
# Create a GCS bucket, e.g. my-maxdiffusion-outputs and set this to "gs://my-maxdiffusion-outputs/"
63+
base_output_directory: ""
64+
output_dir: "/mnt/disks/data/tmp/maxdiffusion"
65+
# Local path for the generated image. If empty, an image named
66+
# {run_name}_output_{seed}.png is written to the working directory.
67+
output_file: "/mnt/disks/data/tmp/zimage.png"
68+
69+
# Hardware
70+
hardware: 'tpu' # Supported hardware types are 'tpu', 'gpu'
71+
skip_jax_distributed_system: True
72+
73+
# Parallelism
74+
mesh_axes: ['data', 'fsdp', 'context', 'tensor']
75+
# Z-Image runs sequence parallel (activations sharded over `context`) with
76+
# data parallelism; the tensor axis is left at 1. `attention_sharding_uniform`
77+
# adds the matching q/kv sequence rules on top of these.
78+
logical_axis_rules: [
79+
['batch', ['data', 'fsdp']],
80+
['activation_batch', ['data', 'fsdp']],
81+
['activation_length', 'context'],
82+
['activation_heads', 'tensor'],
83+
['heads', 'tensor'],
84+
['mlp', 'tensor'],
85+
['embed', 'fsdp'],
86+
['norm', 'tensor'],
87+
['out_channels', 'tensor'],
88+
]
89+
data_sharding: [['data', 'fsdp', 'context', 'tensor']]
90+
91+
# One axis for each parallelism type may hold a placeholder (-1)
92+
# value to auto-shard based on available slices and devices.
93+
#
94+
# Measured on v6e-8, Turbo, 1024x1024, 9 steps, 8 images (per_device_batch_size 1.0):
95+
# data=8 2.45s 0.307s/image <- default
96+
# data=4 ctx=2 2.60s 0.325s/image
97+
# data=2 ctx=4 3.13s 0.391s/image
98+
# ctx=8 4.15s 0.519s/image
99+
# fsdp=8 9.43s 1.179s/image
100+
# The 12B of weights fit in HBM, so replicating them and giving each chip a
101+
# whole image beats every sharded layout: no collectives at all. fsdp re-gathers
102+
# every weight on all 9 denoising steps with no backward pass to amortize it.
103+
#
104+
# For single-image *latency* instead of throughput, generate one image
105+
# (per_device_batch_size: 0.125) with ici_data_parallelism: 1 and
106+
# ici_context_parallelism: -1, which splits one image's sequence 8 ways: 0.587s.
107+
dcn_data_parallelism: 1
108+
dcn_fsdp_parallelism: 1
109+
dcn_context_parallelism: 1
110+
dcn_tensor_parallelism: 1
111+
ici_data_parallelism: -1 # recommended ICI axis to be auto-sharded
112+
ici_fsdp_parallelism: 1
113+
ici_context_parallelism: 1
114+
ici_tensor_parallelism: 1
115+
116+
allow_split_physical_axes: False
117+
118+
# Generation parameters
119+
prompt: "A red fox reading a book in a quiet library, cinematic light"
120+
height: 1024
121+
width: 1024
122+
num_inference_steps: 50
123+
# The released Z-Image configuration runs without classifier free guidance.
124+
guidance_scale: 0.0
125+
seed: 42
126+
# Images per device. Fractional values are allowed: on 8 devices,
127+
# 0.125 generates a single image.
128+
per_device_batch_size: 1.0
129+
# Latents decoded per VAE call. The decoder is replicated and its
130+
# activations are full-resolution, so raising this trades memory for speed.
131+
vae_decode_chunk: 1
132+
133+
# Profiling
134+
# generate_zimage runs an un-profiled warmup pass (compile), then a clean
135+
# generation pass, and only then a profiled pass of profiler_steps steps.
136+
enable_profiler: False
137+
profiler_steps: 5
138+
139+
# ML Diagnostics settings
140+
enable_ml_diagnostics: False
141+
profiler_gcs_path: ""
142+
enable_ondemand_xprof: False
143+
144+
# Directory for the Orbax cache of the converted pipeline (transformer + VAE +
145+
# text encoder). '' disables it. The first run loads from Diffusers and writes
146+
# the cache; later runs restore from it and skip the torch -> flax conversion.
147+
pretrained_orbax_dir: ""
148+
149+
# Keys below are required by pyconfig but unused for Z-Image inference.
150+
jax_cache_dir: ""
151+
dataset_name: ""
152+
dataset_save_location: "/mnt/disks/data/tmp"
153+
learning_rate_schedule_steps: 1
154+
max_train_steps: 1
155+
compile_topology_num_slices: -1
156+
quantization_local_shard_count: -1
157+
use_qwix_quantization: False
Lines changed: 156 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,156 @@
1+
# Copyright 2026 Google LLC
2+
#
3+
# Licensed under the Apache License, Version 2.0 (the "License");
4+
# you may not use this file except in compliance with the License.
5+
# You may obtain a copy of the License at
6+
#
7+
# https://www.apache.org/licenses/LICENSE-2.0
8+
#
9+
# Unless required by applicable law or agreed to in writing, software
10+
# distributed under the License is distributed on an "AS IS" BASIS,
11+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12+
# See the License for the specific language governing permissions and
13+
# limitations under the License.
14+
15+
# Official Z-Image-Turbo inference defaults. The base model uses the same
16+
# transformer; only its recommended denoising-step count differs.
17+
18+
run_name: 'zimage-turbo'
19+
20+
# If true save config to GCS in {base_output_directory}/{run_name}/
21+
save_config_to_gcs: False
22+
23+
pretrained_model_name_or_path: 'Tongyi-MAI/Z-Image-Turbo'
24+
model_name: z_image
25+
model_type: 'T2I'
26+
27+
unet_checkpoint: ''
28+
# This will convert the weights to this dtype.
29+
# When running inference on TPUv5e, use weights_dtype: 'bfloat16'
30+
weights_dtype: 'bfloat16'
31+
# This sets the layer's dtype in the model. Ex: nn.Dense(dtype=activations_dtype)
32+
activations_dtype: 'bfloat16'
33+
34+
# Maximum sequence length for the text encoder
35+
max_sequence_length: 512
36+
# offloads text encoder after text encoding to save memory.
37+
offload_encoders: True
38+
39+
# Attention
40+
attention: 'flash' # Supported attention: dot_product, flash, tokamax_flash, cudnn_flash_te, ring, tokamax_ring, ulysses, ulysses_custom, ulysses_ring
41+
# Maxdiffusion has 2 types of attention sharding strategies:
42+
# 1. attention_sharding_uniform = True : same sequence sharding rules applied for q in both (self and cross attention)
43+
# 2. attention_sharding_uniform = False : Heads are sharded uniformly across devices for self attention while sequence is
44+
# sharded in cross attention q.
45+
attention_sharding_uniform: True
46+
# Empty dict lets the kernel pick defaults. Example v6e override:
47+
# flash_block_sizes: {
48+
# "block_q" : 1024,
49+
# "block_kv_compute" : 1024,
50+
# "block_kv" : 1024,
51+
# "block_q_dkv" : 1024,
52+
# "block_kv_dkv" : 1024,
53+
# "block_kv_dkv_compute" : 1024,
54+
# "block_q_dq" : 1024,
55+
# "block_kv_dq" : 1024,
56+
# "use_fused_bwd_kernel": False,
57+
# }
58+
flash_block_sizes: {}
59+
60+
# Output directory
61+
# Create a GCS bucket, e.g. my-maxdiffusion-outputs and set this to "gs://my-maxdiffusion-outputs/"
62+
base_output_directory: ""
63+
output_dir: "/mnt/disks/data/tmp/maxdiffusion"
64+
# Local path for the generated image. If empty, an image named
65+
# {run_name}_output_{seed}.png is written to the working directory.
66+
output_file: "/mnt/disks/data/tmp/zimage-turbo.png"
67+
68+
# Hardware
69+
hardware: 'tpu' # Supported hardware types are 'tpu', 'gpu'
70+
skip_jax_distributed_system: True
71+
72+
# Parallelism
73+
mesh_axes: ['data', 'fsdp', 'context', 'tensor']
74+
# Z-Image runs sequence parallel (activations sharded over `context`) with
75+
# data parallelism; the tensor axis is left at 1. `attention_sharding_uniform`
76+
# adds the matching q/kv sequence rules on top of these.
77+
logical_axis_rules: [
78+
['batch', ['data', 'fsdp']],
79+
['activation_batch', ['data', 'fsdp']],
80+
['activation_length', 'context'],
81+
['activation_heads', 'tensor'],
82+
['heads', 'tensor'],
83+
['mlp', 'tensor'],
84+
['embed', 'fsdp'],
85+
['norm', 'tensor'],
86+
['out_channels', 'tensor'],
87+
]
88+
data_sharding: [['data', 'fsdp', 'context', 'tensor']]
89+
90+
# One axis for each parallelism type may hold a placeholder (-1)
91+
# value to auto-shard based on available slices and devices.
92+
#
93+
# Measured on v6e-8, Turbo, 1024x1024, 9 steps, 8 images (per_device_batch_size 1.0):
94+
# data=8 2.45s 0.307s/image <- default
95+
# data=4 ctx=2 2.60s 0.325s/image
96+
# data=2 ctx=4 3.13s 0.391s/image
97+
# ctx=8 4.15s 0.519s/image
98+
# fsdp=8 9.43s 1.179s/image
99+
# The 12B of weights fit in HBM, so replicating them and giving each chip a
100+
# whole image beats every sharded layout: no collectives at all. fsdp re-gathers
101+
# every weight on all 9 denoising steps with no backward pass to amortize it.
102+
#
103+
# For single-image *latency* instead of throughput, generate one image
104+
# (per_device_batch_size: 0.125) with ici_data_parallelism: 1 and
105+
# ici_context_parallelism: -1, which splits one image's sequence 8 ways: 0.587s.
106+
dcn_data_parallelism: 1
107+
dcn_fsdp_parallelism: 1
108+
dcn_context_parallelism: 1
109+
dcn_tensor_parallelism: 1
110+
ici_data_parallelism: -1 # recommended ICI axis to be auto-sharded
111+
ici_fsdp_parallelism: 1
112+
ici_context_parallelism: 1
113+
ici_tensor_parallelism: 1
114+
115+
allow_split_physical_axes: False
116+
117+
# Generation parameters
118+
prompt: "A red fox reading a book in a quiet library, cinematic light"
119+
height: 1024
120+
width: 1024
121+
num_inference_steps: 9
122+
# The published Turbo configuration uses guidance_scale=0.
123+
guidance_scale: 0.0
124+
seed: 42
125+
# Images per device. Fractional values are allowed: on 8 devices,
126+
# 0.125 generates a single image.
127+
per_device_batch_size: 1.0
128+
# Latents decoded per VAE call. The decoder is replicated and its
129+
# activations are full-resolution, so raising this trades memory for speed.
130+
vae_decode_chunk: 1
131+
132+
# Profiling
133+
# generate_zimage runs an un-profiled warmup pass (compile), then a clean
134+
# generation pass, and only then a profiled pass of profiler_steps steps.
135+
enable_profiler: False
136+
profiler_steps: 5
137+
138+
# ML Diagnostics settings
139+
enable_ml_diagnostics: False
140+
profiler_gcs_path: ""
141+
enable_ondemand_xprof: False
142+
143+
# Directory for the Orbax cache of the converted pipeline (transformer + VAE +
144+
# text encoder). '' disables it. The first run loads from Diffusers and writes
145+
# the cache; later runs restore from it and skip the torch -> flax conversion.
146+
pretrained_orbax_dir: ""
147+
148+
# Keys below are required by pyconfig but unused for Z-Image inference.
149+
jax_cache_dir: ""
150+
dataset_name: ""
151+
dataset_save_location: "/mnt/disks/data/tmp"
152+
learning_rate_schedule_steps: 1
153+
max_train_steps: 1
154+
compile_topology_num_slices: -1
155+
quantization_local_shard_count: -1
156+
use_qwix_quantization: False

src/maxdiffusion/generate_flux2klein.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -37,7 +37,8 @@
3737

3838
from maxdiffusion.models.flux.transformers.transformer_flux_flax import Flux2KleinTransformer2DModel
3939
from maxdiffusion.models.vae_flax import FlaxAutoencoderKL
40-
from maxdiffusion.models.qwen3_flax import FlaxQwen3Config, FlaxQwen3Model, load_and_convert_qwen3_weights
40+
from maxdiffusion.models.qwen3_flax import FlaxQwen3Config, FlaxQwen3Model
41+
from maxdiffusion.models.qwen3_utils import load_and_convert_qwen3_weights
4142
from maxdiffusion.schedulers.scheduling_flow_match_flax import FlaxFlowMatchScheduler
4243

4344

0 commit comments

Comments
 (0)