Skip to content

utils

utils

mean_nonpadding(values: jax.Array, padding_mask: Optional[jax.Array]) -> jax.Array

Average one scalar per token, excluding padded positions.

state_cosine(left: jax.Array, right: jax.Array) -> jax.Array

Compute per-token cosine similarity between two residual states.