# # 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}")