Skip to content

vllm_omni.diffusion.models.sensenova_u1.paged_decode

Paged KV cache for SenseNova-U1 autoregressive decode, and a CUDA graph over it.

Autoregressive decode dominates a think request -- 80% of the wall clock, at about 20 ms per token for 42 layers -- and almost all of the GPU idle in that window is Python/ATen dispatch rather than any CUDA call. A CUDA graph retires those dispatches, but capture needs static shapes, and the default cache grows K/V with torch.cat on every step.

Padding the cache to a bucket and masking the tail does give static shapes, but the mask is what costs: measured on an A800 at real decode shapes, a masked bucket attention runs 7.96-11.55 ms per step against 1.03 ms for an exact-length unmasked one, which is more than the graph saves.

The way out is a paged cache. flash_attn_varlen_func takes the used length as a tensor (seqused_k) alongside a block_table, so the buffers stay bucket-sized and capturable while the kernel reads only the valid prefix. One captured graph then serves every length in the bucket.

Scope, and when to delete this. The cache is model-local on purpose: DiffusionKVCacheManager reserves once per scheduler request, and ARDiffusionModelRunner lives under vllm_omni/experimental. SenseNova runs its whole autoregressive loop inside one pipeline forward, so the decode steps never surface to the scheduler and it cannot admit, pool, evict, reuse a prefix for, or continuously batch them. What is here is therefore one set of buffers behind an identity block table, reused by whichever request fits them, not general paged-KV support: it holds while the pipeline serves one sequence per forward, and is released when sleep discards the memory it captured against. Once the decode loop is driven by a scheduler-visible runner, or the pipeline handles more than one sequence per forward, delete this file and use the manager instead.

BLOCK_SIZE module-attribute

BLOCK_SIZE = 16

BUCKETS module-attribute

BUCKETS = (512, 1024, 2048, 4096, 8192)

TAIL_STEP module-attribute

TAIL_STEP = 2048

logger module-attribute

logger = init_logger(__name__)

DecodeGraphRunner

One captured decode step per KV bucket, replayed for every token.

The whole point of the paged cache is that everything the step reads which varies -- the token, its position, and how much of the cache is live -- sits in device tensors. So a single capture serves every step in a bucket: fill the tensors, replay, read the logits out of the static output.

A capture is invalidated when the cache reallocates (tracked by PagedDecodeCache.generation), which happens once per bucket boundary.

cache instance-attribute

cache = cache

captures instance-attribute

captures = 0

device instance-attribute

device = device

indexes instance-attribute

indexes = torch.zeros(3, 1, dtype=torch.long, device=device)

input_ids instance-attribute

input_ids = torch.zeros(
    1, 1, dtype=torch.long, device=device
)

lm instance-attribute

lm = language_model

step

step(token, t_index)

Run one decode step. Returns the static logits tensor.

PagedDecodeCache

Per-layer paged K/V for single-token decode.

Layout follows what flash_attn_varlen_func accepts with a block table: (num_blocks, BLOCK_SIZE, kv_heads, head_dim). seqused holds the number of valid tokens and is the only thing that changes between decode steps, which is what lets a captured graph be replayed unchanged.

block_table instance-attribute

block_table = torch.arange(
    nblocks, device=device, dtype=torch.int32
).unsqueeze(0)

bucket instance-attribute

bucket = _bucket_for(length)

cu_seqlens_q instance-attribute

cu_seqlens_q = torch.tensor(
    [0, 1], device=device, dtype=torch.int32
)

device instance-attribute

device = device

dtype instance-attribute

dtype = dtype

generation instance-attribute

generation = 0

head_dim instance-attribute

head_dim = head_dim

k instance-attribute

k = [
    torch.zeros(shape, device=device, dtype=dtype)
    for _ in range(num_layers)
]

kv_heads instance-attribute

kv_heads = kv_heads

length property

length: int

pos instance-attribute

pos = torch.zeros(1, device=device, dtype=torch.int64)

seqused instance-attribute

seqused = torch.zeros(1, device=device, dtype=torch.int32)

v instance-attribute

v = [
    torch.zeros(shape, device=device, dtype=dtype)
    for _ in range(num_layers)
]

attend

attend(
    layer_idx,
    query_bhsd,
    key_bhsd,
    value_bhsd,
    softmax_scale,
)

Append this step's K/V and attend over the whole valid prefix.

Writes at length - 1 because the caller has already advanced the length for this step. Both the slot and the attended length come from device tensors, so a captured graph reads whatever they hold at replay rather than whatever they held at capture.

from_dynamic_cache classmethod

from_dynamic_cache(
    cache, num_layers, device, dtype, min_length=0
)

Allocate buffers for a prefill cache and copy it in.

grow

grow(length: int) -> bool

Move to the next bucket, keeping the tokens already stored.

Everything is built into locals first and published in one step. A resize that failed part way -- an allocation for a later layer, say -- would otherwise leave new K beside old V under the generation a captured graph is keyed on, so the next request would keep selecting that graph and replay it over storage that had just been freed.

load_prefix

load_prefix(cache) -> bool

Copy a prefill cache into the paged buffers.

Only the first seqused rows are ever read, so the tail is left as whatever the previous request wrote. Returns False for an empty cache.

reusable_for

reusable_for(
    num_layers, kv_heads, head_dim, dtype, length
) -> bool

Can this cache serve another request without reallocating?

Reusing the buffers is what lets captured graphs survive across requests: a capture binds the addresses it recorded, so a fresh cache forces a fresh capture.

set_length

set_length(n: int) -> None

to_dynamic_cache

to_dynamic_cache(cache) -> None

Write the decoded K/V back so the generation stage sees one cache.

Only decode is paged; the DiT stage that follows reads the ordinary cache, so the two are reconciled once at the hand-off rather than kept in sync on every step.

dynamic_lora_wrappers_present

dynamic_lora_wrappers_present(module) -> bool

True once DiffusionLoRAManager has wrapped layers under module.

The manager replaces linear layers with BaseLayerWithLoRA and then binds, rescales or resets slot 0 per request; the wrappers stay in the tree afterwards. Capture records the module tree and the branches that ran during it, so a graph taken before an adapter was bound would replay without the adapter's matmuls and one taken with it bound would survive its removal -- both silently, and neither visible in the shape, dtype and bucket that decide reuse. Nothing captured may outlive a request once these exist.

The distilled LoRA is not affected: it is fused one-way into the weights and leaves no wrapper behind.

paged_decode_supported

paged_decode_supported(
    device: device, head_dim: int
) -> bool