TPU-friendly port of QwenImage with PyTorch/XLA
Sayak Paul, Hugging Face 🤗�
Senior Research Engineer
Agenda
Assuming a diffusion-based latent-space image generation model
Diffusion models ≠ single models
Text encoders
“A cat holding a sign that says hello world”
Diffusion models ≠ single models
Text encoders
Noisy latents
Embeddings
“A cat holding a sign that says hello world”
Diffusion models ≠ single models
Text encoders
Noisy latents
Embeddings
“A cat holding a sign that says hello world”
Scheduler
Timestep
Diffusion UNet / Transformer
Diffusion models ≠ single models
Text encoders
Refined latents
Decoder
“A cat holding a sign that says hello world”
Noisy latents
Embeddings
Scheduler
Timestep
Diffusion UNet / Transformer
Diffusers 🧨
Diffusers 🧨
PyTorch/XLA into Diffusers
xm.mark_step() is used when XLA device is available
PyTorch/XLA into Diffusers
torch_xla.experimental.custom_kernel.flash_attention for optimized attention computation
The QwenImage pipeline
Prompt
(Qwen2.5-VL)
Noisy latents + embeddings (Transformer)
Refined latents
(VAE)
The QwenImage pipeline
Prompt
(Qwen2.5-VL)
Noisy latents + embeddings (Transformer)
Refined latents
(VAE)
CPU
XLA
XLA
QwenImage DiT
QwenImage DiT
Doesn’t fully fit in a TPU v6e-8 — needs to be sharded:
Sharding the DiT
A 2D mesh distributes the DiT parameters along the largest parameter dimension.
Sharding the DiT
xs.Mesh(np.arange(8), (2, 4), ("data", "model"))
model axis →
d�a�t�a
↓
TPU0
TPU1
TPU2
TPU3
TPU4
TPU5
TPU6
TPU7
The "model" axis has size 4, so each large DiT weight is split into 4 shards.
2D mesh
Weight example
W.shape = [4096, 16384]
largest dim = 16384 → shard dim = 1
spec = [None, "model"]
m0 / TPU0,4 → 0:4096
m1 / TPU1,5 → 4096:8192
m2 / TPU2,6 → 8192:12288
m3 / TPU3,7 → 12288:16384
Equivalent reading: W is split column-wise into 4 chunks.
Other gotchas
XLA rejects x[:, :, -2:, :, :] when temporal dimension is 1; clamp the slice length.
encode_prompt can drop an all-true prompt_embeds_mask; recreate it to avoid silently disabling true CFG.
Benchmark snapshot on TPU v6e-8
First run compilation
~95 min
Steady-state inference
~20 sec
Both numbers are for 50 denoising steps. Compiled graphs are cached at /tmp/data/compiler_cache_tRiLlium_eXp and reused across runs.
9
Profiling revealed …
TODO:
9
Exciting days ahead!
9
Slides
Code