Skip to content

Commit 7d5e1d0

Browse files
committed
Optimize Z-Image-Turbo serving pipeline, attention layouts, and add sharding specs
- 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. - Add logical sharding specs and registry for Z-Image transformer.
1 parent bda41e1 commit 7d5e1d0

23 files changed

Lines changed: 2831 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: 421 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: 160 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,160 @@
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: False
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+
sharding:
92+
transformer: 'default'
93+
94+
# One axis for each parallelism type may hold a placeholder (-1)
95+
# value to auto-shard based on available slices and devices.
96+
#
97+
# Measured on v6e-8, Turbo, 1024x1024, 9 steps, 8 images (per_device_batch_size 1.0):
98+
# data=8 2.45s 0.307s/image <- default
99+
# data=4 ctx=2 2.60s 0.325s/image
100+
# data=2 ctx=4 3.13s 0.391s/image
101+
# ctx=8 4.15s 0.519s/image
102+
# fsdp=8 9.43s 1.179s/image
103+
# The 12B of weights fit in HBM, so replicating them and giving each chip a
104+
# whole image beats every sharded layout: no collectives at all. fsdp re-gathers
105+
# every weight on all 9 denoising steps with no backward pass to amortize it.
106+
#
107+
# For single-image *latency* instead of throughput, generate one image
108+
# (per_device_batch_size: 0.125) with ici_data_parallelism: 1 and
109+
# ici_context_parallelism: -1, which splits one image's sequence 8 ways: 0.587s.
110+
dcn_data_parallelism: 1
111+
dcn_fsdp_parallelism: 1
112+
dcn_context_parallelism: 1
113+
dcn_tensor_parallelism: 1
114+
ici_data_parallelism: -1 # recommended ICI axis to be auto-sharded
115+
ici_fsdp_parallelism: 1
116+
ici_context_parallelism: 1
117+
ici_tensor_parallelism: 1
118+
119+
allow_split_physical_axes: False
120+
121+
# Generation parameters
122+
prompt: "A red fox reading a book in a quiet library, cinematic light"
123+
height: 1024
124+
width: 1024
125+
num_inference_steps: 50
126+
# The released Z-Image configuration runs without classifier free guidance.
127+
guidance_scale: 0.0
128+
seed: 42
129+
# Images per device. Fractional values are allowed: on 8 devices,
130+
# 0.125 generates a single image.
131+
per_device_batch_size: 1.0
132+
# Latents decoded per VAE call. The decoder is replicated and its
133+
# activations are full-resolution, so raising this trades memory for speed.
134+
vae_decode_chunk: 1
135+
136+
# Profiling
137+
# generate_zimage runs an un-profiled warmup pass (compile), then a clean
138+
# generation pass, and only then a profiled pass of profiler_steps steps.
139+
enable_profiler: False
140+
profiler_steps: 5
141+
142+
# ML Diagnostics settings
143+
enable_ml_diagnostics: False
144+
profiler_gcs_path: ""
145+
enable_ondemand_xprof: False
146+
147+
# Directory for the Orbax cache of the converted pipeline (transformer + VAE +
148+
# text encoder). '' disables it. The first run loads from Diffusers and writes
149+
# the cache; later runs restore from it and skip the torch -> flax conversion.
150+
pretrained_orbax_dir: ""
151+
152+
# Keys below are required by pyconfig but unused for Z-Image inference.
153+
jax_cache_dir: ""
154+
dataset_name: ""
155+
dataset_save_location: "/mnt/disks/data/tmp"
156+
learning_rate_schedule_steps: 1
157+
max_train_steps: 1
158+
compile_topology_num_slices: -1
159+
quantization_local_shard_count: -1
160+
use_qwix_quantization: False
Lines changed: 159 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,159 @@
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: False
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+
sharding:
91+
transformer: 'default'
92+
93+
# One axis for each parallelism type may hold a placeholder (-1)
94+
# value to auto-shard based on available slices and devices.
95+
#
96+
# Measured on v6e-8, Turbo, 1024x1024, 9 steps, 8 images (per_device_batch_size 1.0):
97+
# data=8 2.45s 0.307s/image <- default
98+
# data=4 ctx=2 2.60s 0.325s/image
99+
# data=2 ctx=4 3.13s 0.391s/image
100+
# ctx=8 4.15s 0.519s/image
101+
# fsdp=8 9.43s 1.179s/image
102+
# The 12B of weights fit in HBM, so replicating them and giving each chip a
103+
# whole image beats every sharded layout: no collectives at all. fsdp re-gathers
104+
# every weight on all 9 denoising steps with no backward pass to amortize it.
105+
#
106+
# For single-image *latency* instead of throughput, generate one image
107+
# (per_device_batch_size: 0.125) with ici_data_parallelism: 1 and
108+
# ici_context_parallelism: -1, which splits one image's sequence 8 ways: 0.587s.
109+
dcn_data_parallelism: 1
110+
dcn_fsdp_parallelism: 1
111+
dcn_context_parallelism: 1
112+
dcn_tensor_parallelism: 1
113+
ici_data_parallelism: -1 # recommended ICI axis to be auto-sharded
114+
ici_fsdp_parallelism: 1
115+
ici_context_parallelism: 1
116+
ici_tensor_parallelism: 1
117+
118+
allow_split_physical_axes: False
119+
120+
# Generation parameters
121+
prompt: "A red fox reading a book in a quiet library, cinematic light"
122+
height: 1024
123+
width: 1024
124+
num_inference_steps: 9
125+
# The published Turbo configuration uses guidance_scale=0.
126+
guidance_scale: 0.0
127+
seed: 42
128+
# Images per device. Fractional values are allowed: on 8 devices,
129+
# 0.125 generates a single image.
130+
per_device_batch_size: 1.0
131+
# Latents decoded per VAE call. The decoder is replicated and its
132+
# activations are full-resolution, so raising this trades memory for speed.
133+
vae_decode_chunk: 1
134+
135+
# Profiling
136+
# generate_zimage runs an un-profiled warmup pass (compile), then a clean
137+
# generation pass, and only then a profiled pass of profiler_steps steps.
138+
enable_profiler: False
139+
profiler_steps: 5
140+
141+
# ML Diagnostics settings
142+
enable_ml_diagnostics: False
143+
profiler_gcs_path: ""
144+
enable_ondemand_xprof: False
145+
146+
# Directory for the Orbax cache of the converted pipeline (transformer + VAE +
147+
# text encoder). '' disables it. The first run loads from Diffusers and writes
148+
# the cache; later runs restore from it and skip the torch -> flax conversion.
149+
pretrained_orbax_dir: ""
150+
151+
# Keys below are required by pyconfig but unused for Z-Image inference.
152+
jax_cache_dir: ""
153+
dataset_name: ""
154+
dataset_save_location: "/mnt/disks/data/tmp"
155+
learning_rate_schedule_steps: 1
156+
max_train_steps: 1
157+
compile_topology_num_slices: -1
158+
quantization_local_shard_count: -1
159+
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)