nemo_automodel.components.models.qwen3_8_flash_next.fa4_qsa
nemo_automodel.components.models.qwen3_8_flash_next.fa4_qsa
SM90 QSA using PyTorch preprocessing and FlashAttention 4.
Optional FA4/CuTe dependencies are loaded on the first CUDA call. The only CuTe code owned here is the mask callback compiled into FA4’s kernels.
Module Contents
Functions
API
Classify query/key blocks without host synchronization.
Parameters:
Uint8 tensor [batch, queries, keys].
Number of physical keys in one block; queries use 128.
Returns: torch.Tensor
Int32 [batch, ceil(queries/128), ceil(keys/block_keys)] tile classes:
Compact partial/full block columns for each row.
Parameters:
Integer tensor [batch, block_rows, block_columns], with values zero (absent), one (partial) or two (full).
Returns: torch.Tensor
Int32 partial counts [batch, 1, block_rows], partial column indices
Cache optional FA4 entry points and the mask callback, never tensors.
Build a byte membership table and FA4 forward/reverse block lists.
Parameters:
Signed int32/int64 tensor [batch, queries, routes_per_query]. Duplicate IDs collapse; invalid IDs, including large int64 values, are discarded before indexing.
Positive number of physical key/value positions.
Returns: torch.Tensor
Uint8 membership [batch, queries, kv_length], with each row padded to
Evaluate exact route-set GQA using FA4 on SM90.
Duplicate routes select a token once. Invalid IDs are ignored. Empty query rows have zero output and query gradients. Routes encode causality, documents and padding; no additional triangular mask is imposed. FA4 owns first-order autograd; higher-order gradients and deterministic backward are unsupported. Only discrete preprocessing is torch-compiled.
Parameters:
BF16 CUDA tensor [batch, queries, query_heads, 256]. Arbitrary strides are accepted; non-unit final strides are copied.
BF16 CUDA tensor [batch, keys, kv_heads, 256]. May contain gathered global K/V while query is a local CP slice. query_heads must be a positive multiple of kv_heads.
BF16 CUDA tensor with key’s shape and device.
Signed int32/int64 CUDA tensor [batch, queries, routes_per_query] in physical K/V coordinates. Noncontiguous inputs are accepted. Dimensions must be nonempty.
Finite positive score multiplier; defaults to 1/sqrt(256).
Returns: torch.Tensor
Independent BF16 CUDA tensor [batch, queries, query_heads, 256].