IMPORTANT: To view this page as Markdown, append `.md` to the URL (e.g. /get-started.md). For the complete documentation index, see llms.txt.
Skip to main content
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)

source

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 (PagedCacheValues) – The paged KV cache.
  • layer_idx (TensorValue) – The layer index.
  • n_heads (int) – The number of query attention heads.
  • interleaved (bool) – Whether freqs_cis uses interleaved (re, im) format.
  • position_ids (TensorValue | None) – The optional ragged 2D array of position IDs. If None, defaults to cache_length + token_idx for each token. When num_sections > 1, mrope_section must be provided. Shape: [num_sections, total_seq_len].
  • mrope_section (list[int] | None) – The optional list of ints indicating the section of the head_dim to apply RoPE to. Must be used with position_ids.
  • fuse (bool) – If True (the default), emits a single fused custom op. If False, 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 (with k_norm_weight and rms_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 with position_ids.
  • k_norm_weight (TensorValue | None) – The per-head RMSNorm gamma [head_dim] for K (see q_norm_weight).
  • rms_norm_eps (float | None) – The epsilon for the fused qk-norm; required when q_norm_weight is set.
  • k_eq_v (bool) – When True (only valid with q_norm_weight), V has no own projection and reuses K’s: qkv is [q|k] (no V region) and the kernel reads the K head for both the K and V stores, sharing the norm reduction. When False (the default), qkv is [q|k|v].

Returns:

The roped Q output, [total_seq_len, n_heads * head_dim].

Return type:

Tensor