ironic-mirror / python /xrex_unified /cap_functions.py
SNAPKITTYWEST's picture
push from SNAPKITTYWEST/ironic-mirror
677e207 verified
Raw
History Blame Contribute Delete
3.6 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.
"""
Cap Functions for Attention Logit Capping.
Includes the novel Asymmetric Log-Sigmoid Cap (ALSC):
ALSC(x; cap, alpha, beta) = cap * [sigma(alpha(x/cap + beta)) - sigma(alpha*beta)]
/ [sigma(alpha*beta) * (1 - sigma(alpha*beta))]
Properties of ALSC vs existing methods:
- tanh: symmetric, bound [-cap, cap], gradient vanishes exponentially
- soft_sign: symmetric, bound (-cap, cap), gradient decays as 1/x^2
- ALSC: ASYMMETRIC, bound [0, cap], tunable gradient via alpha, dead zone via beta
Analytic inverse exists. Matches attention's natural asymmetry.
"""
from dataclasses import dataclass
from typing import Optional, Literal
import jax
import jax.numpy as jnp
CapMethod = Literal["tanh", "soft_sign", "alsc", "none"]
@dataclass(frozen=True)
class CapParams:
cap: float = 0.0
alpha: float = 4.0
beta: float = 0.0
def validate(self, method: CapMethod):
if method == "alsc":
if self.cap <= 0:
raise ValueError("ALSC requires cap > 0")
if self.alpha <= 0:
raise ValueError("ALSC requires alpha > 0")
elif method in ("tanh", "soft_sign"):
if self.cap <= 0:
raise ValueError(f"{method} requires cap > 0")
def _sigmoid(x: jax.Array) -> jax.Array:
return jax.nn.sigmoid(x)
def _alsc_forward(x: jax.Array, params: CapParams) -> jax.Array:
cap, alpha, beta = params.cap, params.alpha, params.beta
z = alpha * (x / cap + beta)
sigma_z = _sigmoid(z)
sigma_beta = _sigmoid(alpha * beta)
denom = sigma_beta * (1 - sigma_beta) + 1e-12
return cap * (sigma_z - sigma_beta) / denom
def _alsc_inverse(y: jax.Array, params: CapParams) -> jax.Array:
cap, alpha, beta = params.cap, params.alpha, params.beta
sigma_beta = _sigmoid(alpha * beta)
denom = sigma_beta * (1 - sigma_beta) + 1e-12
sigma_z = y / cap * denom + sigma_beta
sigma_z = jnp.clip(sigma_z, 1e-7, 1 - 1e-7)
z = jnp.log(sigma_z / (1 - sigma_z))
return cap * (z / alpha - beta)
def _alsc_grad(x: jax.Array, params: CapParams) -> jax.Array:
cap, alpha, beta = params.cap, params.alpha, params.beta
z = alpha * (x / cap + beta)
sigma_z = _sigmoid(z)
sigma_beta = _sigmoid(alpha * beta)
denom = sigma_beta * (1 - sigma_beta) + 1e-12
return (alpha / cap) * sigma_z * (1 - sigma_z) / denom
def cap_forward(qk: jax.Array, method: CapMethod, params: CapParams) -> jax.Array:
if method == "none" or params.cap <= 0:
return qk
elif method == "tanh":
return params.cap * jnp.tanh(qk / params.cap)
elif method == "soft_sign":
return qk / (1 + jnp.abs(qk) / params.cap)
elif method == "alsc":
return _alsc_forward(qk, params)
else:
raise ValueError(f"Unknown cap method: {method}")
def cap_grad(qk: jax.Array, method: CapMethod, params: CapParams,
qk_capped: Optional[jax.Array] = None) -> jax.Array:
if method == "none" or params.cap <= 0:
return jnp.ones_like(qk)
elif method == "tanh":
if qk_capped is not None:
return 1 - (qk_capped / params.cap) ** 2
return 1 - jnp.tanh(qk / params.cap) ** 2
elif method == "soft_sign":
return 1 / (1 + jnp.abs(qk) / params.cap) ** 2
elif method == "alsc":
return _alsc_grad(qk, params)
else:
raise ValueError(f"Unknown cap method: {method}")