ReferenceFull Library ReferenceNemo AutomodelNemo AutomodelComponentsModelsGlm Moe DsaKernelsnemo_automodel.components.models.glm_moe_dsa.kernels.cudnn_dsa

nemo_automodel.components.models.glm_moe_dsa.kernels.cudnn_dsa

View as Markdown

cuDNN and FlashMLA kernels for the split GLM-5.2 DSA path.

Module Contents

Classes

NameDescription
CudnnDsaPackedMetadataReusable THD metadata for local-query/global-key cuDNN DSA.

Functions

NameDescription
_compact_and_sort_indicesCanonicalize global indices [T, K] and return valid lengths [T].
_require_availableRaise when either optional runtime required by the split kernel is absent.
_require_cuda_tensorsValidate that arbitrary-layout input tensors share one SM90+ CUDA device.
_topk_wrapper_chunkedSelect top-k for FP32 scores [T, S_max] using causal lengths [T].
_unpack_packed_metadataValidate reusable packed metadata without a CUDA-to-host synchronization.
_validate_topkValidate GLM-5.2’s fixed sparse-selection width.
cudnn_indexer_topkCompute GLM-5.2 packed-THD indexer top-k with cuDNN Frontend.
cudnn_sparse_attentionRun GLM-5.2 sparse MLA through the shared cuDNN adapter.
is_cudnn_dsa_availableReturn whether both optional libraries required by cuDNN DSA import.
prepare_cudnn_dsa_packed_metadataBuild segmented THD metadata for local queries and gathered global keys.

Data

_ATTENTION_HEAD_DIM

_INDEX_HEAD_DIM

_MAX_TOPK

_TOPK_ROW_ALIGNMENT

_TOPK_SCRATCH_INT32_FACTOR

_TOPK_SCRATCH_LIMIT_BYTES

API

class nemo_automodel.components.models.glm_moe_dsa.kernels.cudnn_dsa.CudnnDsaPackedMetadata(
starts: torch.Tensor,
causal_lengths: torch.Tensor,
query_valid: torch.Tensor,
valid_row_indices: torch.Tensor | None,
segment_cu_q: torch.Tensor,
segment_cu_k: torch.Tensor,
q_causal_offsets: torch.Tensor,
key_source_indices: torch.Tensor | None,
max_seqlen_q: int,
max_seqlen_k: int,
total_key_tokens: int,
all_rows_nonempty: bool
)
Dataclass

Reusable THD metadata for local-query/global-key cuDNN DSA.

all_rows_nonempty
bool

Whether every query row has at least one valid key.

causal_lengths
Tensor

Number of real document keys visible to each query, int32 [T_q].

key_source_indices
Tensor | None

Optional int64 [segmented_key_tokens] indices into the gathered padded-storage K/V tensor.

max_seqlen_k
int

Maximum repacked key-prefix segment length.

max_seqlen_q
int

Maximum local query-segment length.

q_causal_offsets
Tensor

Uncompressed document-relative position of each segment’s first local query, int32 [segments].

query_valid
Tensor

Boolean [T_q] mask for real, non-padding query rows.

segment_cu_k
Tensor

Repacked key-prefix segment offsets, int32 [segments + 1].

segment_cu_q
Tensor

Local query-segment offsets, int32 [segments + 1].

starts
Tensor

Global padded-storage start of each query’s document, int32 [T_q].

total_key_tokens
int

Number of rows in the gathered padded-storage K/V tensor.

valid_row_indices
Tensor | None

Optional int64 indices of valid queries, [T_valid]; None when every query is valid.

nemo_automodel.components.models.glm_moe_dsa.kernels.cudnn_dsa._compact_and_sort_indices(
indices: torch.Tensor,
key_count: int
) -> tuple[torch.Tensor, torch.Tensor]

Canonicalize global indices [T, K] and return valid lengths [T].

nemo_automodel.components.models.glm_moe_dsa.kernels.cudnn_dsa._require_available() -> None

