Skip to content

vllm_omni.diffusion.models.auk.vae_cudagraph

CUDA graph replay for the AuK codec decode.

The decoder launches about a thousand small kernels per clip, so it is bound by launch overhead rather than by the GPU. This wrapper replays the whole decode as CUDA graphs, in three tiers:

Compiled bucket graphs are built at startup by warmup(): decode is passed through torch.compile so Inductor fuses the elementwise chains, then one graph is captured per bucket in compile_shapes. Shorter clips are right-padded to their bucket with the normalized latent that decodes to raw zero, which is what the eager decode's own convolution padding sees, so the padding only reaches the last frames through the upsamplers' lookahead. The fused kernels differ from eager in rounding order, so these graphs are close to but not bit-identical with the eager decode.

Clips longer than tile_frames are decoded in tiles of that many frames (by default the largest bucket, so every tile replays the same compiled graph). Neighbouring tiles overlap by the decoder's receptive field, taken from AuKVAE.decode_context_frames(), and each tile only contributes the frames that lie outside that halo, so the stitched waveform equals the whole-clip decode up to floating-point order. Memory and startup cost stay constant however long the clip; the price is the halo recomputed per tile (about an eighth of a 512-frame tile). decode_tiles() yields the tiles as they finish so a caller can stream them.

Plain graphs are captured on demand, one per exact latent length, for windows no compiled bucket serves (tiling off, or a bucket whose compile failed). They replay the same kernels as eager and are bit-identical with it. They share one wrapper-private graph pool, so they are never evicted one at a time: when max_graphs is reached the whole generation is retired and a fresh pool is started, and no surviving graph can hold an address that a destroyed graph released.

DEFAULT_COMPILE_SHAPES module-attribute

DEFAULT_COMPILE_SHAPES: tuple[int, ...] = (128, 256, 512)

logger module-attribute

logger = init_logger(__name__)

AuKVAEDecodeGraph

Replay AuKVAE.decode for one clip per call.

compile_shapes instance-attribute

compile_shapes = sorted(
    {int(size) for size in compile_shapes if int(size) > 0}
)

context_frames instance-attribute

context_frames = vae.decode_context_frames()

enabled instance-attribute

enabled = bool(enabled)

frame_alignment instance-attribute

frame_alignment = max(1, int(frame_alignment))

last_mode instance-attribute

last_mode: str | None = None

max_graphs instance-attribute

max_graphs = max(1, int(max_graphs))

tile_frames instance-attribute

tile_frames = max(0, int(tile_frames))

vae instance-attribute

vae = vae

compiled_bucket

compiled_bucket(frames: int) -> int | None

The smallest compiled bucket that holds frames latents, or None past the largest.

decode_tiles

decode_tiles(latents: Tensor)

Yield (emit_start_frame, waveform) per tile, in order, for a [1, frames, latent_dim] clip.

The waveforms are views into graph buffers where a graph served the tile; copy before the next tile if they must outlive the iteration.

warmup

warmup(device: device | str) -> None

Compile the decode and capture every bucket in compile_shapes.

Meant for service startup, since each bucket costs one Inductor compilation. A failure only logs a warning and the bucket is served by a plain graph instead.

plan_tiles

plan_tiles(
    frames: int,
    tile: int,
    left: int,
    right: int,
    sizes: Sequence[int] = (),
) -> list[tuple[int, int, int, int]]

Split a clip into decode windows: (window_start, window_frames, emit_start, emit_end).

Windows are tile frames wide. A window contributes the frames that are at least left frames after its start and right frames before its end, unless that edge is the clip's own edge, so each emitted frame has its full decoder context inside its window. The last window is pinned to the clip's end; it shrinks to the smallest of sizes (the compiled buckets) that still holds the remaining frames plus the left context, so a clip just past one tile does not pay for a second full one. A clip that fits in one tile is a single window of its own length.