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.