Raise when either optional runtime required by the split kernel is absent.

nemo_automodel.components.models.glm_moe_dsa.kernels.cudnn_dsa._require_cuda_tensors(
operation: str,
tensors: torch.Tensor = ()
) -> tuple[int, int]

Validate that arbitrary-layout input tensors share one SM90+ CUDA device.

nemo_automodel.components.models.glm_moe_dsa.kernels.cudnn_dsa._topk_wrapper_chunked(
scores: torch.Tensor,
seq_lens: torch.Tensor,
topk: int
) -> torch.Tensor

Select top-k for FP32 scores [T, S_max] using causal lengths [T].

nemo_automodel.components.models.glm_moe_dsa.kernels.cudnn_dsa._unpack_packed_metadata(
total_query_tokens: int,
total_key_tokens: int,
device: torch.device

Validate reusable packed metadata without a CUDA-to-host synchronization.

Parameters:

packed_metadata
CudnnDsaPackedMetadata

Metadata whose per-query fields have shape [T_q] and whose key-source indices address the gathered padded-storage K/V tensor.

total_query_tokens
int

Expected local query row count T_q.

total_key_tokens
int

Expected gathered padded-storage K/V row count T_k.

device
torch.device

CUDA device shared by the metadata and kernel inputs.

Returns: CudnnDsaPackedMetadata

The validated metadata object, unchanged.

nemo_automodel.components.models.glm_moe_dsa.kernels.cudnn_dsa._validate_topk(
index_topk: int
) -> None

Validate GLM-5.2’s fixed sparse-selection width.

nemo_automodel.components.models.glm_moe_dsa.kernels.cudnn_dsa.cudnn_indexer_topk(
index_q: torch.Tensor,
index_k: torch.Tensor,
head_weights: torch.Tensor,
cu_seqlens: torch.Tensor,
index_topk: int,
query_indices: torch.Tensor | None = None,
cu_seqlens_padded: torch.Tensor | None = None,
) -> torch.Tensor

Compute GLM-5.2 packed-THD indexer top-k with cuDNN Frontend.

Parameters:

index_q
torch.Tensor

Rank-local indexer query in THD layout, BF16 [T_q, H_index, 128].

index_k
torch.Tensor

Gathered global indexer key in TD layout, BF16 [T_k, 128].

head_weights
torch.Tensor

Already-scaled per-head weights in TH layout, FP32 or BF16 [T, H_index]. The caller owns both H_index**-0.5 and 128**-0.5 scaling; this function does not rescale them.

cu_seqlens
torch.Tensor

Compact packed-sequence offsets, CUDA int32 [num_sequences + 1].

index_topk
int

Fixed output width K in [1, 2048].

query_indices
torch.Tensor | NoneDefaults to None

Optional contiguous global padded coordinates for local query rows, [T_q]. Absent means CP=1 identity coordinates.

cu_seqlens_padded
torch.Tensor | NoneDefaults to None

Optional global padded packed-layout offsets.

packed_metadata
CudnnDsaPackedMetadata | NoneDefaults to None

Optional metadata object returned by prepare_cudnn_dsa_packed_metadata. Supplying it avoids rebuilding and synchronizing the same metadata in every full-indexer layer.

Returns: torch.Tensor

CUDA int32 top-k indices in global padded-storage THD coordinates,

Raises:

  • RuntimeError: If optional kernels, CUDA, or SM90+ are unavailable.
  • TypeError: If tensor dtypes or index_topk are invalid.
  • ValueError: If tensor shapes or compact packed metadata are invalid.
nemo_automodel.components.models.glm_moe_dsa.kernels.cudnn_dsa.cudnn_sparse_attention(
q: torch.Tensor,
kv_latent: torch.Tensor,
topk_indices: torch.Tensor,
softmax_scale: float,
topk_length: torch.Tensor | None = None,
all_rows_nonempty: bool = False,
valid_row_indices: torch.Tensor | None = None
) -> torch.Tensor

Run GLM-5.2 sparse MLA through the shared cuDNN adapter.

Parameters:

q
torch.Tensor

Contiguous CUDA BF16 absorbed query tensor of shape [query_tokens, heads, 576].

kv_latent
torch.Tensor

Contiguous CUDA BF16 latent K/V tensor of shape [key_tokens, 1, 576].

topk_indices
torch.Tensor

Contiguous CUDA int32 tensor of shape [query_tokens, 1, sparse_width] with global K/V coordinates and invalid entries marked -1.

softmax_scale
float

Scale forwarded unchanged to FlashMLA and cuDNN backward.

topk_length
torch.Tensor | NoneDefaults to None

Optional contiguous CUDA int32 valid-prefix lengths of shape [query_tokens].

all_rows_nonempty
boolDefaults to False

Whether every query has a positive valid-prefix length.

valid_row_indices
torch.Tensor | NoneDefaults to None

Optional contiguous CUDA int64 indices of nonempty queries with shape [valid_query_tokens].

Returns: torch.Tensor

Contiguous CUDA BF16 latent output tensor of shape

nemo_automodel.components.models.glm_moe_dsa.kernels.cudnn_dsa.is_cudnn_dsa_available() -> bool

Return whether both optional libraries required by cuDNN DSA import.

nemo_automodel.components.models.glm_moe_dsa.kernels.cudnn_dsa.prepare_cudnn_dsa_packed_metadata(
cu_seqlens: torch.Tensor,
total_key_tokens: int,
max_seqlen: int | torch.Tensor | None = None,
query_indices: torch.Tensor | None = None,
cu_seqlens_padded: torch.Tensor | None = None,
padding_mask: torch.Tensor | None = None

Build segmented THD metadata for local queries and gathered global keys.

This validation intentionally performs one device-to-host synchronization. The model prepares the object once per pipeline stage and reuses it across every indexer and shared-attention layer.

Parameters:

cu_seqlens
torch.Tensor

Cumulative compact real-token lengths, int32 tensor of shape [sequences + 1].

total_key_tokens
int

Number of tokens in the gathered, padded THD key tensor.

max_seqlen
int | torch.Tensor | NoneDefaults to None

Optional precomputed maximum sequence length as a Python integer or scalar integer tensor. Its value is checked against cu_seqlens.

query_indices
torch.Tensor | NoneDefaults to None

Optional contiguous integer tensor of shape [T_q] with global padded-storage coordinates for the local query rows. When absent, queries cover all global key tokens (CP=1).

cu_seqlens_padded
torch.Tensor | NoneDefaults to None

Optional cumulative padded-storage boundaries, int32 tensor of shape [sequences + 1]. When absent, cu_seqlens also defines storage coordinates.

padding_mask
torch.Tensor | NoneDefaults to None

Optional local boolean padding mask of shape [T_q]. This remains authoritative when THD preprocessing has absorbed trailing pack padding into cu_seqlens.

Returns: CudnnDsaPackedMetadata

Metadata with per-query global padded-storage starts and real causal lengths,

nemo_automodel.components.models.glm_moe_dsa.kernels.cudnn_dsa._ATTENTION_HEAD_DIM = 576
nemo_automodel.components.models.glm_moe_dsa.kernels.cudnn_dsa._INDEX_HEAD_DIM = 128
nemo_automodel.components.models.glm_moe_dsa.kernels.cudnn_dsa._MAX_TOPK = 2048
nemo_automodel.components.models.glm_moe_dsa.kernels.cudnn_dsa._TOPK_ROW_ALIGNMENT = 512
nemo_automodel.components.models.glm_moe_dsa.kernels.cudnn_dsa._TOPK_SCRATCH_INT32_FACTOR = 2
nemo_automodel.components.models.glm_moe_dsa.kernels.cudnn_dsa._TOPK_SCRATCH_LIMIT_BYTES = 2 * 1024 * 1024 * 1024