nemo_automodel.components.models.glm_moe_dsa.kernels.cudnn_dsa
nemo_automodel.components.models.glm_moe_dsa.kernels.cudnn_dsa
cuDNN and FlashMLA kernels for the split GLM-5.2 DSA path.
Module Contents
Classes
Functions
Data
API
Reusable THD metadata for local-query/global-key cuDNN DSA.
Whether every query row has at least one valid key.
Number of real document keys visible to each query, int32
[T_q].
Optional int64 [segmented_key_tokens] indices into
the gathered padded-storage K/V tensor.
Maximum repacked key-prefix segment length.
Maximum local query-segment length.
Uncompressed document-relative position of each segment’s
first local query, int32 [segments].
Boolean [T_q] mask for real, non-padding query rows.
Repacked key-prefix segment offsets, int32
[segments + 1].
Local query-segment offsets, int32 [segments + 1].
Global padded-storage start of each query’s document, int32
[T_q].
Number of rows in the gathered padded-storage K/V tensor.
Optional int64 indices of valid queries, [T_valid];
None when every query is valid.
Canonicalize global indices [T, K] and return valid lengths [T].
Raise when either optional runtime required by the split kernel is absent.
Validate that arbitrary-layout input tensors share one SM90+ CUDA device.
Select top-k for FP32 scores [T, S_max] using causal lengths [T].
Validate reusable packed metadata without a CUDA-to-host synchronization.
Parameters:
Metadata whose per-query fields have shape [T_q] and
whose key-source indices address the gathered padded-storage K/V tensor.
Expected local query row count T_q.
Expected gathered padded-storage K/V row count T_k.
CUDA device shared by the metadata and kernel inputs.
Returns: CudnnDsaPackedMetadata
The validated metadata object, unchanged.
Validate GLM-5.2’s fixed sparse-selection width.
Compute GLM-5.2 packed-THD indexer top-k with cuDNN Frontend.
Parameters:
Rank-local indexer query in THD layout, BF16 [T_q, H_index, 128].
Gathered global indexer key in TD layout, BF16 [T_k, 128].
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.
Compact packed-sequence offsets, CUDA int32
[num_sequences + 1].
Fixed output width K in [1, 2048].
Optional contiguous global padded coordinates for local query
rows, [T_q]. Absent means CP=1 identity coordinates.
Optional global padded packed-layout offsets.
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 orindex_topkare invalid.ValueError: If tensor shapes or compact packed metadata are invalid.
Run GLM-5.2 sparse MLA through the shared cuDNN adapter.
Parameters:
Contiguous CUDA BF16 absorbed query tensor of shape
[query_tokens, heads, 576].
Contiguous CUDA BF16 latent K/V tensor of shape
[key_tokens, 1, 576].
Contiguous CUDA int32 tensor of shape
[query_tokens, 1, sparse_width] with global K/V coordinates
and invalid entries marked -1.
Scale forwarded unchanged to FlashMLA and cuDNN backward.
Optional contiguous CUDA int32 valid-prefix lengths of shape
[query_tokens].
Whether every query has a positive valid-prefix length.
Optional contiguous CUDA int64 indices of nonempty
queries with shape [valid_query_tokens].
Returns: torch.Tensor
Contiguous CUDA BF16 latent output tensor of shape
Return whether both optional libraries required by cuDNN DSA import.
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:
Cumulative compact real-token lengths, int32 tensor of shape
[sequences + 1].
Number of tokens in the gathered, padded THD key tensor.
Optional precomputed maximum sequence length as a Python integer or
scalar integer tensor. Its value is checked against cu_seqlens.
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).
Optional cumulative padded-storage boundaries, int32 tensor
of shape [sequences + 1]. When absent, cu_seqlens also defines
storage coordinates.
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,