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