ironic-mirror / python /xrex_unified /softmax_state.py
SNAPKITTYWEST's picture
push from SNAPKITTYWEST/ironic-mirror
677e207 verified
Raw
History Blame Contribute Delete
1.59 kB
#
# 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]