Download python/xrex_unified/cap_functions.py from Snapkitty/ironic-mirror: direct link, hf CLI and curl.
- Browser
- Download file 3.6 kB
-
https://huggingface.co/Snapkitty/ironic-mirror/resolve/main/python/xrex_unified/cap_functions.py
- Command line
-
hf download hf://Snapkitty/ironic-mirror/python/xrex_unified/cap_functions.py
-
curl -L -o cap_functions.py https://huggingface.co/Snapkitty/ironic-mirror/resolve/main/python/xrex_unified/cap_functions.py
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"] | |
| 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}") | |