Models / sdxl-base-unet

SDXL Base UNet

The 2.6B-parameter denoising UNet of Stable Diffusion XL base: residual blocks and spatial transformers at three resolutions, cross-attending to a 2048-wide text embedding, conditioned on the timestep and on SDXL's pooled-text-plus-micro-conditioning vector.

2.6B parameterssdxlCreativeML Open RAIL++-Mimage-generationvisionconvolutionaldiffusiontext-to-imageunet

NVIDIA H100 80GB HBM3 · median of 10 runs after 3 warm-ups · measured 2026-09-28T20:26:01+00:00

Linnet 0.1.0, PyTorch 2.14.0, JAX 0.11.2, transformers 5.17.0, diffusers 0.40.0, ONNX Runtime 1.30.0, Python 3.12.3, driver 580.126.09

Linnet against the stack it replaces

Each side in its fastest configuration, on the same GPU and checkpoint. A speed-up is how many times the other's speed; an export is measured against the original checkpoint in the same engine.

vs the reference, in PyTorch1.01× fasterdenoising step
  • Linnet, CUDA graphs34.7 msdiffusers, compiled35.0 ms

Each row also carries its distance from reference: the largest absolute difference between its output and diffusers (eager)'s on the same input. bf16 outputs differ by rounding (about 0.1 on logits near 16), so a small number is expected; it is there so that a fast wrong answer cannot look like a win.

Denoising step

ms, lower is better; the best bar is red.

PyTorchJAXONNX Runtime and TensorRT050100msdiffusers48.8 msdiffusers, compiled35.0 msLinnet, generated PyTorch56.3 msLinnet, inductor35.7 msLinnet, CUDA graphs34.7 msLinnet, XLA (StableHLO)53.7 msLinnet, XLA (generated JAX)39.3 msLinnet, ONNX Runtime f32141 msLinnet, ONNX Runtime f1690.8 msLinnet, ONNX Runtime bf16100 msLinnet, TensorRT f32141 msLinnet, TensorRT f1691.4 msLinnet, TensorRT bf16100 ms

Load time

s, lower is better; the best bar is red.

PyTorchJAXONNX Runtime and TensorRT0102030sdiffusers31.1 sdiffusers, compiled8.13 sLinnet, generated PyTorch17.0 sLinnet, inductor17.2 sLinnet, CUDA graphs19.8 sLinnet, XLA (StableHLO)22.2 sLinnet, XLA (generated JAX)25.0 sLinnet, ONNX Runtime f320.30 sLinnet, ONNX Runtime f160.29 sLinnet, ONNX Runtime bf160.29 sLinnet, TensorRT f320.30 sLinnet, TensorRT f160.29 sLinnet, TensorRT bf160.29 s

Peak GPU memory

What the driver reports the process holding at its peak, in GiB; lower is better; the best bar is red.

PyTorchJAXONNX Runtime and TensorRT01020GiBdiffusers6.29 GiBdiffusers, compiled6.30 GiBLinnet, generated PyTorch6.50 GiBLinnet, inductor6.28 GiBLinnet, CUDA graphs6.42 GiBLinnet, XLA (StableHLO)6.31 GiBLinnet, XLA (generated JAX)5.24 GiBLinnet, ONNX Runtime f3224.3 GiBLinnet, ONNX Runtime f1614.1 GiBLinnet, ONNX Runtime bf1614.6 GiBLinnet, TensorRT f3224.3 GiBLinnet, TensorRT f1614.1 GiBLinnet, TensorRT bf1614.6 GiB

Every number

MethodDenoising stepLoad timePeak GPU memoryAgainst the stack it replacesDistance from referenceNotes
PyTorch
diffusers48.8 ms31.1 s6.29 GiB0 (the reference)one step, 128x128 latent (1024-pixel image), batch 2, 77-token context, bf16
diffusers, compiled35.0 ms8.13 s6.30 GiB0.0313one step, 128x128 latent (1024-pixel image), batch 2, 77-token context, bf16
Linnet, generated PyTorch56.3 ms17.0 s6.50 GiB0.62× diffusers, compiled0.0156one step, 128x128 latent (1024-pixel image), batch 2, 77-token context, bf16
Linnet, inductor35.7 ms17.2 s6.28 GiB0.98× diffusers, compiled0.0313one step, 128x128 latent (1024-pixel image), batch 2, 77-token context, bf16
Linnet, CUDA graphs34.7 ms19.8 s6.42 GiB1.01× diffusers, compiled0.0313one step, 128x128 latent (1024-pixel image), batch 2, 77-token context, bf16
JAX
Linnet, XLA (StableHLO)53.7 ms22.2 s6.31 GiB0.0313one step, 128x128 latent (1024-pixel image), batch 2, 77-token context, bf16; peak memory is JAX's allocator peak; the driver shows its pool, which grows in whole regions
Linnet, XLA (generated JAX)39.3 ms25.0 s5.24 GiB0.0313one step, 128x128 latent (1024-pixel image), batch 2, 77-token context, bf16; peak memory is JAX's allocator peak; the driver shows its pool, which grows in whole regions
ONNX Runtime and TensorRT
Linnet, ONNX Runtime f32141 ms0.30 s24.3 GiB0.0303f32; ONNX Runtime's CUDA execution provider; first calls build the sessions
Linnet, ONNX Runtime f1690.8 ms0.29 s14.1 GiB0.0313f16; ONNX Runtime's CUDA execution provider; first calls build the sessions
Linnet, ONNX Runtime bf16100 ms0.29 s14.6 GiB0.0313bf16; ONNX Runtime's CUDA execution provider; first calls build the sessions
Linnet, TensorRT f32141 ms0.30 s24.3 GiB0.0303f32; ONNX Runtime's TensorRT execution provider; first calls build the sessions
Linnet, TensorRT f1691.4 ms0.29 s14.1 GiB0.0313f16; ONNX Runtime's TensorRT execution provider; first calls build the sessions
Linnet, TensorRT bf16100 ms0.29 s14.6 GiB0.0313bf16; ONNX Runtime's TensorRT execution provider; first calls build the sessions