Upload fractus/nn/stats.py with huggingface_hub
Browse files- fractus/nn/stats.py +45 -0
fractus/nn/stats.py
ADDED
|
@@ -0,0 +1,45 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Numerical utilities for fractus.
|
| 2 |
+
|
| 3 |
+
Ported from the original system (src/math/stats.rs) in pure PyTorch, differentiable.
|
| 4 |
+
|
| 5 |
+
elu_plus_one : strictly positive feature map for linear attention.
|
| 6 |
+
φ(x, α) = x + 1 if x > 0
|
| 7 |
+
= α(e^x - 1) + 1 otherwise
|
| 8 |
+
With α=1 (default), φ is strictly positive (min e^x > 0 for x→-∞,
|
| 9 |
+
= 1 at x=0). This positivity guarantees that the denominator of causal
|
| 10 |
+
linear attention stays well-defined.
|
| 11 |
+
|
| 12 |
+
stable_softmax : softmax with max subtraction (no overflow).
|
| 13 |
+
"""
|
| 14 |
+
|
| 15 |
+
import torch
|
| 16 |
+
|
| 17 |
+
|
| 18 |
+
def elu_plus_one(x: torch.Tensor, alpha: float = 1.0) -> torch.Tensor:
|
| 19 |
+
"""ELU+1 strictly positive feature map, differentiable.
|
| 20 |
+
|
| 21 |
+
Args:
|
| 22 |
+
x : tensor of arbitrary shape.
|
| 23 |
+
alpha : ELU coefficient (1.0 by default, as in the original).
|
| 24 |
+
Returns:
|
| 25 |
+
tensor of the same shape, strictly positive.
|
| 26 |
+
"""
|
| 27 |
+
# We use the direct formula (differentiable via torch.where):
|
| 28 |
+
# positive branch: x + 1; negative branch: alpha * (exp(x) - 1) + 1.
|
| 29 |
+
pos = x + 1.0
|
| 30 |
+
neg = alpha * (torch.exp(x) - 1.0) + 1.0
|
| 31 |
+
return torch.where(x > 0, pos, neg)
|
| 32 |
+
|
| 33 |
+
|
| 34 |
+
def stable_softmax(logits: torch.Tensor, dim: int = -1) -> torch.Tensor:
|
| 35 |
+
"""Numerically stable softmax (max subtraction).
|
| 36 |
+
|
| 37 |
+
If the exponential sum is < 1e-10, returns the uniform 1/N
|
| 38 |
+
(limit behavior inherited from the original stats.rs:56-57).
|
| 39 |
+
"""
|
| 40 |
+
max_logits, _ = logits.max(dim=dim, keepdim=True)
|
| 41 |
+
exp = torch.exp(logits - max_logits)
|
| 42 |
+
denom = exp.sum(dim=dim, keepdim=True)
|
| 43 |
+
# Limit behavior: uniform if denom ~ 0.
|
| 44 |
+
uniform = torch.full_like(exp, 1.0 / exp.shape[dim])
|
| 45 |
+
return torch.where(denom > 1e-10, exp / denom, uniform)
|