JAX 的 SciPy 兼容模块 jax.scipy 完全指南:从特殊函数到稀疏线性代数
2026/9/10 2:08:31 网站建设 项目流程

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.numpyjax.lax的协作方式,从而在科学计算与机器学习混合场景中写出既贴近 SciPy 习惯又能享受自动微分与 GPU/TPU 加速的代码。

一、jax.scipy 模块总览与设计定位

jax.scipy的目标是"API 兼容 SciPy,底层对接 JAX 运行时":调用方按 SciPy 的书写习惯组织代码,实际执行却发生在 JAX 的 tracing 与 XLA 编译管线中,因此天然支持jax.jitjax.gradjax.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.funmspecial.fresnel),这些是从上游 SciPy 移植的独立实现。

二、jax.scipy.fft:离散余弦变换族

jax.scipy.fft当前提供DCT(离散余弦变换)及其逆变换的 1D / N-D 版本,共 4 个函数:dctdctnidctidctn

dct / dctn 的参数语义

以 jax/_src/scipy/fft.py 中的dct为例,其完整签名如下:

dct(x, type=2, n=None, axis=-1, norm=None)
参数类型默认值语义
x数组必填输入数据,支持实数与复数
typeint2变换类型,当前仅支持 type=2,传入其他值抛出NotImplementedError
nintx.shape[axis]变换长度;大于输入长度时零填充(lax.pad),小于输入长度时截断
axisint-1沿哪个轴做变换,经canonicalize_axis归一化
normstrNone归一化模式,取值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):

  1. _dct_interleave将输入按奇偶下标拆开、翻转并拼接,构造出适合 FFT 的交错序列;
  2. 调用jnp_fft.fft计算快速傅里叶变换;
  3. 乘以旋转因子_W4(N, k) = exp(-.5j * π * k / N)并取实部乘以 2;
  4. norm="ortho"时通过_dct_ortho_norm施加正交归一化因子。

复数输入会被拆成实部/虚部分别变换后再用lax.complex重组。idctdct的逆过程:先在频域除以旋转因子、乘以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等间距分布;
  • dxx缺省时的采样间距,默认1.0
  • axis:积分轴,默认最后一轴。

从 jax/_src/scipy/integrate.py 的源码看,trapezoid@jit(static_argnames=('axis',))装饰,并直接委托给jax.numpy.trapezoid(即lax_numpy.trapezoid)——这是jax.scipyjax.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):

分解类lulu_factorlu_solveqrqr_multiplycholeskycho_factorcho_solvesvdeigheigh_tridiagonalschurrsf2csfhessenbergpolar

求解类solvesolve_triangularsolve_sylvesterinvdet

矩阵函数expmexpm_frechetsqrtmfunm

结构矩阵构造器block_diagcirculantcompanionfiedlerfiedler_companionconvolution_matrixhadamardhankelhelmerthilbertinvhilbertinvpascallesliepascaltoeplitzdft

其中dft用于构造离散傅里叶变换矩阵,hilbert/invhilbert生成病态条件数著名的 Hilbert 矩阵(常用于数值稳定性测试),toeplitz/hankel/circulant等服务于卷积与 Toeplitz 系统。所有这些操作都基于jax.numpy.linalglax原语实现,因而可被自动微分(如对svdeigh求导)并可嵌入jit计算图。

五、jax.scipy.signal:信号处理

jax.scipy.signal提供 10 个函数,覆盖卷积、相关与频谱分析:

  • 卷积/相关fftconvolveconvolveconvolve2dcorrelatecorrelate2d
  • 时频分析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.specialfresnel来自jax._src.third_party.scipy.special。按用途可归纳为:

概率/统计相关ndtr(正态 CDF)、ndtri(正态分位数)、log_ndtrerferfcerfcxerfinvgammaincgammainccgammalngammasgndigammapolygammabetabetaincbetalnmultigammalnloggammaowens_tlogitexpit

软最大/KL 散度softmaxlog_softmaxlogsumexpkl_divrel_entrentrxlogyxlog1py

组合与计数combfactorialpochbernoulli(伯努利数)

贝塞尔/超越函数i0i0ei1i1e(修正贝塞尔函数)、sph_harm_y(球谐函数)、hyp1f1hyp2f1(超几何函数)、wofz(Faddeeva 函数)、dawsnfresnelsiciexp1expiexpnspencezetaboxcoxboxcox1plogit

退化/移除的函数lpmnlpmn_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 个通用统计函数。

分布清单与可用方法

每个分布(如normbetagammapoisson)提供 PDF/PMF、CDF、分位数、生存函数等方法的子集。文档索引中列出的分布及方法可归纳为:

  • 连续分布(pdf / logpdf / cdf / logcdf / sf / logsf / ppf / isf)norm(完整八件套)、cauchygumbel_lgumbel_rparetotruncnormuniformlaplace(cdf/logpdf/pdf)、logistic(cdf/isf/logpdf/pdf/ppf/sf)、expongammachi2betavonmiseswrapcauchygennormt(logpdf/pdf)、dirichletmultivariate_normal(logpdf/pdf)
  • 离散分布(logpmf / pmf)bernoulli(另有 cdf/ppf)、binombetabinomgeommultinomialnbinompoisson(另有 cdf/entropy)
  • 通用统计量mode(众数)、rankdata(秩变换)、sem(标准误),定义于 jax/_src/scipy/stats/_core.py
  • 核密度估计gaussian_kde,定义于 jax/_src/scipy/stats/kde.py,提供evaluatepdflogpdfresampleintegrate_gaussianintegrate_box_1dintegrate_kde等方法

使用要点

各分布模块独立成文件(如 jax/scipy/stats/norm.py、jax/scipy/stats/poisson.py),可通过jax.scipy.stats.norm.pdf(x, loc, scale)这类 SciPy 风格调用,同时全部基于jax.laxspecial函数实现,因此概率计算可 JIT、可求导(如对pdf的对数似然求梯度以做最大似然估计)。gaussian_kderesample依赖 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:无约束优化

提供minimizeOptimizeResults,实现在 jax/_src/scipy/optimize/(含bfgs.py_lbfgs.pyline_search.pyminimize.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 底层体系紧密耦合:

  1. 复用jax.numpy实现:如integrate.trapezoid直接委托jax.numpy.trapezoidfft.dct复用jax._src.numpy.fftsparse.linalgjnp.dot/jnp.vdot/einsum搭建内积与矩阵乘法。
  2. 构建在lax原语之上dct中的lax.slice_in_dimlax.revlax.padlax.expand_dimssparse.linalg中的lax.Precision.HIGHEST,都是直接调用 XLA 级原语,保证算子可编译、可求导。
  3. jax.random协作stats.gaussian_kde.resample依赖显式随机键,遵循 JAX "显式 PRNG" 的全局约定。

使用边界(以当前仓库源码为准)

  • fft.dct/dctn/idct/idctntype参数仅支持 2,其他类型抛出NotImplementedErrornorm仅支持None/"backward"/"ortho",传入"forward"会抛ValueError
  • special.lpmnspecial.lpmn_values已弃用且无替代方案,新代码应避免使用;
  • scipy.signalfftconvolve是 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提供可微概率分布,linalgsparse.linalg提供可编译的稠密/稀疏求解,fftsignalndimageinterpolate覆盖信号与图像处理,optimizecluster补齐经典数值算法。使用时注意各函数的实现边界(如 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),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询