1 of 22

​

​

TPU-friendly port of QwenImage with PyTorch/XLA

Sayak Paul, Hugging Face 🤗�

Senior Research Engineer

​

​

2 of 22

Agenda

  • Quick intro to image generation
  • Diffusers and PyTorch/XLA
  • The QwenImage pipeline
  • Running on TPU v6e-8 and SPMD
  • Results

3 of 22

Assuming a diffusion-based latent-space image generation model

4 of 22

Diffusion models ≠ single models

Text encoders

“A cat holding a sign that says hello world”

5 of 22

Diffusion models ≠ single models

Text encoders

Noisy latents

Embeddings

“A cat holding a sign that says hello world”

6 of 22

Diffusion models ≠ single models

Text encoders

Noisy latents

Embeddings

“A cat holding a sign that says hello world”

Scheduler

Timestep

Diffusion UNet / Transformer

7 of 22

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

8 of 22

Diffusers 🧨

  • Open-source Python library for using and customizing open & SOTA {image,video,audio,text} generation models
  • Prioritizes performance and accessibility
  • PyTorch/XLA-compatible

9 of 22

Diffusers 🧨

10 of 22

PyTorch/XLA into Diffusers

xm.mark_step() is used when XLA device is available

11 of 22

PyTorch/XLA into Diffusers

torch_xla.experimental.custom_kernel.flash_attention for optimized attention computation

12 of 22

The QwenImage pipeline

Prompt

(Qwen2.5-VL)

Noisy latents + embeddings (Transformer)

Refined latents

(VAE)

13 of 22

The QwenImage pipeline

Prompt

(Qwen2.5-VL)

Noisy latents + embeddings (Transformer)

Refined latents

(VAE)

CPU

XLA

XLA

14 of 22

QwenImage DiT

15 of 22

QwenImage DiT

Doesn’t fully fit in a TPU v6e-8 — needs to be sharded:

  • Handle memory
  • Handle XLA compilation graph cache
  • Mitigating XLA violations

16 of 22

Sharding the DiT

A 2D mesh distributes the DiT parameters along the largest parameter dimension.

  • Avoids full transformer replication
  • Keeps Diffusers components largely intact
  • Uses single-process SPMD instead of xmp.spawn
  • Lets the compiler optimize a stable graph

17 of 22

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.

18 of 22

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.

19 of 22

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

20 of 22

Profiling revealed …

  • H2D transfers are small
  • Most graphs are dense

​

TODO:

  • Less H2D syncs
  • Less recompiles
  • Better sharding choices

9

21 of 22

Exciting days ahead!

9

22 of 22

Slides

Code