Skip to content

Repository files navigation

保角层次 Softmax(Conformal Hierarchical Softmax)

用保角(共形)变换改进 softmax 的计算结构,使复杂度接近 O(log n)。 完整数学推导见 softmax_logn.pdf(作者:魏勇勤,日期:2026 年 8 月 30 日)。

核心思想

  1. 保角视角:指数映射 exp: (R, +) → (R₊, ×) 是实数加群到正实数乘群的共形同构, softmax 恰为"共形提升 + 径向投影"的复合;
  2. 结合律:LSE(S∪T) = LSE(LSE(S), LSE(T)),把归一化计算组织成二叉树(每次合并 O(1));
  3. 结果(建树 O(n log n),一次性):
    • 全量输出的并行深度 O(log n);
    • 建树后单样本 p_i 查询 O(1);
    • top-k 查询 O(k log k);
    • 层次(路径)softmax 每样本 O(log n),且路径乘积精确等于 softmax 值;
    • 剪枝近似 softmax:相对误差 η ≤ δ 可保证,结构化(衰减/稀疏)数据上为次线性。

诚实声明:对任意输入,精确输出全部 n 个概率的下界为 Ω(n)(PDF 定理 1)。 "接近 O(log n)" 按上述五种精确含义成立(PDF 定义 1)。

文件清单

文件 说明
softmax_logn.py 核心实现
benchmark_softmax_logn.py 基准测试与实验图
softmax_logn.pdf 数学推导(作者 魏勇勤,2026 年 8 月 30 日)
softmax_logn.tex PDF 源文件(可重编译)
readme.md 本文档

安装

Python ≥ 3.9

pip install numpy scipy matplotlib

快速开始

import numpy as np
from softmax_logn import ConformalSoftmaxTree

x = np.array([1.0, 2.0, 0.5, -1.0])
tree = ConformalSoftmaxTree(x, split="conformal_gap")   # 或 split="balanced"

tree.all_probs()            # 全量精确 softmax,O(n)
tree.prob(1)                # 单样本查询,O(1)
tree.path_prob(1)           # 层次(望远镜)路径,O(log n),精确等于 softmax 值
idx, prob = tree.topk(2)    # top-2(按 logit 降序),O(k log k)

res = tree.approx_softmax(delta=1e-3)   # 剪枝近似:相对误差 η ≤ δ
res.kept_idx, res.kept_prob             # 保留位置与近似概率
res.eta_exact                           # 实际精确相对误差(等式,非上界)
res.visited                             # 访问节点数(代价度量)

API

class ConformalSoftmaxTree(x, split="balanced")

  • x:一维 logits(允许 -inf 表示掩码;不允许 NaN / +inf)
  • split
    • "balanced":按排序秩中位分割,深度 ≤ ⌈log₂n⌉;
    • "conformal_gap":平衡窗口内保角最大间隙二分(间隙 = 权重双曲距离),深度 ≤ log_{4/3} n,聚类更利于剪枝近似
  • 方法:
    • prob(i) → float:单样本精确 softmax,O(1)
    • path_prob(i) → float:层次路径乘积,O(log n),精确
    • topk(k) → (indices, probs):按 logit 降序,O(k log k)
    • all_probs() → ndarray:全量精确 softmax,O(n)
    • approx_softmax(delta=1e-3, max_iters=50, eps_min=1e-12) → ApproxResult
  • 属性:lse_root(LSE(x))、Z(配分函数,logit 极大时可为 inf)、depth、n、x

模块函数

  • build(x, split="balanced"):构建树
  • softmax_ref(x):数值稳定参考实现
  • mobius(x, a, b, c, d):实 Möbius 变换(PSL(2,R) 共形自同构群)
  • cross_ratio(a, b, c, d):四点交比(Möbius 不变量)

复杂度一览(详见 PDF 定理 1–7)

操作 复杂度 条件
建树(排序) O(n log n) 任意输入
建树 O(n) 工作,O(log n) 深度 logits 已排序 / 并行
全量输出 O(n) Ω(n) 下界
单样本 p_i O(1) 建树后
层次 softmax / 样本 O(log n) 建树后,精确
top-k O(k log k) 建树后
剪枝近似 O(K log n),K ≤ 1/ε η ≤ δ 保证;衰减/稀疏数据 K = O(log(1/ε))

与标准 softmax 的对比(n = 2²⁰,Python 3.13 实测)

操作 标准 softmax(scipy) 本方法(建树后)
全量输出一次 13.9 ms,O(n) 建树 4.2 s + 输出 6.1 ms
单样本 p_i 查询 每次重算 13.5 ms,O(n) ≈2.3 µs,O(1)
top-32 查询 argpartition 9.0 ms,O(n) ≈0.09 ms,O(k log k)
层次分类每样本 O(n) O(log n),精确
剪枝近似(长尾数据) 无此模式 次线性,η ≤ δ
内存 8 MB 64 MB(8 倍)

对照结论

  • 一次性输出整个向量:本方法无改进(总成本高约 300 倍),请用标准实现;
  • 单样本查询快约 6000 倍、top-32 快约 100 倍,且 O(1) / O(k log k) 对 O(n) 渐近占优;
  • 层次分类与带界剪枝近似是标准实现不具备的能力;
  • 内存约为输入的 8 倍,是主要代价;数值稳定性与精度两者持平(机器精度量级)。

盈亏平衡点(建树一次性约 4.2 s,n = 2²⁰):单样本查询约 300 次、top-32 查询约 500 次后开始划算;n 越大、查询越多,优势越大。

验证与基准

python softmax_logn.py            # 自检:精确性 / 望远镜 / top-k / 共形不变性
python benchmark_softmax_logn.py  # 基准 + 生成 fig_scaling.pdf / fig_approx.pdf / fig_depth.pdf

预期结果:全量输出与 scipy.special.softmax 的最大误差 ~1e-17(机器精度); 几何衰减数据(x_j = -j)上 n = 2²⁰ 的剪枝近似仅访问 49 个节点。

重新编译 PDF(可选)

需要 tectonic(仓库不附带二进制,请自行安装):

tectonic softmax_logn.tex

适用场景

  • 固定 logits 的反复查询(缓存 softmax、检索重排序)
  • 树型类别层次上的分类(每样本 O(log n) 精确层次 softmax)
  • 并行 / GPU 归约(并行深度 O(log n))
  • 长尾 / 稀疏 logits 的近似 softmax

作者与引用

魏勇勤(2026 年 8 月 30 日)。参考文献见 softmax_logn.pdf。

About

保角层次 Softmax:接近 O(log n) 的 Softmax 计算方法(论文 PDF/LaTeX + 实现 + 基准)

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages