PyTorch torch.special 模块全解析:从 SciPy 对齐的特殊函数到张量化实现
【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch
torch.special 是 PyTorch 中面向科学计算与统计建模的特殊函数模块,其 API 设计以 SciPy 的 scipy.special 为参照,为贝塞尔函数、误差函数、伽马族函数、正交多项式与概率分布辅助函数提供统一的张量化入口。本文以 docs/source/special.md 为骨架,结合 torch/special/init.py 的完整实现与底层算子注册、测试用例,带你掌握该模块的全部函数清单、数值语义、用法示例与源码级原理,可直接用于概率编程、信号处理、物理模拟等实战场景。
模块定位:PyTorch 中的 SciPy special 等价物
特殊函数(Special Functions)是一组在数学物理、统计推断与数值算法中反复出现的高阶函数,例如贝塞尔函数、伽马函数、误差函数等。SciPy 社区将这些函数统一收纳在scipy.special命名空间下,而 PyTorch 以相同理念提供了torch.special:
- 统一的张量接口:所有函数接受 Tensor(或 Tensor 与标量的混合)输入,支持自动求导、GPU 加速与广播语义;
- 与 SciPy 对齐的命名与行为:文档原文明确指出该模块"modeled after SciPy's special module"(以 SciPy 的 special 模块为模型),函数命名(如
ndtr、gammainc、xlogy)与 SciPy 保持一致,便于从 NumPy/SciPy 工作流平滑迁移; - Alias 语义:部分函数是 PyTorch 既有算子的别名,部分则是独立的原生算子。
从源码看,模块的对外导出集中在 torch/special/init.py 的__all__列表中,共 56 个公开函数,每个函数通过_add_docstr将底层 C++ 算子(torch._C._special命名空间下的special_*)包装为带完整文档的 Python 入口:
from torch._C import _add_docstr, _special psi = _add_docstr( _special.special_psi, r"""psi(input, *, out=None) -> Tensor ...""", )这种"Python 文档层 + C++ 算子层"的两段式结构,意味着你看到的每个函数在 ATen 层都有对应的原生实现(注册于 aten/src/ATen/native/native_functions.yaml,例如special_ndtr、special_gammainc、special_multigammaln分别注册在第 13700、13947、13959 行附近),并配有 CPU 与 CUDA 两套 kernel(如 aten/src/ATen/native/cpu/UnaryOpsKernel.cpp、aten/src/ATen/native/cuda/UnarySpecialOpsKernel.cu)。
函数全景:56 个 API 按数学家族分类
原文档以autofunction指令列出了全部函数,这里按数学归属整理成速查表,便于按需定位:
| 家族 | 函数 | 典型用途 |
|---|---|---|
| 伽马与多伽马 | gammaln、digamma、polygamma、psi、multigammaln、gammainc、gammaincc | 概率密度归一化、分布矩、多变量统计 |
| 误差函数族 | erf、erfc、erfcx、erfinv | 高斯积分、Q 函数、反误差求解 |
| 高斯分布辅助 | ndtr、ndtri、log_ndtr | 正态 CDF、分位数函数、log-CDF |
| 贝塞尔函数 | bessel_j0、bessel_j1、bessel_y0、bessel_y1、modified_bessel_i0、modified_bessel_i1、modified_bessel_k0、modified_bessel_k1、scaled_modified_bessel_k0、scaled_modified_bessel_k1、spherical_bessel_j0 | 波动方程、扩散问题、天线与声学 |
| 艾里函数 | airy_ai | 量子力学 WKB、光学衍射 |
| 正交多项式 | chebyshev_polynomial_t/u/v/w、shifted_chebyshev_polynomial_t/u/v/w、hermite_polynomial_h/he、laguerre_polynomial_l、legendre_polynomial_p | 数值逼近、谱方法、积分求积 |
| 对数与熵相关 | log1p、expm1、exp2、expit、logit、entr、logsumexp、softmax、log_softmax、sinc、round | 数值稳定性处理、激活函数、信息熵 |
| 双参特殊运算 | zeta、xlogy、xlog1py | Hurwitz ζ 函数、鲁棒乘积对数 |
其中psi、log1p、round、logsumexp是显式的别名:psi等价于digamma(torch/special/init.py 中注明 "Alias for torch.special.digamma"),logsumexp等价于torch.logsumexp,log1p等价于torch.log1p,round等价于torch.round。
常用函数深入讲解
3.1 误差函数族:erf / erfc / erfcx / erfinv
误差函数族是概率论与偏微分方程中最常见的特殊函数:
erf(x) = (2/√π)∫₀ˣ e^(-t²) dt,用于计算正态分布的累积概率区间;erfc(x) = 1 - erf(x),即互补误差函数,常用于通信系统的误码率(BER)计算;erfcx(x) = e^(x²)·erfc(x),缩放版互补误差函数,避免大正 x 时erfc下溢;erfinv在 (-1, 1) 区间内满足erfinv(erf(x)) = x,是反误差求解的核心。
源码中定义与示例(torch/special/init.py):
>>> torch.special.erf(torch.tensor([0, -1., 10.])) tensor([ 0.0000, -0.8427, 1.0000]) >>> torch.special.erfc(torch.tensor([0, -1., 10.])) tensor([ 1.0000, 1.8427, 0.0000]) >>> torch.special.erfcx(torch.tensor([0, -1., 10.])) tensor([ 1.0000, 5.0090, 0.0561]) >>> torch.special.erfinv(torch.tensor([0, 0.5, -1.])) tensor([ 0.0000, 0.4769, -inf])注意erfinv(-1)返回-inf,这与定义域边界行为一致。
3.2 高斯分布辅助函数:ndtr / ndtri / log_ndtr
这三个函数让 PyTorch 在原生张量层面直接拥有正态分布的 CDF、分位数与 log-CDF:
ndtr(x)计算标准正态 PDF 从 -∞ 到 x 的积分,即正态分布累积分布函数;ndtri(p)是其反函数,即正态分布的分位数函数(quantile function),满足ndtri(p) = √2·erf⁻¹(2p - 1);log_ndtr(x)计算log(ndtr(x)),用于概率的 log-space 运算,避免概率下溢。
>>> torch.special.ndtr(torch.tensor([-3., -2, -1, 0, 1, 2, 3])) tensor([0.0013, 0.0228, 0.1587, 0.5000, 0.8413, 0.9772, 0.9987]) >>> torch.special.ndtri(torch.tensor([0, 0.25, 0.5, 0.75, 1])) tensor([ -inf, -0.6745, 0.0000, 0.6745, inf]) >>> torch.special.log_ndtr(torch.tensor([-3., -2, -1, 0, 1, 2, 3])) tensor([-6.6077 -3.7832 -1.841 -0.6931 -0.1728 -0.023 -0.0014])这三个函数在 test/test_unary_ufuncs.py 中有专门的 SciPy 对拍测试:test_special_ndtr_vs_scipy与test_special_log_ndtr_vs_scipy使用torch.linspace(-10, 10, ...)及 dtype 的min/max/eps/tiny极值,将 PyTorch 结果与scipy.special.ndtr/log_ndtr逐一比对,验证数值等价性。这从测试层面印证了"与 SciPy 对齐"的设计目标。
3.3 伽马函数族:gammaln / digamma / psi / polygamma / multigammaln / gammainc / gammaincc
伽马族是贝叶斯统计与分布计算的基础:
gammaln计算ln|Γ(x)|,以对数域规避阶乘爆炸(torch/special/init.py);digamma(x) = Γ'(x)/Γ(x)是伽马函数的对数导数,psi是它的别名;polygamma(n, x)计算 digamma 的第 n 阶导数,n 必须为非负整数(torch/special/init.py);multigammaln(a, p)计算维数为 p 的多变量 log-伽马函数,其中常数项C = log(π)·p(p-1)/4,并要求所有元素大于 (p-1)/2,否则行为未定义(torch/special/init.py);gammainc/gammaincc分别计算正则化下/上不完全伽马函数,两者对同一对输入之和恒为 1,可互为校验(torch/special/init.py)。
>>> torch.special.gammaln(torch.arange(0.5, 2, 0.5)) tensor([ 0.5724, 0.0000, -0.1208]) >>> torch.special.digamma(torch.tensor([1, 0.5])) tensor([-0.5772, -1.9635]) >>> torch.special.polygamma(1, torch.tensor([1, 0.5])) tensor([1.64493, 4.9348]) >>> a1 = torch.tensor([4.0]); a2 = torch.tensor([3.0, 4.0, 5.0]) >>> torch.special.gammainc(a1, a2) + torch.special.gammaincc(a1, a2) tensor([1., 1., 1.])两点需要留意:
digamma在 0 处的行为:从 PyTorch 1.8 起返回-Inf(此前返回 NaN),这是文档明确记录的破坏性变更;- 反向传播限制:
gammainc/gammaincc目前不支持对第一个输入(a)的反向传播,文档明确提示"backward pass with respect to input is not yet supported"。这意味着这两个函数可安全参与对第二参数(x)的求导,但对第一参数求梯度会失败,设计损失函数时需规避。
3.4 贝塞尔函数族:J / Y / I / K 与缩放版本
PyTorch 覆盖了四类贝塞尔函数的前两阶:
- 第一类:
bessel_j0、bessel_j1(J₀、J₁); - 第二类:
bessel_y0、bessel_y1(Y₀、Y₁); - 修正第一类:
modified_bessel_i0、modified_bessel_i1(I₀、I₁,即i0、i1); - 修正第二类:
modified_bessel_k0、modified_bessel_k1(K₀、K₁); - 缩放修正第二类:
scaled_modified_bessel_k0、scaled_modified_bessel_k1,通过指数缩放避免大参数下的溢出; - 球贝塞尔:
spherical_bessel_j0。
i0的幂级数定义为I₀(x) = Σₖ (x²/4)ᵏ/(k!)²(torch/special/init.py),i0e(x) = exp(-|x|)·i0(x)是其指数缩放版(torch/special/init.py),其余i1/i1e类似。缩放版本的引入正是为了保证数值稳定性——当 |x| 很大时,i0会以指数速度增长,先取对数尺度再缩放可保持中间结果在有限浮点范围内。
>>> torch.special.i0(torch.arange(5, dtype=torch.float32)) tensor([ 1.0000, 1.2661, 2.2796, 4.8808, 11.3019]) >>> torch.special.i0e(torch.arange(5, dtype=torch.float32)) tensor([1.0000, 0.4658, 0.3085, 0.2430, 0.2070])3.5 正交多项式:Chebyshev / Hermite / Laguerre / Legendre
模块内置四族经典正交多项式,均以(input, n)双参数调用,n为张量形式的次数:
- Chebyshev 第一至第四类:
chebyshev_polynomial_t/u/v/w,以及定义在 [0,1] 区间的shifted_chebyshev_polynomial_t/u/v/w; - Hermite:物理学家版本
hermite_polynomial_h(Hₙ)与概率论版本hermite_polynomial_he(Heₙ); - Laguerre:
laguerre_polynomial_l(Lₙ); - Legendre:
legendre_polynomial_p(Pₙ)。
源码文档详细刻画了数值求值策略。以chebyshev_polynomial_t为例(torch/special/init.py):
- n = 0 时返回 1;n = 1 时返回 x;
- 当 n < 6 或 |x| > 1 时,使用三项递推
Tₙ₊₁(x) = 2x·Tₙ(x) - Tₙ₋₁(x); - 否则改用显式三角公式
Tₙ(x) = cos(n·arccos(x))。
chebyshev_polynomial_u采用相同的递推骨架,但在 |x| ≤ 1 且 n ≥ 6 时切换为sin((n+1)·arccos(x))/sin(arccos(x))。Hermite、Laguerre、Legendre 同样基于递推实现。这种"低阶用递推、高阶用闭式公式"的分段策略,本质上是在浮点误差累积(递推)与三角精度(公式)之间做权衡:递推在阶数高时会累积舍入误差,而显式公式在定义域内更稳定。
3.6 数值稳定性工具:logit / expit / xlogy / xlog1py / logsumexp / entr
这一组函数专注于"大数运算中的稳定性",是机器学习损失函数与概率计算的常客:
expit(x) = 1/(1+e^(-x)),即 logistic sigmoid,与logit互为反函数(torch/special/init.py);logit(input, eps=None):当eps给定(如 1e-6)时,先将输入裁剪到[eps, 1-eps]再计算ln(z/(1-z)),防止对数取到 0 或负值;当eps=None且输入超出 (0,1) 时产生 NaN(torch/special/init.py);xlogy(input, other)计算input·log(other),特例处理:other为 NaN 时输出 NaN,input为 0 时输出 0(避免 0·(-∞) = NaN),与 SciPy 的scipy.special.xlogy行为一致(torch/special/init.py);xlog1py(input, other)计算input·log1p(other),将log(1+x)合并为数值上更稳的log1p(torch/special/init.py);entr按分段定义-x·ln(x)(x>0)、0(x=0)、-∞(x<0)逐元素计算熵(torch/special/init.py);logsumexp(input, dim, keepdim=False)是torch.logsumexp的别名,以max平移技巧保证数值稳定。
>>> torch.special.expit(torch.randn(4)) tensor([0.7153, 0.7481, 0.2920, 0.1458]) >>> x = torch.zeros(5,) >>> y = torch.tensor([-1, 0, 1, float('inf'), float('nan')]) >>> torch.special.xlogy(x, y) tensor([0., 0., 0., 0., nan]) >>> torch.special.xlog1py(x, y) tensor([0., 0., 0., 0., nan]) >>> torch.special.entr(torch.arange(-0.5, 1, 0.5)) tensor([ -inf, 0.0000, 0.3466])注意xlogy/xlog1py支持Tensor 与标量混用(xlogy(x, 4)、xlogy(2, y)均合法),但要求二者至少有一个是 Tensor,且自动应用广播语义。
3.7 其他实用函数:zeta / softmax / log_softmax / sinc / exp2 / expm1
zeta(input, other)计算 Hurwitz ζ 函数ζ(x, q) = Σₖ 1/(k+q)ˣ,当 q=1 时退化为 Riemann ζ 函数(torch/special/init.py)。该函数同样支持张量-标量混用与广播:
>>> x = torch.tensor([2., 4.]) >>> torch.special.zeta(x, 1) tensor([1.6449, 1.0823]) >>> torch.special.zeta(x, torch.tensor([1., 2.])) tensor([1.6449, 0.0823])softmax(input, dim, *, dtype=None)与log_softmax:按 dim 维归一化到 [0,1] 且和为 1;log_softmax直接在 log 域计算,数学上等价于log(softmax(x))但更慢的两次运算会产生精度损失,文档强调应直接使用本函数(torch/special/init.py)。二者均支持dtype参数,可先将输入提升精度再运算以防溢出;sinc计算归一化 sinc:x=0 时取 1,否则sin(πx)/(πx)(torch/special/init.py),是信号处理中的插值核;exp2(x) = 2ˣ、expm1(x) = eˣ - 1,后者对小 x 比exp(x)-1更精确(torch/special/init.py)。
实战用例:概率与信号处理场景组合
将上述函数组合即可完成端到端的统计计算。例如,构造一个带数值稳定的正态 log-CDF 与熵的统计管线:
import torch # 1. 标准正态分位数 → 生成指定分位的阈值 p = torch.tensor([0.025, 0.5, 0.975]) z = torch.special.ndtri(p) # tensor([-1.9600, 0.0000, 1.9600]) # 2. 大负数区域的 log-CDF(直接 log(ndtr) 会下溢,log_ndtr 保持稳定) x = torch.linspace(-20, -1, 5) log_cdf = torch.special.log_ndtr(x) # 3. 概率分布熵与伽马族:Beta 分布的对数归一化常数 a = torch.tensor([1.0, 2.0, 3.0]) b = torch.tensor([2.0, 2.0, 1.0]) log_beta = (torch.special.gammaln(a) + torch.special.gammaln(b) - torch.special.gammaln(a + b)) # 4. 通信误码率:利用 erfc 计算 Q 函数 Q(x) = 0.5 * erfc(x / sqrt(2)) snr = torch.tensor([0.0, 3.0, 6.0]) ber = 0.5 * torch.special.erfc(snr / torch.sqrt(torch.tensor(2.0))) # 5. 熵计算:entr 直接给出逐元素 -x ln x probs = torch.tensor([0.25, 0.25, 0.25, 0.25]) entropy = torch.special.entr(probs).sum()这类组合充分利用了模块"一次调用即完成逐元素数值计算 + 自动微分追踪"的特性,无需手动实现级数展开或查表。
源码结构导读
如果你想深入底层,可按以下路径继续探索:
- Python 包装层:torch/special/init.py —— 全部 56 个函数的文档、示例与
_add_docstr绑定,是理解 API 语义的第一站; - 算子注册:aten/src/ATen/native/native_functions.yaml —— 搜索
special_前缀可找到special_ndtr、special_gammainc、special_multigammaln等原生算子的 schema 声明; - CPU/CUDA kernel:aten/src/ATen/native/cpu/UnaryOpsKernel.cpp、aten/src/ATen/native/cuda/UnarySpecialOpsKernel.cu、aten/src/ATen/native/cpu/airy_ai.cpp —— 具体数值算法实现(如 airy_ai 独立成文件);
- 测试对拍:test/test_unary_ufuncs.py ——
test_special_ndtr_vs_scipy、test_special_log_ndtr_vs_scipy等用例,验证与 SciPy 的数值一致性。
使用注意事项总结
- 输入类型:绝大多数函数要求浮点张量输入;
softmax/log_softmax额外要求显式dim; - out 参数:几乎所有逐元素函数都支持
out=None可选输出张量(Keyword-only),避免额外分配; - 广播:双参数函数(
zeta、xlogy、xlog1py、gammainc、gammaincc)支持广播到公共形状,且至少一个参数必须是 Tensor; - 梯度限制:
gammainc/gammaincc对第一参数不支持反向传播; - 行为边界:
digamma(0) = -Inf(PyTorch 1.8+),erfinv(±1) = ±Inf,multigammaln要求元素 > (p-1)/2,polygamma的 n 必须为非负整数; - 别名等价:
psi ≡ digamma、logsumexp ≡ torch.logsumexp、log1p ≡ torch.log1p、round ≡ torch.round,按需选择即可。
以上内容完整覆盖了 docs/source/special.md 中列出的全部 56 个函数,并补充了源码层实现细节、SciPy 对拍测试证据与实战组合示例,可作为 PyTorch 特殊函数编程的速查手册使用。
【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考