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)
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
gammabefore 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
gammabefore casting back to the cache dtype. - per_head_norm (bool) – Whether to normalize each head separately. Defaults
to
True.
-
Return type:
-
None