For the complete documentation index, see llms.txt. Markdown versions of all pages are available by appending .md to any URL (e.g. /get-started.md).
Python function
rope_split_store_ragged
rope_split_store_ragged()
max.experimental.nn.common_layers.functional_kernels.rope_split_store_ragged(kv_params, qkv, input_row_offsets, freqs_cis, kv_collection, layer_idx, n_heads, interleaved=True, position_ids=None, mrope_section=None, fuse=True, q_out_dtype=None, q_norm_weight=None, k_norm_weight=None, rms_norm_eps=None, k_eq_v=False)
Applies RoPE to Q and K from a flat QKV buffer and stores K/V to the cache.
Reads from a flat QKV matmul output, applies RoPE to the Q and K regions, stores K/V to the paged KV cache, and writes the roped Q to the output.
-
Parameters:
-
- kv_params (KVCacheParams) – The KV cache parameters.
- qkv (TensorValue) – The flat QKV matmul output,
[total_seq_len, q_dim + k_dim + v_dim]. - input_row_offsets (TensorValue) – The ragged offsets,
[batch_size + 1]. - freqs_cis (TensorValue) – The RoPE frequencies,
[max_seq_len, head_dim]. - kv_collection (KVCacheInputsPerDevice[TensorValue, BufferValue]) – The paged KV cache.
- layer_idx (TensorValue) – The layer index.
- n_heads (int) – The number of query attention heads.
- interleaved (bool) – Whether
freqs_cisuses interleaved (re, im) format. - position_ids (TensorValue | None) – The optional ragged 2D array of position IDs. If
None, defaults tocache_length + token_idxfor each token. Whennum_sections > 1,mrope_sectionmust be provided. Shape:[num_sections, total_seq_len]. - mrope_section (list[int] | None) – The optional list of ints indicating the section of
the
head_dimto apply RoPE to. Must be used withposition_ids. - fuse (bool) – If
True(the default), emits a single fused custom op. IfFalse, emits separate split, rope, and store ops for testing graph compiler fusion. - q_out_dtype (DType | None) – The dtype for the roped Q output. Defaults to
qkv.dtype. - q_norm_weight (TensorValue | None) – The optional per-head RMSNorm gamma
[head_dim]for Q. When given (withk_norm_weightandrms_norm_eps), the per-head Q/K/V RMS-norm is fused into the op (Q/K use their gammas, V is a bare norm), removing the separate norm ops. Mutually exclusive withposition_ids. - k_norm_weight (TensorValue | None) – The per-head RMSNorm gamma
[head_dim]for K (seeq_norm_weight). - rms_norm_eps (float | None) – The epsilon for the fused qk-norm; required when
q_norm_weightis set. - k_eq_v (bool) – When
True(only valid withq_norm_weight), V has no own projection and reuses K’s:qkvis[q|k](no V region) and the kernel reads the K head for both the K and V stores, sharing the norm reduction. WhenFalse(the default),qkvis[q|k|v].
-
Returns:
-
The roped Q output,
[total_seq_len, n_heads * head_dim]. -
Return type: