nemo_automodel.components.models.qwen3_8_flash_next.fa4_qsa

View as Markdown

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

NameDescription
_block_kindsClassify query/key blocks without host synchronization.
_compact_blocksCompact partial/full block columns for each row.
_load_fa4Cache optional FA4 entry points and the mask callback, never tensors.
_preprocessBuild a byte membership table and FA4 forward/reverse block lists.
fa4_sparse_gqa_attentionEvaluate exact route-set GQA using FA4 on SM90.

API

nemo_automodel.components.models.qwen3_8_flash_next.fa4_qsa._block_kinds(
membership: torch.Tensor,
block_keys: int
) -> torch.Tensor

Classify query/key blocks without host synchronization.

Parameters:

membership
torch.Tensor

Uint8 tensor [batch, queries, keys].

block_keys
int

Number of physical keys in one block; queries use 128.

Returns: torch.Tensor

Int32 [batch, ceil(queries/128), ceil(keys/block_keys)] tile classes:

nemo_automodel.components.models.qwen3_8_flash_next.fa4_qsa._compact_blocks(
kinds: torch.Tensor
) -> tuple[torch.Tensor, ...]

Compact partial/full block columns for each row.

Parameters:

kinds
torch.Tensor

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

nemo_automodel.components.models.qwen3_8_flash_next.fa4_qsa._load_fa4() -> tuple[collections.abc.Callable, type, collections.abc.Callable]

Cache optional FA4 entry points and the mask callback, never tensors.

nemo_automodel.components.models.qwen3_8_flash_next.fa4_qsa._preprocess(
routes: torch.Tensor,
kv_length: int
) -> tuple[torch.Tensor, ...]

Build a byte membership table and FA4 forward/reverse block lists.

Parameters:

routes
torch.Tensor

Signed int32/int64 tensor [batch, queries, routes_per_query]. Duplicate IDs collapse; invalid IDs, including large int64 values, are discarded before indexing.

kv_length
int

Positive number of physical key/value positions.

Returns: torch.Tensor

Uint8 membership [batch, queries, kv_length], with each row padded to

nemo_automodel.components.models.qwen3_8_flash_next.fa4_qsa.fa4_sparse_gqa_attention(
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
selected_token_ids: torch.Tensor,
softmax_scale: float | None = None
) -> torch.Tensor

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:

query
torch.Tensor

BF16 CUDA tensor [batch, queries, query_heads, 256]. Arbitrary strides are accepted; non-unit final strides are copied.

key
torch.Tensor

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.

value
torch.Tensor

BF16 CUDA tensor with key’s shape and device.

selected_token_ids
torch.Tensor

Signed int32/int64 CUDA tensor [batch, queries, routes_per_query] in physical K/V coordinates. Noncontiguous inputs are accepted. Dimensions must be nonempty.

softmax_scale
float | NoneDefaults to None

Finite positive score multiplier; defaults to 1/sqrt(256).

Returns: torch.Tensor

Independent BF16 CUDA tensor [batch, queries, query_heads, 256].