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.
AuKVAEDecodeGraph ¶
Replay AuKVAE.decode for one clip per call.
compile_shapes instance-attribute ¶
compiled_bucket ¶
The smallest compiled bucket that holds frames latents, or None past the largest.
decode_tiles ¶
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.
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.