|
1 | | -# Official Z-Image inference defaults. Z-Image and Turbo share the denoiser. |
2 | | -pretrained_model_name_or_path: "Tongyi-MAI/Z-Image" |
| 1 | +# Copyright 2025 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 |
3 | 117 | prompt: "A red fox reading a book in a quiet library, cinematic light" |
4 | | -resolution: 1024 |
| 118 | +height: 1024 |
| 119 | +width: 1024 |
5 | 120 | num_inference_steps: 50 |
| 121 | +# The released Z-Image configuration runs without classifier free guidance. |
| 122 | +guidance_scale: 0.0 |
6 | 123 | seed: 42 |
7 | | -output_file: "/tmp/zimage.png" |
8 | | -weights_dtype: "bfloat16" |
9 | | -activations_dtype: "bfloat16" |
10 | | -attention: "flash" |
11 | | -flash_block_sizes: {} |
12 | | -hardware: "tpu" |
13 | | -skip_jax_distributed_system: true |
14 | | -run_name: "zimage" |
15 | | -output_dir: "/tmp/maxdiffusion" |
16 | | -save_config_to_gcs: false |
| 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. |
17 | 148 | jax_cache_dir: "" |
18 | | -unet_checkpoint: "" |
19 | 149 | dataset_name: "" |
20 | | -dataset_save_location: "/tmp" |
21 | | -per_device_batch_size: 1 |
| 150 | +dataset_save_location: "/mnt/disks/data/tmp" |
22 | 151 | learning_rate_schedule_steps: 1 |
23 | 152 | max_train_steps: 1 |
24 | 153 | compile_topology_num_slices: -1 |
25 | 154 | quantization_local_shard_count: -1 |
26 | | -use_qwix_quantization: false |
27 | | -attention_sharding_uniform: true |
28 | | -logical_axis_rules: [["batch", "data"], ["activation_length", "context"], ["activation_heads", "tensor"], ["heads", "tensor"], ["mlp", "tensor"], ["embed", "fsdp"], ["norm", "tensor"], ["out_channels", "tensor"]] |
29 | | -data_sharding: [["data", "fsdp", "context", "tensor"]] |
30 | | -mesh_axes: ["data", "fsdp", "context", "tensor"] |
31 | | -dcn_data_parallelism: 1 |
32 | | -dcn_fsdp_parallelism: 1 |
33 | | -dcn_context_parallelism: 1 |
34 | | -dcn_tensor_parallelism: 1 |
35 | | -ici_data_parallelism: 1 |
36 | | -ici_fsdp_parallelism: 1 |
37 | | -ici_context_parallelism: 1 |
38 | | -ici_tensor_parallelism: -1 |
39 | | -allow_split_physical_axes: false |
| 155 | +use_qwix_quantization: False |
0 commit comments