|
| 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 | + |
| 38 | +# Attention |
| 39 | +attention: 'flash' # Supported attention: dot_product, flash, tokamax_flash, cudnn_flash_te, ring, tokamax_ring, ulysses, ulysses_custom, ulysses_ring |
| 40 | +# Maxdiffusion has 2 types of attention sharding strategies: |
| 41 | +# 1. attention_sharding_uniform = True : same sequence sharding rules applied for q in both (self and cross attention) |
| 42 | +# 2. attention_sharding_uniform = False : Heads are sharded uniformly across devices for self attention while sequence is |
| 43 | +# sharded in cross attention q. |
| 44 | +attention_sharding_uniform: True |
| 45 | +# Empty dict lets the kernel pick defaults. Example v6e override: |
| 46 | +# flash_block_sizes: { |
| 47 | +# "block_q" : 1024, |
| 48 | +# "block_kv_compute" : 1024, |
| 49 | +# "block_kv" : 1024, |
| 50 | +# "block_q_dkv" : 1024, |
| 51 | +# "block_kv_dkv" : 1024, |
| 52 | +# "block_kv_dkv_compute" : 1024, |
| 53 | +# "block_q_dq" : 1024, |
| 54 | +# "block_kv_dq" : 1024, |
| 55 | +# "use_fused_bwd_kernel": False, |
| 56 | +# } |
| 57 | +flash_block_sizes: {} |
| 58 | + |
| 59 | +# Output directory |
| 60 | +# Create a GCS bucket, e.g. my-maxdiffusion-outputs and set this to "gs://my-maxdiffusion-outputs/" |
| 61 | +base_output_directory: "" |
| 62 | +output_dir: "/mnt/disks/data/tmp/maxdiffusion" |
| 63 | +# Local path for the generated image. If empty, an image named |
| 64 | +# {run_name}_output_{seed}.png is written to the working directory. |
| 65 | +output_file: "/mnt/disks/data/tmp/zimage.png" |
| 66 | + |
| 67 | +# Hardware |
| 68 | +hardware: 'tpu' # Supported hardware types are 'tpu', 'gpu' |
| 69 | +skip_jax_distributed_system: True |
| 70 | + |
| 71 | +# Parallelism |
| 72 | +mesh_axes: ['data', 'fsdp', 'context', 'tensor'] |
| 73 | +# Z-Image runs sequence parallel (activations sharded over `context`) with |
| 74 | +# data parallelism; the tensor axis is left at 1. `attention_sharding_uniform` |
| 75 | +# adds the matching q/kv sequence rules on top of these. |
| 76 | +logical_axis_rules: [ |
| 77 | + ['batch', ['data', 'fsdp']], |
| 78 | + ['activation_batch', ['data', 'fsdp']], |
| 79 | + ['activation_length', 'context'], |
| 80 | + ['activation_heads', 'tensor'], |
| 81 | + ['heads', 'tensor'], |
| 82 | + ['mlp', 'tensor'], |
| 83 | + ['embed', 'fsdp'], |
| 84 | + ['norm', 'tensor'], |
| 85 | + ['out_channels', 'tensor'], |
| 86 | + ] |
| 87 | +data_sharding: [['data', 'fsdp', 'context', 'tensor']] |
| 88 | + |
| 89 | +# One axis for each parallelism type may hold a placeholder (-1) |
| 90 | +# value to auto-shard based on available slices and devices. |
| 91 | +# |
| 92 | +# Measured on v6e-8, Turbo, 1024x1024, 9 steps, 8 images (per_device_batch_size 1.0): |
| 93 | +# data=8 2.45s 0.307s/image <- default |
| 94 | +# data=4 ctx=2 2.60s 0.325s/image |
| 95 | +# data=2 ctx=4 3.13s 0.391s/image |
| 96 | +# ctx=8 4.15s 0.519s/image |
| 97 | +# fsdp=8 9.43s 1.179s/image |
| 98 | +# The 12B of weights fit in HBM, so replicating them and giving each chip a |
| 99 | +# whole image beats every sharded layout: no collectives at all. fsdp re-gathers |
| 100 | +# every weight on all 9 denoising steps with no backward pass to amortize it. |
| 101 | +# |
| 102 | +# For single-image *latency* instead of throughput, generate one image |
| 103 | +# (per_device_batch_size: 0.125) with ici_data_parallelism: 1 and |
| 104 | +# ici_context_parallelism: -1, which splits one image's sequence 8 ways: 0.587s. |
| 105 | +dcn_data_parallelism: 1 |
| 106 | +dcn_fsdp_parallelism: 1 |
| 107 | +dcn_context_parallelism: 1 |
| 108 | +dcn_tensor_parallelism: 1 |
| 109 | +ici_data_parallelism: -1 # recommended ICI axis to be auto-sharded |
| 110 | +ici_fsdp_parallelism: 1 |
| 111 | +ici_context_parallelism: 1 |
| 112 | +ici_tensor_parallelism: 1 |
| 113 | + |
| 114 | +allow_split_physical_axes: False |
| 115 | + |
| 116 | +# Generation parameters |
| 117 | +prompt: "A red fox reading a book in a quiet library, cinematic light" |
| 118 | +height: 1024 |
| 119 | +width: 1024 |
| 120 | +num_inference_steps: 50 |
| 121 | +# The released Z-Image configuration runs without classifier free guidance. |
| 122 | +guidance_scale: 0.0 |
| 123 | +seed: 42 |
| 124 | +# Images per device. Fractional values are allowed: on 8 devices, |
| 125 | +# 0.125 generates a single image. |
| 126 | +per_device_batch_size: 1.0 |
| 127 | +# Latents decoded per VAE call. The decoder is replicated and its |
| 128 | +# activations are full-resolution, so raising this trades memory for speed. |
| 129 | +vae_decode_chunk: 1 |
| 130 | + |
| 131 | +# Profiling |
| 132 | +# generate_zimage runs an un-profiled warmup pass (compile), then a clean |
| 133 | +# generation pass, and only then a profiled pass of profiler_steps steps. |
| 134 | +enable_profiler: False |
| 135 | +profiler_steps: 5 |
| 136 | + |
| 137 | +# ML Diagnostics settings |
| 138 | +enable_ml_diagnostics: False |
| 139 | +profiler_gcs_path: "" |
| 140 | +enable_ondemand_xprof: False |
| 141 | + |
| 142 | +# Directory for the Orbax cache of the converted pipeline (transformer + VAE + |
| 143 | +# text encoder). '' disables it. The first run loads from Diffusers and writes |
| 144 | +# the cache; later runs restore from it and skip the torch -> flax conversion. |
| 145 | +pretrained_orbax_dir: "" |
| 146 | + |
| 147 | +# Keys below are required by pyconfig but unused for Z-Image inference. |
| 148 | +jax_cache_dir: "" |
| 149 | +dataset_name: "" |
| 150 | +dataset_save_location: "/mnt/disks/data/tmp" |
| 151 | +learning_rate_schedule_steps: 1 |
| 152 | +max_train_steps: 1 |
| 153 | +compile_topology_num_slices: -1 |
| 154 | +quantization_local_shard_count: -1 |
| 155 | +use_qwix_quantization: False |
0 commit comments