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

rms_norm_key_cache

rms_norm_key_cache()​

max.experimental.nn.common_layers.functional_kernels.rms_norm_key_cache(kv_params, kv_collection, gamma, epsilon, layer_idx, total_seq_len, input_row_offsets, weight_offset, rms_norm_cols=None, multiply_before_cast=True, per_head_norm=True)

source

Applies RMSNorm to the new entries in the KV cache.

When per_head_norm is True (the default), RMSNorm is applied separately to each head. In this mode, gamma should have size [head_dim] and normalization occurs across the head_dim dimensions within each head.

When per_head_norm is False, RMSNorm is applied per token across all heads. In this mode, gamma should have size [n_kv_heads * head_dim] and normalization occurs across all dimensions for each token.

The size of the gamma tensor determines how many dimensions will be normalized. If gamma’s size doesn’t match the expected size based on the per_head_norm setting, rms_norm_cols must be explicitly specified to confirm the intention to normalize only a subset of dimensions.

The KV cache collection itself isn’t aware of the new cache entries until the cache length increment, which happens after the model forward, so input_row_offsets does this bookkeeping.

Parameters:

  • kv_params (KVCacheParams) – The KV cache parameters.
  • kv_collection (KVCacheInputsPerDevice[TensorValue, BufferValue]) – The paged KV cache holding the entries to normalize.
  • gamma (TensorValue) – The RMSNorm weight.
  • epsilon (float | floating[Any]) – The epsilon added inside the normalization.
  • layer_idx (TensorValue) – The index of the layer being normalized.
  • total_seq_len (Dim) – The total sequence length of the ragged batch.
  • input_row_offsets (TensorValue) – The ragged offsets delimiting the new entries.
  • weight_offset (float | floating[Any]) – The offset added to gamma before the multiply.
  • rms_norm_cols (int | None) – The number of columns to normalize. Required when gamma’s size doesn’t match the expected size, to confirm the intention to normalize only a subset of dimensions.
  • multiply_before_cast (bool) – Whether to multiply by gamma before casting back to the cache dtype.
  • per_head_norm (bool) – Whether to normalize each head separately. Defaults to True.

Return type:

None