资讯动态

JAX 的 SciPy 兼容模块 jax.scipy 完全指南:从特殊函数到稀疏线性代数

发布时间:2026/9/10 2:08:31 来源:尧图企业网站定制
JAX 的 SciPy 兼容模块 jax.scipy 完全指南从特殊函数到稀疏线性代数【免费下载链接】jaxComposable transformations of PythonNumPy 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, type2, nNone, axis-1, normNone)参数类型默认值语义x数组必填输入数据支持实数与复数typeint2变换类型当前仅支持 type2传入其他值抛出NotImplementedErrornintx.shape[axis]变换长度大于输入长度时零填充lax.pad小于输入长度时截断axisint-1沿哪个轴做变换经canonicalize_axis归一化normstrNone归一化模式取值None/backward/ortho默认等价于backwardforward未实现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)并取实部乘以 2normortho时通过_dct_ortho_norm施加正交归一化因子。复数输入会被拆成实部/虚部分别变换后再用lax.complex重组。idct是dct的逆过程先在频域除以旋转因子、乘以2N再ifft后经_dct_deinterleave还原奇偶位。因此文档示例中jnp.allclose(x, idct(dct(x)))返回Array(True, dtypebool)验证了正逆变换的闭环性质。三、jax.scipy.integrate梯形法则数值积分jax.scipy.integrate目前只暴露一个函数trapezoidSciPy 1.6 中trapz的替代名实现复合梯形法则trapezoid(y, xNone, dx1.0, axis-1)y被积分数据数组x采样点坐标缺省时按dx等间距分布dxx缺省时的采样间距默认1.0axis积分轴默认最后一轴。从 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]以dx1.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、welchWelch 功率谱估计、csd互谱密度预处理detrend去除趋势以 jax/_src/scipy/signal.py 中的fftconvolve为例签名与参数为fftconvolve(in1, in2, modefull, axesNone)in1/in2两个输入数组要求in1.ndim in2.ndimmode输出尺寸控制三选一——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可以是任意嵌套结构数组、字典、列表而不仅是单个向量内积运算统一使用precisionlax.Precision.HIGHEST见_dot partial(jnp.dot, precisionlax.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_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超几何函数、wofzFaddeeva 函数、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 / isfnorm完整八件套、cauchy、gumbel_l、gumbel_r、pareto、truncnorm、uniform、laplacecdf/logpdf/pdf、logisticcdf/isf/logpdf/pdf/ppf/sf、expon、gamma、chi2、beta、vonmises、wrapcauchy、gennorm、tlogpdf/pdf、dirichlet、multivariate_normallogpdf/pdf离散分布logpmf / pmfbernoulli另有 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、spatialjax.scipy.cluster向量量化vqvector 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旋转与插值Rotation3D 旋转对象支持矩阵、四元数、欧拉角、旋转向量等多种表示之间的转换Slerp球面线性插值用于旋转的平滑过渡在机器人运动规划、骨骼动画插帧中很实用。十、与 jax.numpy / jax.lax 的关系及使用边界jax.scipy并不是一个孤立命名空间它与 JAX 底层体系紧密耦合复用jax.numpy实现如integrate.trapezoid直接委托jax.numpy.trapezoidfft.dct复用jax._src.numpy.fftsparse.linalg用jnp.dot/jnp.vdot/einsum搭建内积与矩阵乘法。构建在lax原语之上dct中的lax.slice_in_dim、lax.rev、lax.pad、lax.expand_dimssparse.linalg中的lax.Precision.HIGHEST都是直接调用 XLA 级原语保证算子可编译、可求导。与jax.random协作stats.gaussian_kde.resample依赖显式随机键遵循 JAX 显式 PRNG 的全局约定。使用边界以当前仓库源码为准fft.dct/dctn/idct/idctn的type参数仅支持 2其他类型抛出NotImplementedErrornorm仅支持None/backward/ortho传入forward会抛ValueErrorspecial.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) FFTDCT 正逆变换闭环 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, dtypejnp.float32) x jsp.sparse.linalg.cg(A, b, maxiter100)[0] assert jnp.allclose(A x, b, atol1e-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]), methodBFGS) 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 PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价