| |
| |
| |
|
|
| |
| |
| """ |
| 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] |
|
|