JAX 的 SciPy 兼容模块 jax.scipy 完全指南:从特殊函数到稀疏线性代数
【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax
导读
jax.scipy是 JAX 中对 SciPy 科学计算栈的差异化重实现,将 SciPy 经典 API(线性代数、FFT、信号处理、统计分布、稀疏迭代求解器等)全部迁移到 JAX 的可微、可向量化、可 JIT 的jax.Array体系之上。阅读本文你将掌握jax.scipy的全部子模块布局与函数清单、关键 API 的参数语义与源码实现原理、以及与jax.numpy、jax.lax的协作方式,从而在科学计算与机器学习混合场景中写出既贴近 SciPy 习惯又能享受自动微分与 GPU/TPU 加速的代码。
一、jax.scipy 模块总览与设计定位
jax.scipy的目标是"API 兼容 SciPy,底层对接 JAX 运行时":调用方按 SciPy 的书写习惯组织代码,实际执行却发生在 JAX 的 tracing 与 XLA 编译管线中,因此天然支持jax.jit、jax.grad、jax.vmap等组合变换。
从 jax/scipy/init.py 可以看出,模块采用**懒加载(lazy loading)**机制:通过jax._src.lazy_loader.attach挂载了 10 个顶层子模块,即
interpolate(插值)linalg(线性代数)ndimage(图像处理)signal(信号处理)sparse(稀疏矩阵)special(特殊函数)stats(统计分布)fft(离散余弦变换)cluster(聚类)integrate(数值积分)
TYPE_CHECKING分支中的显式导入(如from jax.scipy import linalg as linalg)保证了类型检查器能正确解析命名空间,这也对应了仓库注释中强调的 "PEP 484 要求import <name> as <name>才能导出符号" 的约定。
从实现层次看,公开 API 大多定义在 jax/_src/scipy/ 下的同名文件中(如 jax/_src/scipy/linalg.py、jax/_src/scipy/special.py),公共层只是做 re-export;少量函数来自jax._src.third_party.scipy(例如linalg.funm、special.fresnel),这些是从上游 SciPy 移植的独立实现。
二、jax.scipy.fft:离散余弦变换族
jax.scipy.fft当前提供DCT(离散余弦变换)及其逆变换的 1D / N-D 版本,共 4 个函数:dct、dctn、idct、idctn。
dct / dctn 的参数语义
以 jax/_src/scipy/fft.py 中的dct为例,其完整签名如下:
dct(x, type=2, n=None, axis=-1, norm=None)| 参数 | 类型 | 默认值 | 语义 |
|---|---|---|---|
x | 数组 | 必填 | 输入数据,支持实数与复数 |
type | int | 2 | 变换类型,当前仅支持 type=2,传入其他值抛出NotImplementedError |
n | int | x.shape[axis] | 变换长度;大于输入长度时零填充(lax.pad),小于输入长度时截断 |
axis | int | -1 | 沿哪个轴做变换,经canonicalize_axis归一化 |
norm | str | None | 归一化模式,取值None/"backward"/"ortho",默认等价于"backward";"forward"未实现 |
dctn额外提供s(结果形状)与axes(变换轴序列)参数:axes缺省时使用最后len(s)个轴;两者都缺省时沿全部轴变换。高维实现由 1D/2D DCT 组合而成。
源码级实现原理
从 jax/_src/scipy/fft.py 可以看到,DCT 并非直接实现余弦求和,而是复用 FFT 完成(注释引用了 John Makhoul 1980 年的经典论文A Fast Cosine Transform in One and Two Dimensions):
_dct_interleave将输入按奇偶下标拆开、翻转并拼接,构造出适合 FFT 的交错序列;- 调用
jnp_fft.fft计算快速傅里叶变换; - 乘以旋转因子
_W4(N, k) = exp(-.5j * π * k / N)并取实部乘以 2; norm="ortho"时通过_dct_ortho_norm施加正交归一化因子。
复数输入会被拆成实部/虚部分别变换后再用lax.complex重组。idct是dct的逆过程:先在频域除以旋转因子、乘以2N,再ifft后经_dct_deinterleave还原奇偶位。因此文档示例中jnp.allclose(x, idct(dct(x)))返回Array(True, dtype=bool),验证了正逆变换的闭环性质。
三、jax.scipy.integrate:梯形法则数值积分
jax.scipy.integrate目前只暴露一个函数trapezoid(SciPy 1.6+ 中trapz的替代名),实现复合梯形法则:
trapezoid(y, x=None, dx=1.0, axis=-1)y:被积分数据数组;x:采样点坐标,缺省时按dx等间距分布;dx:x缺省时的采样间距,默认1.0;axis:积分轴,默认最后一轴。
从 jax/_src/scipy/integrate.py 的源码看,trapezoid被@jit(static_argnames=('axis',))装饰,并直接委托给jax.numpy.trapezoid(即lax_numpy.trapezoid)——这是jax.scipy与jax.numpy底层复用的典型例子。文档中的数值示例:对y = [1,2,3,2,3,2,1]以dx=1.0积分得到13.0;用不规则网格x = [0,2,5,7,10,15,20]积分得到43.0;对sin²在[0, 2π]上 1000 点积分,结果与π在浮点精度内一致(allclose返回True)。
四、jax.scipy.linalg:线性代数工具箱
jax.scipy.linalg是子模块中体量最大的部分,提供34 个函数,完整清单见 jax/scipy/linalg.py,全部从jax._src.scipy.linalgre-export(唯一例外是funm,来自jax._src.third_party.scipy.linalg):
分解类:lu、lu_factor、lu_solve、qr、qr_multiply、cholesky、cho_factor、cho_solve、svd、eigh、eigh_tridiagonal、schur、rsf2csf、hessenberg、polar
求解类:solve、solve_triangular、solve_sylvester、inv、det
矩阵函数:expm、expm_frechet、sqrtm、funm
结构矩阵构造器:block_diag、circulant、companion、fiedler、fiedler_companion、convolution_matrix、hadamard、hankel、helmert、hilbert、invhilbert、invpascal、leslie、pascal、toeplitz、dft
其中dft用于构造离散傅里叶变换矩阵,hilbert/invhilbert生成病态条件数著名的 Hilbert 矩阵(常用于数值稳定性测试),toeplitz/hankel/circulant等服务于卷积与 Toeplitz 系统。所有这些操作都基于jax.numpy.linalg与lax原语实现,因而可被自动微分(如对svd、eigh求导)并可嵌入jit计算图。
五、jax.scipy.signal:信号处理
jax.scipy.signal提供 10 个函数,覆盖卷积、相关与频谱分析:
- 卷积/相关:
fftconvolve、convolve、convolve2d、correlate、correlate2d - 时频分析:
stft(短时傅里叶变换)、istft(逆 STFT)、welch(Welch 功率谱估计)、csd(互谱密度) - 预处理:
detrend(去除趋势)
以 jax/_src/scipy/signal.py 中的fftconvolve为例,签名与参数为:
fftconvolve(in1, in2, mode="full", axes=None)in1/in2:两个输入数组,要求in1.ndim == in2.ndim;mode:输出尺寸控制,三选一——"full"(默认,完整卷积)、"same"(输出与in1同尺寸的中心部分)、"valid"(仅保留不依赖边缘填充的部分);axes:可选,指定沿哪些轴做卷积。
fftconvolve内部依赖jax._src.numpy.fft,将时域卷积转为频域乘法再逆变换,适合大卷积核场景;convolve系列则提供直接(非 FFT)实现,参数与 NumPy/SciPy 语义对齐。源码开头的ModeString类型别名(Literal["full", "same", "valid"])也以类型标注形式固化了mode的合法取值。
六、jax.scipy.sparse.linalg:稀疏迭代求解器
jax.scipy.sparse.linalg提供 3 个 Krylov 子空间迭代求解器:cg(共轭梯度)、gmres(广义最小残差)、bicgstab(双共轭梯度稳定法)。它们用于求解A x = b形式的大规模稀疏线性系统,尤其适合矩阵以算子形式给出、无需显式存储的场景。
从 jax/_src/scipy/sparse/linalg.py 源码可以看出实现的关键设计:
- 内部以pytree 抽象处理未知量
x:_vdot_tree、_norm、_mul、_dot_tree等辅助函数对 pytree 逐叶做内积、范数与标量乘,因此x可以是任意嵌套结构(数组、字典、列表)而不仅是单个向量; - 内积运算统一使用
precision=lax.Precision.HIGHEST(见_dot = partial(jnp.dot, precision=lax.Precision.HIGHEST)),并专门实现了_vdot_real_part——对复数输入只保留实部内积以保证z^H M z的实值性,这是 CG 类算法收敛判据正确性的前提; - 代码中出现的
tree_util.Partial用于携带操作符函数,说明A既可以传矩阵也可以传"黑盒线性算子"(callable),这与 SciPy 的LinearOperator用法对应。
这些求解器全部可 JIT 编译,因此适用于在jit内部完成内层线性求解的算法(如 Gauss-Newton 步、隐式时间积分等)。
七、jax.scipy.special:特殊函数库
jax.scipy.special是数值计算中最常用的子模块,从 jax/scipy/special.py 可见其完整导出清单共50 余个函数,其中大部分来自jax._src.scipy.special,fresnel来自jax._src.third_party.scipy.special。按用途可归纳为:
概率/统计相关:ndtr(正态 CDF)、ndtri(正态分位数)、log_ndtr、erf、erfc、erfcx、erfinv、gammainc、gammaincc、gammaln、gammasgn、digamma、polygamma、beta、betainc、betaln、multigammaln、loggamma、owens_t、logit、expit
软最大/KL 散度:softmax、log_softmax、logsumexp、kl_div、rel_entr、entr、xlogy、xlog1py
组合与计数:comb、factorial、poch、bernoulli(伯努利数)
贝塞尔/超越函数:i0、i0e、i1、i1e(修正贝塞尔函数)、sph_harm_y(球谐函数)、hyp1f1、hyp2f1(超几何函数)、wofz(Faddeeva 函数)、dawsn、fresnel、sici、exp1、expi、expn、spence、zeta、boxcox、boxcox1p、logit等
退化/移除的函数:lpmn与lpmn_values(连带勒让德函数)已被标记为弃用——在 jax/scipy/special.py 的_deprecations表中,二者于 2024 年 1 月加入弃用名单,提示语为 "lpmn is deprecated; no replacement is planned",访问时会触发deprecation_getattr警告,但为了向后兼容仍可通过__getattr__拿到旧实现。
得益于这些函数全部由 JAX 原语构建,jax.grad(softmax)、jax.jit(gammaln)等组合可以直接使用,这是相比直接调用 SciPy 数值库的最大优势。
八、jax.scipy.stats:概率分布家族
jax.scipy.stats是覆盖最广的子模块,从 jax/scipy/stats/init.py 可见共26 个分布模块与 3 个通用统计函数。
分布清单与可用方法
每个分布(如norm、beta、gamma、poisson)提供 PDF/PMF、CDF、分位数、生存函数等方法的子集。文档索引中列出的分布及方法可归纳为:
- 连续分布(pdf / logpdf / cdf / logcdf / sf / logsf / ppf / isf):
norm(完整八件套)、cauchy、gumbel_l、gumbel_r、pareto、truncnorm、uniform、laplace(cdf/logpdf/pdf)、logistic(cdf/isf/logpdf/pdf/ppf/sf)、expon、gamma、chi2、beta、vonmises、wrapcauchy、gennorm、t(logpdf/pdf)、dirichlet、multivariate_normal(logpdf/pdf) - 离散分布(logpmf / pmf):
bernoulli(另有 cdf/ppf)、binom、betabinom、geom、multinomial、nbinom、poisson(另有 cdf/entropy) - 通用统计量:
mode(众数)、rankdata(秩变换)、sem(标准误),定义于 jax/_src/scipy/stats/_core.py - 核密度估计:
gaussian_kde,定义于 jax/_src/scipy/stats/kde.py,提供evaluate、pdf、logpdf、resample、integrate_gaussian、integrate_box_1d、integrate_kde等方法
使用要点
各分布模块独立成文件(如 jax/scipy/stats/norm.py、jax/scipy/stats/poisson.py),可通过jax.scipy.stats.norm.pdf(x, loc, scale)这类 SciPy 风格调用,同时全部基于jax.lax与special函数实现,因此概率计算可 JIT、可求导(如对pdf的对数似然求梯度以做最大似然估计)。gaussian_kde的resample依赖 JAX 的随机数键(jax.random),与 JAX 的显式 RNG 体系保持一致。
九、其余子模块:cluster、interpolate、ndimage、optimize、spatial
jax.scipy.cluster:向量量化
vq(vector quantization)将观测向量映射到码本中最接近的码字,对应 SciPy 的cluster.vq.vq,实现在 jax/_src/scipy/cluster/vq.py。jax.scipy.cluster.vq常用于 K-Means 的量化步骤,可 JIT 化以加速大规模码本分配。
jax.scipy.interpolate:规则网格插值
RegularGridInterpolator提供 N 维规则网格上的插值器,支持在任意查询点处取值。它的核心价值在于可自动微分:在物理模拟或可微渲染中,将网格场(如速度场、密度场)插值到粒子位置的操作可以被jax.grad反向传播。
jax.scipy.ndimage:坐标映射采样
map_coordinates按给定坐标对 N 维数组做采样(样条/线性插值),是图像变形、可微空间变换网络(STN)的基石。与interpolate一样,它对坐标的梯度可自然传递。
jax.scipy.optimize:无约束优化
提供minimize与OptimizeResults,实现在 jax/_src/scipy/optimize/(含bfgs.py、_lbfgs.py、line_search.py、minimize.py)。minimize(fun, x0, method=...)支持 BFGS / L-BFGS 等方法,且目标函数梯度默认由 JAX 自动微分提供(或可通过jax.value_and_grad组合),OptimizeResults对象携带最优值、迭代信息等结果字段。这是 JAX 场景中不依赖外部 SciPy 的常用优化入口。
jax.scipy.spatial.transform:旋转与插值
Rotation:3D 旋转对象,支持矩阵、四元数、欧拉角、旋转向量等多种表示之间的转换;Slerp:球面线性插值,用于旋转的平滑过渡,在机器人运动规划、骨骼动画插帧中很实用。
十、与 jax.numpy / jax.lax 的关系及使用边界
jax.scipy并不是一个孤立命名空间,它与 JAX 底层体系紧密耦合:
- 复用
jax.numpy实现:如integrate.trapezoid直接委托jax.numpy.trapezoid;fft.dct复用jax._src.numpy.fft;sparse.linalg用jnp.dot/jnp.vdot/einsum搭建内积与矩阵乘法。 - 构建在
lax原语之上:dct中的lax.slice_in_dim、lax.rev、lax.pad、lax.expand_dims,sparse.linalg中的lax.Precision.HIGHEST,都是直接调用 XLA 级原语,保证算子可编译、可求导。 - 与
jax.random协作:stats.gaussian_kde.resample依赖显式随机键,遵循 JAX "显式 PRNG" 的全局约定。
使用边界(以当前仓库源码为准):
fft.dct/dctn/idct/idctn的type参数仅支持 2,其他类型抛出NotImplementedError;norm仅支持None/"backward"/"ortho",传入"forward"会抛ValueError;special.lpmn、special.lpmn_values已弃用且无替代方案,新代码应避免使用;scipy.signal的fftconvolve是 FFT 近似实现,浮点结果可能与直接卷积存在微小差异(源码 docstring 也提示使用jnp.printoptions调整打印精度);- 各子模块覆盖范围是 SciPy 的子集,若需要 SciPy 更完整的生态能力(如
scipy.optimize的全部方法),应结合 JAX 的转换能力自行组合或在 Python 回调边界调用原始 SciPy。
十一、快速上手示例
下面组合jax.scipy的多个子模块,展示其在 JAX 可微管线中的典型用法:
import jax import jax.numpy as jnp import jax.scipy as jsp # 1) 统计:正态分布对数似然(可 JIT + 可 grad) key = jax.random.key(0) data = jax.random.normal(key, (1000,)) def neg_loglik(params): loc, scale = params return -jsp.stats.norm.logpdf(data, loc, scale).sum() print(jax.grad(neg_loglik)((0.0, 1.0))) # 2) FFT:DCT 正逆变换闭环 x = jax.random.normal(key, (8,)) assert jnp.allclose(x, jsp.fft.idct(jsp.fft.dct(x))) # 3) 积分:梯形法则计算定积分 xs = jnp.linspace(0, 2 * jnp.pi, 1000) integral = jsp.integrate.trapezoid(jnp.sin(xs) ** 2, xs) assert jnp.allclose(integral, jnp.pi) # 4) 稀疏求解:CG 解 A x = b A = jnp.eye(4) * 2 + jnp.ones((4, 4)) b = jnp.arange(4, dtype=jnp.float32) x = jsp.sparse.linalg.cg(A, b, maxiter=100)[0] assert jnp.allclose(A @ x, b, atol=1e-4) # 5) 优化:BFGS 最小化 Rosenbrock 函数 result = jsp.optimize.minimize( lambda v: (1 - v[0])**2 + 100 * (v[1] - v[0]**2)**2, jnp.array([0.0, 0.0]), method="BFGS") print(result.x)结语:何时选择 jax.scipy
当你的代码同时需要"SciPy 的数学语义"与"JAX 的自动微分/向量化/JIT"时,jax.scipy是最直接的桥接层:special提供可微特殊函数,stats提供可微概率分布,linalg与sparse.linalg提供可编译的稠密/稀疏求解,fft、signal、ndimage、interpolate覆盖信号与图像处理,optimize与cluster补齐经典数值算法。使用时注意各函数的实现边界(如 DCT 仅 type-2、lpmn已弃用),并善用 jax/_src/scipy/ 下的源码与 docs/jax.scipy.rst 中的完整 API 索引按需检索。
【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考