跳转至

torch-sla

PyTorch Sparse Linear Algebra —— 一个可微分、多后端的稀疏线性方程求解库。

一句话

在 PyTorch 生态里做稀疏线性代数(Ax = b、特征值、范数、adjoint 非线性求解),对 CUDA / CPU 多套后端统一封装,保留梯度可以端到端训练。

为什么做

  • 深度学习框架的稀疏支持单薄:PyTorch 原生稀疏矩阵 API 较弱,生产级求解(尤其 GPU)要自己拼 cuDSS / CuPy / SciPy
  • 科学计算/PINN/FEM 场景需要可微:隐式求解器要跟 torch.autograd 打通
  • 多后端选择标准混乱:device / dtype / nnz / symmetry / posdef 各自影响最优后端,普通用户难以手动选

torch-sla 把这些都封在 SparseTensor.solve(b) 后面,自动选后端、自动微分。

特性

🔥 可微 通过 torch.autograd 全链路梯度支持
🚀 多后端 SciPy / Eigen (CPU) · CuPy / cuDSS / PyTorch-native (CUDA)
📦 批量稀疏张量 [..., M, N, ...] 形状
🎯 属性自动检测 symmetry / positive-definiteness 自动识别
⚡ 自适应调度 根据 device / dtype / problem size 自动选最优 solver
🌐 分布式 Domain decomposition + halo exchange(CFD / FEM 风格)
🧮 非线性求解 adjoint-based Newton / Anderson + 隐式微分

最近改动

v0.2.0 (2026-04): - 后端重构:cuDSS 从 C++ glue 换成 nvmath-python(NVIDIA 官方 Python 绑定),免掉自维护的 C++ 编译层 - CuPy 后端替换了 cuSOLVER,命名统一 - multi-RHS solve 整合进 backend 系统 - 新增 scipy_det 稀疏行列式计算 - dtype 修复(torch.result_type 统一 spmv/spmm 的类型提升规则)

完整 changelog 见 GitHub Releases

快速开始

import torch
from torch_sla import SparseTensor

dense = torch.tensor([[4.0, -1.0,  0.0],
                      [-1.0, 4.0, -1.0],
                      [ 0.0, -1.0, 4.0]], dtype=torch.float64)
A = SparseTensor.from_dense(dense)

b = torch.tensor([1.0, 2.0, 3.0], dtype=torch.float64)
x = A.solve(b)                        # 自动选后端
x = A.solve(b, backend='scipy', method='lu')   # 手动指定

完整示例 / benchmark 看 torchsla.com/examples

相关


上级 · Projects