# # Copyright (c) 2026 BEL ESPRIT D ACCORD TRUST HOLDINGS INC # All rights reserved. # SPDX-License-Identifier: Apache-2.0 # Copyright 2026 X.AI Corp. """ Online Softmax State for Flash Attention. Invariant: l = sum(exp(qk - m)), m = max(qk) per query. """ from dataclasses import dataclass from typing import Optional import jax import jax.numpy as jnp from .cap_functions import CapMethod, CapParams, cap_forward @dataclass class SoftmaxState: m: jax.Array l: jax.Array acc: jax.Array @staticmethod def init(block_q: int, head_dim: int, dtype=jnp.float32) -> "SoftmaxState": return SoftmaxState( m=jnp.full((block_q,), -jnp.inf, dtype=dtype), l=jnp.zeros((block_q,), dtype=dtype), acc=jnp.zeros((block_q, head_dim), dtype=dtype) ) def update(self, qk: jax.Array, v: jax.Array, cap_method: CapMethod, cap_params: CapParams, temp: Optional[jax.Array] = None) -> "SoftmaxState": if temp is not None: qk = qk * temp[..., None] qk_capped = cap_forward(qk, cap_method, cap_params) m_new = jnp.maximum(self.m, jnp.max(qk_capped, axis=1)) alpha = jnp.exp(self.m - m_new) l_new = self.l * alpha + jnp.sum(jnp.exp(qk_capped - m_new[:, None]), axis=1) p = jnp.exp(qk_capped - m_new[:, None]) p = p / l_new[:, None] acc_new = self.acc * alpha[:, None] + p @ v return SoftmaxState(m=m_new, l=l_new, acc=acc_new) def finalize(self) -> jax.Array: return self.acc / self.l[:, None]