> For clean Markdown of any page, append .md to the page URL.
> For a complete documentation index, see https://docs.nvidia.com/nemo/automodel/llms.txt.
> For AI client integration (Claude Code, Cursor, etc.), connect to the MCP server at https://docs.nvidia.com/nemo/automodel/_mcp/server.

# 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

| Name                                                                                                                | Description                                                           |
| ------------------------------------------------------------------------------------------------------------------- | --------------------------------------------------------------------- |
| [`_block_kinds`](#nemo_automodel-components-models-qwen3_8_flash_next-fa4_qsa-_block_kinds)                         | Classify query/key blocks without host synchronization.               |
| [`_compact_blocks`](#nemo_automodel-components-models-qwen3_8_flash_next-fa4_qsa-_compact_blocks)                   | Compact partial/full block columns for each row.                      |
| [`_load_fa4`](#nemo_automodel-components-models-qwen3_8_flash_next-fa4_qsa-_load_fa4)                               | Cache optional FA4 entry points and the mask callback, never tensors. |
| [`_preprocess`](#nemo_automodel-components-models-qwen3_8_flash_next-fa4_qsa-_preprocess)                           | Build a byte membership table and FA4 forward/reverse block lists.    |
| [`fa4_sparse_gqa_attention`](#nemo_automodel-components-models-qwen3_8_flash_next-fa4_qsa-fa4_sparse_gqa_attention) | Evaluate exact route-set GQA using FA4 on SM90.                       |

### API

```python
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:

```python
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

```python
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.

```python
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

```python
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 | None` — default: None

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

---

**Returns:** `torch.Tensor`

Independent BF16 CUDA tensor \[batch, queries, query\_heads, 256].