torch-sla¶
PyTorch Sparse Linear Algebra —— 一个可微分、多后端的稀疏线性方程求解库。
-
arXiv
2601.13994 -
Repo
walkerchi/torch-sla -
PyPI
pip install torch-sla
一句话¶
在 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。
相关¶
- Paper: arXiv:2601.13994(收录于 Publications)
- License: MIT
- Python: 3.8+
↑ 上级 · Projects