Models / sd-vae-ft-mse

SD VAE ft-MSE (decoder)

Stability AI's MSE-finetuned autoencoder for Stable Diffusion, decoder half: group-normalized residual blocks, one spatial self-attention, and nearest-neighbour upsampling turn a 4-channel latent into an RGB image.

49.5M parameterssd-vaeMITimage-generationvisionconvolutionaldiffusionautoencoder

NVIDIA H100 80GB HBM3 · median of 10 runs after 3 warm-ups · measured 2026-09-28T20:13:23+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.00× fasterlatency
  • Linnet, CUDA graphs10.4 msdiffusers, compiled10.4 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.

Latency

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

PyTorchJAXONNX Runtime and TensorRT01020msdiffusers22.0 msdiffusers, compiled10.4 msLinnet, generated PyTorch27.5 msLinnet, inductor10.9 msLinnet, CUDA graphs10.4 msLinnet, XLA (StableHLO)7.47 msLinnet, XLA (generated JAX)7.41 msLinnet, ONNX Runtime f3224.7 msLinnet, ONNX Runtime f1621.0 msLinnet, ONNX Runtime bf1628.1 msLinnet, TensorRT f3213.7 msLinnet, TensorRT f165.41 msLinnet, TensorRT bf1612.5 ms

Throughput

/s, higher is better; the best bar is red.

PyTorchJAXONNX Runtime and TensorRT050100150200/sdiffusers66.2/sdiffusers, compiled110/sLinnet, generated PyTorch39.4/sLinnet, inductor111/sLinnet, CUDA graphs112/sLinnet, XLA (StableHLO)175/sLinnet, XLA (generated JAX)173/sLinnet, ONNX Runtime f3249.1/sLinnet, ONNX Runtime f1669.0/sLinnet, ONNX Runtime bf1642.1/sLinnet, TensorRT f3271.2/sLinnet, TensorRT f16199/sLinnet, TensorRT bf1678.1/s

Load time

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

PyTorchJAXONNX Runtime and TensorRT0246sdiffusers6.71 sdiffusers, compiled4.89 sLinnet, generated PyTorch0.87 sLinnet, inductor0.47 sLinnet, CUDA graphs0.45 sLinnet, XLA (StableHLO)0.66 sLinnet, XLA (generated JAX)0.50 sLinnet, ONNX Runtime f320.27 sLinnet, ONNX Runtime f160.24 sLinnet, ONNX Runtime bf160.24 sLinnet, TensorRT f320.24 sLinnet, TensorRT f160.24 sLinnet, TensorRT bf160.24 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 TensorRT01020GiBdiffusers8.95 GiBdiffusers, compiled7.40 GiBLinnet, generated PyTorch12.2 GiBLinnet, inductor4.82 GiBLinnet, CUDA graphs5.53 GiBLinnet, XLA (StableHLO)3.30 GiBLinnet, XLA (generated JAX)3.30 GiBLinnet, ONNX Runtime f3226.3 GiBLinnet, ONNX Runtime f1613.0 GiBLinnet, ONNX Runtime bf1621.2 GiBLinnet, TensorRT f3211.5 GiBLinnet, TensorRT f168.58 GiBLinnet, TensorRT bf1610.5 GiB

Every number

MethodLatencyThroughputLoad timePeak GPU memoryAgainst the stack it replacesDistance from referenceNotes
PyTorch
diffusers22.0 ms66.2/s6.71 s8.95 GiB0 (the reference)64x64 latent (512-pixel image), bf16; throughput at batch 8
diffusers, compiled10.4 ms110/s4.89 s7.40 GiB0.033264x64 latent (512-pixel image), bf16; throughput at batch 8
Linnet, generated PyTorch27.5 ms39.4/s0.87 s12.2 GiB0.38× diffusers, compiled0.031364x64 latent (512-pixel image), bf16; throughput at batch 8
Linnet, inductor10.9 ms111/s0.47 s4.82 GiB0.95× diffusers, compiled0.033264x64 latent (512-pixel image), bf16; throughput at batch 8
Linnet, CUDA graphs10.4 ms112/s0.45 s5.53 GiB1.00× diffusers, compiled0.033264x64 latent (512-pixel image), bf16; throughput at batch 8
JAX
Linnet, XLA (StableHLO)7.47 ms175/s0.66 s3.30 GiB0.044964x64 latent (512-pixel image), bf16; throughput at batch 8; peak memory is JAX's allocator peak; the driver shows its pool, which grows in whole regions
Linnet, XLA (generated JAX)7.41 ms173/s0.50 s3.30 GiB0.052764x64 latent (512-pixel image), bf16; throughput at batch 8; peak memory is JAX's allocator peak; the driver shows its pool, which grows in whole regions
ONNX Runtime and TensorRT
Linnet, ONNX Runtime f3224.7 ms49.1/s0.27 s26.3 GiB0.0397f32; ONNX Runtime's CUDA execution provider; first calls build the sessions
Linnet, ONNX Runtime f1621.0 ms69.0/s0.24 s13.0 GiB0.038f16; ONNX Runtime's CUDA execution provider; first calls build the sessions
Linnet, ONNX Runtime bf1628.1 ms42.1/s0.24 s21.2 GiB0.0566bf16; ONNX Runtime's CUDA execution provider; first calls build the sessions
Linnet, TensorRT f3213.7 ms71.2/s0.24 s11.5 GiB0.0398f32; ONNX Runtime's TensorRT execution provider; first calls build the sessions
Linnet, TensorRT f165.41 ms199/s0.24 s8.58 GiB0.0432f16; ONNX Runtime's TensorRT execution provider; first calls build the sessions
Linnet, TensorRT bf1612.5 ms78.1/s0.24 s10.5 GiB0.0352bf16; ONNX Runtime's TensorRT execution provider; first calls build the sessions