Skip to content

vllm_omni.diffusion.models.ltx2.ops.fna

TileLang FNA for the LTX-2.5 DiffVAE stage-5 schedule.

This intentionally reuses NATTEN's token permutation. It only replaces the middle BF16 attention kernel for Q tiles (4, 4, 4), KV tiles (4, 4, 8), head_dim=64, stride=dilation=1, and scale=1.0.

HEADS module-attribute

HEADS = 4

HEAD_DIM module-attribute

HEAD_DIM = 64

KV_BUCKET_LIMITS module-attribute

KV_BUCKET_LIMITS = (32, 48, 60, 75)

KV_TILE module-attribute

KV_TILE = (4, 4, 8)

KV_TILE_ID_BITS module-attribute

KV_TILE_ID_BITS = 15

KV_TOKENS module-attribute

KV_TOKENS = math.prod(KV_TILE)

LOG2_E module-attribute

LOG2_E = 1.4426950408889634

PATTERN_TABLE_SIZE module-attribute

PATTERN_TABLE_SIZE = 16

Q_TILE module-attribute

Q_TILE = (4, 4, 4)

Q_TOKENS module-attribute

Q_TOKENS = math.prod(Q_TILE)

SHAPE_CACHE_SIZE module-attribute

SHAPE_CACHE_SIZE = 16

Problem dataclass

frames instance-attribute

frames: int

height instance-attribute

height: int

kv_grid property

kv_grid: tuple[int, int, int]

kv_length property

kv_length: int

kv_padded property

kv_padded: tuple[int, int, int]

q_grid property

q_grid: tuple[int, int, int]

q_length property

q_length: int

q_padded property

q_padded: tuple[int, int, int]

shape property

shape: tuple[int, int, int]

width instance-attribute

width: int

window instance-attribute

window: tuple[int, int, int]

build_kv_tile_buckets cached

build_kv_tile_buckets(
    problem: Problem,
) -> tuple[tuple[Tensor, Tensor, int], ...]

Group Q tiles by an upper bound on their useful KV tile count.

build_kv_tile_map cached

build_kv_tile_map(
    problem: Problem,
) -> tuple[Tensor, Tensor]

Return packed KV tile ids and valid counts for each permuted Q tile.

build_kv_tile_metadata cached

build_kv_tile_metadata(
    problem: Problem,
) -> tuple[Tensor, Tensor, Tensor, Tensor]

Pack KV tile id and three separable mask pattern ids into int32.

compile_kernel cached

compile_kernel(max_kv_tiles: int, *, target: str = 'auto')

Compile one T/H/W-dynamic kernel for a fixed KV-count ceiling.

fna3d_tilelang

fna3d_tilelang(
    query: Tensor,
    key: Tensor,
    value: Tensor,
    *,
    shape: tuple[int, int, int],
    window: tuple[int, int, int],
) -> Tensor

Run native stage-5 FNA, or raise for unsupported inputs/launch failures.

Q and KV use the documented token-tiled layout. Batch members share spatial metadata; compilation is cached by target architecture and KV-count bucket.