Download python/xrex_unified/softmax_state.py from Snapkitty/ironic-mirror: direct link, hf CLI and curl.
- Browser
- Download file 1.59 kB
-
https://huggingface.co/Snapkitty/ironic-mirror/resolve/main/python/xrex_unified/softmax_state.py
- Command line
-
hf download hf://Snapkitty/ironic-mirror/python/xrex_unified/softmax_state.py
-
curl -L -o softmax_state.py https://huggingface.co/Snapkitty/ironic-mirror/resolve/main/python/xrex_unified/softmax_state.py
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 | |
| class SoftmaxState: | |
| m: jax.Array | |
| l: jax.Array | |
| acc: jax.Array | |
| 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] | |