torchsparsegradutils 0.1.3__tar.gz → 0.2.1__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (110) hide show
  1. torchsparsegradutils-0.2.1/PKG-INFO +703 -0
  2. torchsparsegradutils-0.2.1/README.md +661 -0
  3. torchsparsegradutils-0.2.1/docs/source/conf.py +187 -0
  4. torchsparsegradutils-0.2.1/pyproject.toml +110 -0
  5. torchsparsegradutils-0.2.1/setup.py +5 -0
  6. {torchsparsegradutils-0.1.3 → torchsparsegradutils-0.2.1}/torchsparsegradutils/__init__.py +2 -2
  7. torchsparsegradutils-0.2.1/torchsparsegradutils/benchmarks/__init__.py +10 -0
  8. torchsparsegradutils-0.2.1/torchsparsegradutils/benchmarks/batched_sparse_mm_rand.py +444 -0
  9. torchsparsegradutils-0.2.1/torchsparsegradutils/benchmarks/benchmark_suite.py +77 -0
  10. torchsparsegradutils-0.2.1/torchsparsegradutils/benchmarks/benchmark_utils.py +404 -0
  11. torchsparsegradutils-0.2.1/torchsparsegradutils/benchmarks/sparse_generic_solve_rand.py +309 -0
  12. torchsparsegradutils-0.2.1/torchsparsegradutils/benchmarks/sparse_generic_solve_suite.py +273 -0
  13. torchsparsegradutils-0.2.1/torchsparsegradutils/benchmarks/sparse_mm_rand.py +168 -0
  14. torchsparsegradutils-0.2.1/torchsparsegradutils/benchmarks/sparse_mm_suite.py +160 -0
  15. torchsparsegradutils-0.2.1/torchsparsegradutils/benchmarks/sparse_triangular_solve_rand.py +244 -0
  16. torchsparsegradutils-0.2.1/torchsparsegradutils/benchmarks/sparse_triangular_solve_suitesparse.py +264 -0
  17. torchsparsegradutils-0.2.1/torchsparsegradutils/benchmarks/visualize_benchmark_results.py +1049 -0
  18. {torchsparsegradutils-0.1.3 → torchsparsegradutils-0.2.1}/torchsparsegradutils/cupy/__init__.py +1 -1
  19. torchsparsegradutils-0.2.1/torchsparsegradutils/cupy/cupy_bindings.py +245 -0
  20. torchsparsegradutils-0.2.1/torchsparsegradutils/cupy/cupy_sparse_solve.py +422 -0
  21. torchsparsegradutils-0.2.1/torchsparsegradutils/distributions/__init__.py +3 -0
  22. torchsparsegradutils-0.2.1/torchsparsegradutils/distributions/sparse_multivariate_normal.py +589 -0
  23. torchsparsegradutils-0.2.1/torchsparsegradutils/encoders/__init__.py +44 -0
  24. torchsparsegradutils-0.2.1/torchsparsegradutils/encoders/pairwise_encoder.py +849 -0
  25. torchsparsegradutils-0.2.1/torchsparsegradutils/encoders/pairwise_voxel_encoder.py +118 -0
  26. torchsparsegradutils-0.2.1/torchsparsegradutils/indexed_matmul.py +217 -0
  27. {torchsparsegradutils-0.1.3 → torchsparsegradutils-0.2.1}/torchsparsegradutils/jax/__init__.py +6 -2
  28. torchsparsegradutils-0.2.1/torchsparsegradutils/jax/jax_bindings.py +313 -0
  29. torchsparsegradutils-0.2.1/torchsparsegradutils/jax/jax_sparse_solve.py +258 -0
  30. torchsparsegradutils-0.2.1/torchsparsegradutils/sparse_lstsq.py +271 -0
  31. torchsparsegradutils-0.2.1/torchsparsegradutils/sparse_matmul.py +234 -0
  32. torchsparsegradutils-0.2.1/torchsparsegradutils/sparse_solve.py +514 -0
  33. torchsparsegradutils-0.2.1/torchsparsegradutils/tests/test_bicgstab.py +61 -0
  34. torchsparsegradutils-0.2.1/torchsparsegradutils/tests/test_cupy_bindings.py +123 -0
  35. torchsparsegradutils-0.2.1/torchsparsegradutils/tests/test_cupy_sparse_solve.py +278 -0
  36. torchsparsegradutils-0.2.1/torchsparsegradutils/tests/test_dist_stats_helpers.py +321 -0
  37. torchsparsegradutils-0.2.1/torchsparsegradutils/tests/test_distributions.py +595 -0
  38. torchsparsegradutils-0.2.1/torchsparsegradutils/tests/test_doctests.py +73 -0
  39. {torchsparsegradutils-0.1.3 → torchsparsegradutils-0.2.1}/torchsparsegradutils/tests/test_encoders.py +245 -40
  40. torchsparsegradutils-0.2.1/torchsparsegradutils/tests/test_integration_pairwise_sparse_mvn.py +761 -0
  41. torchsparsegradutils-0.2.1/torchsparsegradutils/tests/test_jax_bindings.py +123 -0
  42. torchsparsegradutils-0.2.1/torchsparsegradutils/tests/test_jax_sparse_solve.py +223 -0
  43. torchsparsegradutils-0.2.1/torchsparsegradutils/tests/test_linear_cg.py +123 -0
  44. torchsparsegradutils-0.2.1/torchsparsegradutils/tests/test_lsmr.py +255 -0
  45. torchsparsegradutils-0.2.1/torchsparsegradutils/tests/test_minres.py +72 -0
  46. torchsparsegradutils-0.2.1/torchsparsegradutils/tests/test_quickstart_guide.py +189 -0
  47. torchsparsegradutils-0.2.1/torchsparsegradutils/tests/test_random.py +923 -0
  48. torchsparsegradutils-0.2.1/torchsparsegradutils/tests/test_sparse_lstsq.py +242 -0
  49. torchsparsegradutils-0.2.1/torchsparsegradutils/tests/test_sparse_matmul.py +396 -0
  50. torchsparsegradutils-0.2.1/torchsparsegradutils/tests/test_sparse_solve.py +271 -0
  51. torchsparsegradutils-0.2.1/torchsparsegradutils/tests/test_sparse_triangular_solve.py +305 -0
  52. torchsparsegradutils-0.2.1/torchsparsegradutils/tests/test_utils.py +290 -0
  53. {torchsparsegradutils-0.1.3 → torchsparsegradutils-0.2.1}/torchsparsegradutils/utils/__init__.py +9 -6
  54. torchsparsegradutils-0.2.1/torchsparsegradutils/utils/bicgstab.py +247 -0
  55. torchsparsegradutils-0.2.1/torchsparsegradutils/utils/dist_stats_helpers.py +373 -0
  56. {torchsparsegradutils-0.1.3 → torchsparsegradutils-0.2.1}/torchsparsegradutils/utils/linear_cg.py +116 -38
  57. {torchsparsegradutils-0.1.3 → torchsparsegradutils-0.2.1}/torchsparsegradutils/utils/lsmr.py +121 -59
  58. {torchsparsegradutils-0.1.3 → torchsparsegradutils-0.2.1}/torchsparsegradutils/utils/minres.py +123 -28
  59. torchsparsegradutils-0.2.1/torchsparsegradutils/utils/random_sparse.py +1371 -0
  60. torchsparsegradutils-0.2.1/torchsparsegradutils/utils/utils.py +914 -0
  61. torchsparsegradutils-0.2.1/torchsparsegradutils.egg-info/PKG-INFO +703 -0
  62. {torchsparsegradutils-0.1.3 → torchsparsegradutils-0.2.1}/torchsparsegradutils.egg-info/SOURCES.txt +19 -0
  63. torchsparsegradutils-0.2.1/torchsparsegradutils.egg-info/requires.txt +24 -0
  64. {torchsparsegradutils-0.1.3 → torchsparsegradutils-0.2.1}/torchsparsegradutils.egg-info/top_level.txt +3 -0
  65. torchsparsegradutils-0.1.3/PKG-INFO +0 -66
  66. torchsparsegradutils-0.1.3/README.md +0 -42
  67. torchsparsegradutils-0.1.3/pyproject.toml +0 -24
  68. torchsparsegradutils-0.1.3/setup.py +0 -42
  69. torchsparsegradutils-0.1.3/torchsparsegradutils/cupy/cupy_bindings.py +0 -77
  70. torchsparsegradutils-0.1.3/torchsparsegradutils/cupy/cupy_sparse_solve.py +0 -107
  71. torchsparsegradutils-0.1.3/torchsparsegradutils/distributions/__init__.py +0 -3
  72. torchsparsegradutils-0.1.3/torchsparsegradutils/distributions/sparse_multivariate_normal.py +0 -198
  73. torchsparsegradutils-0.1.3/torchsparsegradutils/encoders/__init__.py +0 -3
  74. torchsparsegradutils-0.1.3/torchsparsegradutils/encoders/pairwise_voxel_encoder.py +0 -511
  75. torchsparsegradutils-0.1.3/torchsparsegradutils/indexed_matmul.py +0 -117
  76. torchsparsegradutils-0.1.3/torchsparsegradutils/jax/jax_bindings.py +0 -80
  77. torchsparsegradutils-0.1.3/torchsparsegradutils/jax/jax_sparse_solve.py +0 -95
  78. torchsparsegradutils-0.1.3/torchsparsegradutils/sparse_lstsq.py +0 -139
  79. torchsparsegradutils-0.1.3/torchsparsegradutils/sparse_matmul.py +0 -130
  80. torchsparsegradutils-0.1.3/torchsparsegradutils/sparse_solve.py +0 -308
  81. torchsparsegradutils-0.1.3/torchsparsegradutils/tests/test_bicgstab.py +0 -60
  82. torchsparsegradutils-0.1.3/torchsparsegradutils/tests/test_cupy_bindings.py +0 -96
  83. torchsparsegradutils-0.1.3/torchsparsegradutils/tests/test_cupy_sparse_solve.py +0 -86
  84. torchsparsegradutils-0.1.3/torchsparsegradutils/tests/test_distributions.py +0 -247
  85. torchsparsegradutils-0.1.3/torchsparsegradutils/tests/test_jax_bindings.py +0 -110
  86. torchsparsegradutils-0.1.3/torchsparsegradutils/tests/test_jax_sparse_solve.py +0 -77
  87. torchsparsegradutils-0.1.3/torchsparsegradutils/tests/test_linear_cg.py +0 -151
  88. torchsparsegradutils-0.1.3/torchsparsegradutils/tests/test_lsmr.py +0 -315
  89. torchsparsegradutils-0.1.3/torchsparsegradutils/tests/test_minres.py +0 -113
  90. torchsparsegradutils-0.1.3/torchsparsegradutils/tests/test_random.py +0 -476
  91. torchsparsegradutils-0.1.3/torchsparsegradutils/tests/test_sparse_lstsq.py +0 -78
  92. torchsparsegradutils-0.1.3/torchsparsegradutils/tests/test_sparse_matmul.py +0 -187
  93. torchsparsegradutils-0.1.3/torchsparsegradutils/tests/test_sparse_solve.py +0 -598
  94. torchsparsegradutils-0.1.3/torchsparsegradutils/tests/test_utils.py +0 -705
  95. torchsparsegradutils-0.1.3/torchsparsegradutils/utils/bicgstab.py +0 -187
  96. torchsparsegradutils-0.1.3/torchsparsegradutils/utils/random_sparse.py +0 -367
  97. torchsparsegradutils-0.1.3/torchsparsegradutils/utils/utils.py +0 -480
  98. torchsparsegradutils-0.1.3/torchsparsegradutils.egg-info/PKG-INFO +0 -66
  99. torchsparsegradutils-0.1.3/torchsparsegradutils.egg-info/requires.txt +0 -5
  100. {torchsparsegradutils-0.1.3 → torchsparsegradutils-0.2.1}/LICENSE +0 -0
  101. {torchsparsegradutils-0.1.3 → torchsparsegradutils-0.2.1}/MANIFEST.in +0 -0
  102. {torchsparsegradutils-0.1.3 → torchsparsegradutils-0.2.1}/setup.cfg +0 -0
  103. {torchsparsegradutils-0.1.3 → torchsparsegradutils-0.2.1}/torchsparsegradutils/distributions/constraints.py +0 -0
  104. {torchsparsegradutils-0.1.3 → torchsparsegradutils-0.2.1}/torchsparsegradutils/tests/__init__.py +0 -0
  105. {torchsparsegradutils-0.1.3 → torchsparsegradutils-0.2.1}/torchsparsegradutils/tests/test_indexed_matmul.py +1 -1
  106. {torchsparsegradutils-0.1.3 → torchsparsegradutils-0.2.1}/torchsparsegradutils/tests/test_params/czyx_shifts.yaml +0 -0
  107. {torchsparsegradutils-0.1.3 → torchsparsegradutils-0.2.1}/torchsparsegradutils/tests/test_params/pairwise_coo_indices.yaml +0 -0
  108. {torchsparsegradutils-0.1.3 → torchsparsegradutils-0.2.1}/torchsparsegradutils/tests/test_params/xyz_coords.yaml +0 -0
  109. {torchsparsegradutils-0.1.3 → torchsparsegradutils-0.2.1}/torchsparsegradutils.egg-info/dependency_links.txt +0 -0
  110. {torchsparsegradutils-0.1.3 → torchsparsegradutils-0.2.1}/torchsparsegradutils.egg-info/not-zip-safe +0 -0
@@ -0,0 +1,703 @@
1
+ Metadata-Version: 2.4
2
+ Name: torchsparsegradutils
3
+ Version: 0.2.1
4
+ Summary: A collection of utility functions to work with PyTorch sparse tensors
5
+ Author-email: CAI4CAI research group <contact@cai4cai.uk>
6
+ License-Expression: Apache-2.0
7
+ Project-URL: Homepage, https://github.com/cai4cai/torchsparsegradutils
8
+ Project-URL: Documentation, https://torchsparsegradutils.readthedocs.io
9
+ Project-URL: Repository, https://github.com/cai4cai/torchsparsegradutils
10
+ Project-URL: Issues, https://github.com/cai4cai/torchsparsegradutils/issues
11
+ Project-URL: Changelog, https://github.com/cai4cai/torchsparsegradutils/releases
12
+ Keywords: sparse,torch,utility
13
+ Classifier: Operating System :: OS Independent
14
+ Classifier: Programming Language :: Python :: 3.10
15
+ Classifier: Programming Language :: Python :: 3.11
16
+ Classifier: Programming Language :: Python :: 3.12
17
+ Requires-Python: >=3.10
18
+ Description-Content-Type: text/markdown
19
+ License-File: LICENSE
20
+ Requires-Dist: torch>=2.5
21
+ Requires-Dist: scipy
22
+ Provides-Extra: extras
23
+ Requires-Dist: jax; extra == "extras"
24
+ Requires-Dist: cupy; extra == "extras"
25
+ Provides-Extra: docs
26
+ Requires-Dist: sphinx>=7.0.0; extra == "docs"
27
+ Requires-Dist: sphinx-rtd-theme>=1.3.0; extra == "docs"
28
+ Requires-Dist: sphinx-copybutton>=0.5.0; extra == "docs"
29
+ Requires-Dist: myst-parser>=2.0.0; extra == "docs"
30
+ Requires-Dist: sphinx-autobuild>=2021.3.14; extra == "docs"
31
+ Requires-Dist: matplotlib>=3.5.0; extra == "docs"
32
+ Requires-Dist: sphinx-autodoc-typehints>=1.24.0; extra == "docs"
33
+ Requires-Dist: sphinxcontrib-bibtex>=2.5.0; extra == "docs"
34
+ Requires-Dist: sphinx-gallery>=0.13.0; extra == "docs"
35
+ Provides-Extra: dev
36
+ Requires-Dist: pytest; extra == "dev"
37
+ Requires-Dist: pytest-cov; extra == "dev"
38
+ Requires-Dist: black; extra == "dev"
39
+ Requires-Dist: isort; extra == "dev"
40
+ Requires-Dist: flake8; extra == "dev"
41
+ Dynamic: license-file
42
+
43
+ # torchsparsegradutils: Sparsity-preserving gradient utility tools for PyTorch
44
+
45
+ [![PyPI](https://img.shields.io/pypi/v/torchsparsegradutils.svg)](https://pypi.org/project/torchsparsegradutils/) [![Python Versions](https://img.shields.io/pypi/pyversions/torchsparsegradutils.svg)](https://pypi.org/project/torchsparsegradutils/) [![Downloads](https://img.shields.io/pypi/dm/torchsparsegradutils.svg)](https://pypi.org/project/torchsparsegradutils/) ![PyTorch 2.5+](https://img.shields.io/badge/PyTorch-2.5%2B-ee4c2c?logo=pytorch) ![Tested 2.5 / 2.9 / nightly](https://img.shields.io/badge/Tested-2.5%20|%202.9%20|%20nightly-ee4c2c?logo=pytorch) [![Build](https://github.com/cai4cai/torchsparsegradutils/actions/workflows/python-package.yml/badge.svg)](https://github.com/cai4cai/torchsparsegradutils/actions/workflows/python-package.yml) [![Docs](https://readthedocs.org/projects/torchsparsegradutils/badge/?version=latest)](https://readthedocs.org/projects/torchsparsegradutils) [![Code Style: Black](https://img.shields.io/badge/code%20style-black-000000.svg)](https://github.com/psf/black) [![License](https://img.shields.io/github/license/cai4cai/torchsparsegradutils)](LICENSE) [![status](https://joss.theoj.org/papers/6da0e92488d06f70c0a03d0a7cbfba7d/status.svg)](https://joss.theoj.org/papers/6da0e92488d06f70c0a03d0a7cbfba7d)
46
+
47
+ A comprehensive collection of utility functions to work with PyTorch sparse tensors, ensuring memory efficiency and supporting various sparsity-preserving tensor operations with automatic differentiation. This package addresses fundamental gaps in PyTorch's sparse tensor ecosystem, providing essential operations that preserve sparsity in gradients during backpropagation.
48
+
49
+ ## 🚀 Key Features
50
+
51
+ ### Core Sparse Operations with Sparse Gradient Support
52
+
53
+ **Memory-Efficient Sparse Matrix Multiplication**
54
+ - `sparse_mm`: Memory-efficient sparse matrix multiplication with batch support
55
+ - Preserves sparsity in gradients during backpropagation
56
+ - Workaround for [PyTorch issue #41128](https://github.com/pytorch/pytorch/issues/41128)
57
+ - Supports both COO and CSR formats with optional batching
58
+
59
+ **Sparse Linear System Solvers**
60
+ - `sparse_triangular_solve`: Sparse triangular solver with batch support
61
+ - Discussion reference: [PyTorch issue #87358](https://github.com/pytorch/pytorch/issues/87358)
62
+ - `sparse_generic_solve`: Generic sparse linear solver with pluggable backends
63
+ - Tested and benchmarked with CG, BICGSTAB, LSMR and MINRES solvers
64
+
65
+ - `sparse_solve_c4t`: Wrappers around [cupy sparse solvers](https://docs.cupy.dev/en/stable/reference/scipy_sparse_linalg.html#solving-linear-problems)
66
+ - Discussion reference: [Pytorch issue #69538](https://github.com/pytorch/pytorch/issues/69538)
67
+ - Tested and benchmarked with: [CG](https://docs.cupy.dev/en/v9.6.0/reference/generated/cupyx.scipy.sparse.linalg.cg.html), [CGS](https://docs.cupy.dev/en/stable/reference/generated/cupyx.scipy.sparse.linalg.cgs.html#cupyx.scipy.sparse.linalg.cgs), [MINRES](https://docs.cupy.dev/en/stable/reference/generated/cupyx.scipy.sparse.linalg.minres.html#cupyx.scipy.sparse.linalg.minres), [GMRES](https://docs.cupy.dev/en/stable/reference/generated/cupyx.scipy.sparse.linalg.gmres.html#cupyx.scipy.sparse.linalg.gmres), [spsolve](https://docs.cupy.dev/en/stable/reference/generated/cupyx.scipy.sparse.linalg.spsolve.html#cupyx.scipy.sparse.linalg.spsolve) and [spsolve_triangular](https://docs.cupy.dev/en/stable/reference/generated/cupyx.scipy.sparse.linalg.spsolve_triangular.html#cupyx.scipy.sparse.linalg.spsolve_triangular) CuPy solvers
68
+ - `tsgujax.sparse_solve_j4t`: Wrappers around [jax sparse solvers](https://jax.readthedocs.io/en/latest/jax.scipy.html#module-jax.scipy.sparse.linalg)
69
+ - Tested with: CG and BICGSTAB JAX solvers
70
+ - `sparse_generic_lstsq`: Generic sparse linear least-squares solver
71
+
72
+ ### Built-in Iterative Solvers (No External Dependencies)
73
+
74
+ **Pure PyTorch Implementations**
75
+ - **BICGSTAB**: Biconjugate Gradient Stabilized method (ported from [pykrylov](https://github.com/PythonOptimizers/pykrylov))
76
+ - **CG**: Conjugate Gradient method (ported from [cornellius-gp/linear_operator](https://github.com/cornellius-gp/linear_operator))
77
+ - **LSMR**: Least Squares Minimal Residual method (ported from [pytorch-minimize](https://github.com/rfeinman/pytorch-minimize))
78
+ - **MINRES**: Minimal Residual method (ported from [cornellius-gp/linear_operator](https://github.com/cornellius-gp/linear_operator))
79
+
80
+ ### Sparse Multivariate Normal Distributions
81
+
82
+ - **SparseMultivariateNormal**: Structured Gaussian Distribution
83
+ - Implements reparameterised sampling (rsample)
84
+ - Supports leading batch dimension
85
+ - Supports COO and CSR sparse tensors
86
+ - Covariance or precision matrices with LL^T or LDL^T parameterisations.
87
+ - LDL^T parameterization offers numerical stability without SPD constraints
88
+ - **SparseMultivariateNormalNative**:
89
+ - Implements reparameterised sampling (rsample)
90
+ - Uses native `torch.sparse.mm` only
91
+ - Only supports ubatched CSR tensors
92
+ - Covariance LL^T parameterization
93
+
94
+ ### Spatial Encoding Tools
95
+
96
+ **Pairwise Encoder**
97
+ - Encode local neighborhood relationships in nD spatial volumes
98
+ - Multi-channel/class support
99
+ - Configurable neighborhood radius and sparsity patterns
100
+ - Outputs sparse unbatched/batched COO or CSR matrices for downstream processing
101
+ - Optimised for medical imaging and volumetric data applications
102
+
103
+ ### Graph Neural Network Operations
104
+
105
+ **Indexed Matrix Multiplication**
106
+ - `segment_mm`: Segmented matrix multiplication compatible with DGL/PyG
107
+ - `gather_mm`: Gather-based matrix multiplication for graph operations
108
+ - Pure PyTorch implementations as alternatives to [`dgl.ops.segment_mm`](https://docs.dgl.ai/generated/dgl.ops.segment_mm.html), [`pyg_lib.ops.segment_matmul`](https://pyg-lib.readthedocs.io/en/latest/modules/ops.html#pyg_lib.ops.segment_matmul), and [`dgl.ops.gather_mm`](https://docs.dgl.ai/generated/dgl.ops.gather_mm.html)
109
+ - Supports PyTorch >= 2.4 with nested tensor operations
110
+
111
+
112
+
113
+ ## 🛠️ Installation
114
+
115
+ ### Basic Installation
116
+
117
+ The package can be installed using pip:
118
+
119
+ ```bash
120
+ pip install torchsparsegradutils
121
+ ```
122
+
123
+ ### Development Installation
124
+
125
+ For the latest features and development work:
126
+
127
+ ```bash
128
+ pip install git+https://github.com/cai4cai/torchsparsegradutils
129
+ ```
130
+
131
+ ### Optional Dependencies
132
+
133
+ For full functionality, install optional dependencies:
134
+
135
+ ```bash
136
+ # For CuPy sparse solver support (GPU acceleration)
137
+ pip install cupy-cuda12x # Replace with your CUDA version
138
+
139
+ # For JAX sparse solver support
140
+ pip install "jax[cpu]" # CPU version
141
+ pip install "jax[cuda12]" # GPU version (replace with your CUDA version)
142
+
143
+ # For benchmarking and testing
144
+ pip install scipy matplotlib pandas tqdm pytest
145
+ ```
146
+
147
+ ### Requirements
148
+
149
+ - **Python**: ≥ 3.10
150
+ - **PyTorch**: ≥ 2.5 (≥ 2.4 for indexed operations)
151
+ - **Operating Systems**: Linux, macOS, Windows
152
+ - **Hardware**: CPU and CUDA GPU support
153
+
154
+
155
+ ## 📊 Performance Benchmarks
156
+
157
+ Our comprehensive benchmark suite demonstrates significant performance improvements across various sparse operations. All benchmarks were conducted on an NVIDIA GeForce RTX 4090 with PyTorch 2.9.0+cu130. Benchmarks are performed using [Rothberg/cfd2](https://suitesparse-collection-website.herokuapp.com/Rothberg/cfd2) matrix from [SuiteSparse Matrix Collection](https://suitesparse-collection-website.herokuapp.com/)
158
+
159
+ ![Sparse MM Suite Performance (int32/float32 COO)](torchsparsegradutils/benchmarks/benchmark_visualizations/sparse_mm_suite_performance_int32_float32_coo.png)
160
+
161
+ ![Sparse Triangular Solve Suite Performance (int32/float32 COO)](torchsparsegradutils/benchmarks/benchmark_visualizations/triangular_solve_suitesparse_performance_int32_float32_coo.png)
162
+
163
+ ![Sparse Genertic Solve Suite Performance (int32/float32 COO)](torchsparsegradutils/benchmarks/benchmark_visualizations/sparse_solve_suite_performance_int32_float32_coo.png)
164
+
165
+ ## 🚀 Quick Start
166
+
167
+ ### Basic Sparse Matrix Multiplication
168
+
169
+ ```python
170
+ import torch
171
+ from torchsparsegradutils import sparse_mm
172
+
173
+ # Create sparse matrix in COO format
174
+ indices = torch.tensor([[0, 1, 1], [2, 0, 2]], dtype=torch.int64)
175
+ values = torch.tensor([3., 4., 5.], requires_grad=True)
176
+ A = torch.sparse_coo_tensor(indices, values, (2, 3))
177
+
178
+ # Dense matrix
179
+ B = torch.randn(3, 4, requires_grad=True)
180
+
181
+ # Memory-efficient sparse matrix multiplication with gradient support
182
+ C = sparse_mm(A, B)
183
+ loss = C.sum()
184
+ loss.backward() # Gradients preserved in sparse format
185
+
186
+ print(f"A.grad: {A.grad}") # Sparse gradient
187
+ print(f"B.grad: {B.grad}") # Dense gradient
188
+ ```
189
+
190
+ ### Sparse Linear System Solving
191
+
192
+ ```python
193
+ import torch
194
+ from torchsparsegradutils import sparse_triangular_solve, sparse_generic_solve
195
+ from torchsparsegradutils.utils import linear_cg
196
+
197
+ # Create sparse triangular matrix
198
+ A = create_sparse_triangular_matrix() # Your sparse CSR matrix
199
+ b = torch.randn(A.shape[0], requires_grad=True)
200
+
201
+ # Triangular solve (fast for triangular systems)
202
+ x1 = sparse_triangular_solve(A, b, upper=False)
203
+
204
+ # Generic solve with different backends
205
+ x2 = sparse_generic_solve(A, b, solve=linear_cg, tol=1e-6)
206
+
207
+ # Using CuPy backend (if available)
208
+ from torchsparsegradutils.cupy import sparse_solve_c4t
209
+ x3 = sparse_solve_c4t(A, b, solve="cg", tol=1e-6)
210
+ ```
211
+
212
+ ### Sparse Multivariate Normal Distribution
213
+
214
+ ```python
215
+ import torch
216
+ from torchsparsegradutils.distributions import SparseMultivariateNormal
217
+ from torchsparsegradutils.utils.random_sparse import rand_sparse_tri
218
+
219
+ # Create parameters
220
+ batch_size, event_size = 2, 1000
221
+ loc = torch.zeros(batch_size, event_size)
222
+
223
+ # Example 1: LDL^T parameterization (numerically stable for precision matrices)
224
+ # Create sparse lower triangular matrix (unit triangular, no diagonal)
225
+ scale_tril = rand_sparse_tri(
226
+ (batch_size, event_size, event_size),
227
+ nnz=5000, # 5000 non-zeros for 1M parameters (0.5% sparsity)
228
+ layout=torch.sparse_csr,
229
+ upper=False,
230
+ unit_triangular=True # Unit triangular for LDL^T
231
+ )
232
+
233
+ # Diagonal component for LDL^T parameterization
234
+ diagonal = torch.ones(batch_size, event_size) * 0.5
235
+
236
+ # Create distribution with LDL^T parameterization
237
+ dist_ldlt = SparseMultivariateNormal(
238
+ loc=loc,
239
+ diagonal=diagonal,
240
+ scale_tril=scale_tril # Unit lower triangular
241
+ )
242
+
243
+ # Example 2: LL^T parameterization (standard Cholesky)
244
+ scale_tril_chol = rand_sparse_tri(
245
+ (batch_size, event_size, event_size),
246
+ nnz=5000,
247
+ layout=torch.sparse_csr,
248
+ upper=False,
249
+ unit_triangular=False # Include diagonal for LL^T
250
+ )
251
+
252
+ # Create distribution with LL^T parameterization
253
+ dist_chol = SparseMultivariateNormal(
254
+ loc=loc,
255
+ scale_tril=scale_tril_chol # Lower triangular with diagonal
256
+ )
257
+
258
+ # Example 3: Precision matrix parameterization (more stable with LDL^T)
259
+ precision_tril = rand_sparse_tri(
260
+ (batch_size, event_size, event_size),
261
+ nnz=5000,
262
+ layout=torch.sparse_csr,
263
+ upper=False,
264
+ unit_triangular=True
265
+ )
266
+
267
+ precision_diagonal = torch.ones(batch_size, event_size) * 2.0
268
+
269
+ dist_precision = SparseMultivariateNormal(
270
+ loc=loc,
271
+ diagonal=precision_diagonal,
272
+ precision_tril=precision_tril # Unit triangular precision factor
273
+ )
274
+
275
+ # Sample with gradient support
276
+ samples = dist_ldlt.rsample((100,)) # 100 samples
277
+
278
+ # Gradient computation preserves sparsity
279
+ loss = samples.sum()
280
+ loss.backward()
281
+ print(f"Sparse gradient shape: {scale_tril.grad.shape}")
282
+ print(f"Sparse gradient nnz: {scale_tril.grad._nnz()}")
283
+ print(f"Using LDL^T parameterization: {dist_ldlt.is_ldlt_parameterization}")
284
+ ```
285
+
286
+ ### Pairwise Voxel Encoding
287
+
288
+ ```python
289
+ import torch
290
+ from torchsparsegradutils.encoders import PairwiseEncoder
291
+
292
+ # Create 3D volume encoder (channels, height, depth, width)
293
+ volume_shape = (4, 64, 64, 64) # 4 channels, 64x64x64 spatial
294
+ encoder = PairwiseEncoder(
295
+ radius=2.0,
296
+ volume_shape=volume_shape,
297
+ layout=torch.sparse_csr
298
+ )
299
+
300
+ # Generate values for each spatial relationship offset
301
+ num_offsets = len(encoder.offsets)
302
+ values = torch.randn(num_offsets, *volume_shape)
303
+
304
+ # Generate sparse encoding matrix
305
+ sparse_matrix = encoder(values)
306
+
307
+ print(f"Encoded volume shape: {sparse_matrix.shape}")
308
+ print(f"Sparsity: {sparse_matrix._nnz() / sparse_matrix.numel():.3%}")
309
+ print(f"Number of spatial offsets: {num_offsets}")
310
+
311
+ # Use in sparse multivariate normal
312
+ flat_size = 4 * 64 * 64 * 64 # Total flattened size
313
+ dist = SparseMultivariateNormal(
314
+ loc=torch.zeros(flat_size),
315
+ scale_tril=sparse_matrix
316
+ )
317
+ ```
318
+
319
+ #### Spatial Relationship Visualization
320
+
321
+ The encoder creates sparse matrices that encode pairwise spatial relationships within a specified radius. Different channel relationship types affect how channels interact:
322
+
323
+ - **`indep`**: Independent channels (only spatial neighbors within same channel)
324
+ - **`intra`**: Intra-channel relationships (spatial neighbors within same channel)
325
+ - **`inter`**: Inter-channel relationships (spatial neighbors across all channels)
326
+
327
+ **3D Spatial Grid (3×3×3×3) with Different Channel Relations:**
328
+
329
+ <div align="center">
330
+
331
+ **Radius = 1.0**
332
+ ![Spatial Encodings Radius 1](torchsparsegradutils/tests/test_outputs/sparse_encodings_radius_1.png)
333
+
334
+ **Radius = 2.0**
335
+ ![Spatial Encodings Radius 2](torchsparsegradutils/tests/test_outputs/sparse_encodings_radius_2.png)
336
+ <!--
337
+ **Legend for Spatial Offsets:**
338
+ <table>
339
+ <tr>
340
+ <td><img src="torchsparsegradutils/tests/test_outputs/legend_radius_1.png" width="150"/></td>
341
+ <td><img src="torchsparsegradutils/tests/test_outputs/legend_radius_2.png" width="150"/></td>
342
+ </tr>
343
+ <tr>
344
+ <td align="center">Radius 1.0 Offsets</td>
345
+ <td align="center">Radius 2.0 Offsets</td>
346
+ </tr>
347
+ </table> -->
348
+
349
+ </div>
350
+
351
+ Each color represents a different spatial offset (relative position) in the 3D neighborhood. The sparse matrix encodes these relationships efficiently, enabling:
352
+
353
+ - **Local spatial modeling** for volumetric data (medical imaging, 3D computer vision)
354
+ - **Multi-channel feature interaction** in convolutional architectures
355
+ - **Sparse graph construction** from regular grids
356
+ - **Memory-efficient neighborhood encoding** for large volumes
357
+
358
+ **Key Parameters:**
359
+ - `radius`: Spatial neighborhood radius (1.0 = immediate neighbors, 2.0 = extended neighborhood)
360
+ - `volume_shape`: `(channels, height, depth, width)` for 4D volumes
361
+ - `channel_voxel_relation`: Controls cross-channel connectivity patterns
362
+ - `layout`: Output sparse format (`torch.sparse_coo` or `torch.sparse_csr`)
363
+
364
+ ### Indexed Matrix Operations (Graph Neural Networks)
365
+
366
+ ```python
367
+ import torch
368
+ from torchsparsegradutils import segment_mm, gather_mm
369
+
370
+ # Segment matrix multiplication (compatible with DGL/PyG)
371
+ a = torch.randn(15, 10, requires_grad=True) # Node features
372
+ b = torch.randn(3, 10, 5, requires_grad=True) # Edge type embeddings
373
+ seglen_a = torch.tensor([5, 6, 4]) # Segment lengths
374
+
375
+ # Performs: a[0:5] @ b[0], a[5:11] @ b[1], a[11:15] @ b[2]
376
+ result = segment_mm(a, b, seglen_a)
377
+
378
+ # Gather matrix multiplication
379
+ indices = torch.tensor([0, 0, 1, 1, 2])
380
+ a_gathered = torch.randn(5, 10, requires_grad=True)
381
+ result = gather_mm(a_gathered, b, indices)
382
+ ```
383
+
384
+ ### Statistical Distribution Validation
385
+
386
+ ```python
387
+ import torch
388
+ from torch.distributions import MultivariateNormal
389
+ from torchsparsegradutils.utils import mean_hotelling_t2_test, cov_nagao_test
390
+
391
+ # Generate sample data from known distribution
392
+ torch.manual_seed(42)
393
+ true_mean = torch.tensor([[0.0, 0.0]])
394
+ true_cov = torch.eye(2).unsqueeze(0)
395
+ n = 1000
396
+
397
+ # Generate samples and compute statistics
398
+ dist = MultivariateNormal(true_mean.squeeze(0), true_cov.squeeze(0))
399
+ samples = dist.sample((n,)).unsqueeze(1)
400
+ sample_mean = samples.mean(0)
401
+ sample_cov = torch.cov(samples.squeeze(1).T).unsqueeze(0)
402
+
403
+ # Test if sample mean is consistent with hypothesized mean (should pass)
404
+ result, t2_stat, threshold = mean_hotelling_t2_test(
405
+ sample_mean, true_mean, sample_cov, n, confidence_level=0.95
406
+ )
407
+ print(f"Mean test passed: {result.item()}") # True
408
+
409
+ # Test if sample covariance is consistent with hypothesized covariance (should pass)
410
+ result, t_n_stat, threshold = cov_nagao_test(
411
+ sample_cov, true_cov, n, confidence_level=0.95
412
+ )
413
+ print(f"Covariance test passed: {result.item()}") # True
414
+
415
+ # Test against wrong parameters (should fail)
416
+ wrong_mean = true_mean + 1.0 # Significantly different mean
417
+ result, _, _ = mean_hotelling_t2_test(
418
+ sample_mean, wrong_mean, sample_cov, n, confidence_level=0.95
419
+ )
420
+ print(f"Wrong mean test passed: {result.item()}") # False
421
+ ```
422
+
423
+ ## 🧪 Testing and Benchmarks
424
+
425
+ ### Running Tests
426
+
427
+ ```bash
428
+ # Run all tests
429
+ python -m pytest
430
+
431
+ # Run specific test modules
432
+ python -m pytest torchsparsegradutils/tests/test_sparse_matmul.py
433
+ python -m pytest torchsparsegradutils/tests/test_distributions.py
434
+
435
+ # Run with coverage
436
+ python -m pytest --cov=torchsparsegradutils
437
+ ```
438
+
439
+ ### Running Benchmarks
440
+
441
+ The package includes comprehensive benchmarks for performance evaluation:
442
+
443
+ ```bash
444
+ # Sparse matrix multiplication benchmarks
445
+ python -m torchsparsegradutils.benchmarks.sparse_mm_rand
446
+ python -m torchsparsegradutils.benchmarks.batched_sparse_mm_rand
447
+
448
+ # Triangular solver benchmarks
449
+ python -m torchsparsegradutils.benchmarks.sparse_triangular_solve_rand
450
+
451
+ # Generic solver benchmarks
452
+ python -m torchsparsegradutils.benchmarks.sparse_generic_solve_suite
453
+
454
+ # SuiteSparse matrix benchmarks
455
+ python -m torchsparsegradutils.benchmarks.sparse_mm_suite
456
+ ```
457
+
458
+ Results are automatically saved to `torchsparsegradutils/benchmarks/results/` as CSV files.
459
+
460
+ ### Utility Functions
461
+
462
+ #### `torchsparsegradutils.utils.random_sparse`
463
+
464
+ **Sparse Random Matrix Generators**
465
+ - **`rand_sparse(size, nnz, layout=torch.sparse_coo, **kwargs)`**: Generate random sparse matrices with specified layout and properties
466
+ - Supports COO and CSR
467
+ - Supports batch dimension
468
+ - **`rand_sparse_tri(size, nnz, layout=torch.sparse_coo, upper=True, strict=False, **kwargs)`**: Generate random sparse triangular matrices
469
+ - Supports COO and CSR
470
+ - Supports batch dimension
471
+ - Strict triangular (no diagonal) or non-strict (with diagonal values)
472
+ - Option to produce well conditioned matrices and regulate diagonal values
473
+
474
+ - **`make_spd_sparse(n, layout, value_dtype, index_dtype, device, sparsity_ratio=0.5, nz=None)`**: Generate sparse symmetric positive definite (SPD) matrices
475
+
476
+ #### `torchsparsegradutils.utils.utils`
477
+
478
+ **Sparse Matrix Operations**
479
+ - **`sparse_block_diag(*sparse_tensors)`**: Create block diagonal sparse matrix from multiple sparse tensors
480
+ - **`sparse_block_diag_split(sparse_block_diag_tensor, *shapes)`**: Split block diagonal sparse matrix into original sparse tensors
481
+ - **`sparse_eye(size, layout=torch.sparse_coo, **kwargs)`**: Create batched or unbatched sparse identity matrices
482
+ - **`stack_csr(tensors, dim=0)`**: Stack CSR tensors along batch dimension (like torch.stack for CSR)
483
+
484
+ **Sparse Format Conversion**
485
+ - **`convert_coo_to_csr_indices_values(coo_indices, num_rows, values=None)`**: Convert COO indices and values to CSR format, with support for batch dimension
486
+ - **`convert_coo_to_csr(sparse_coo_tensor)`**: Convert COO sparse tensor to CSR format with batch support
487
+
488
+ #### `torchsparsegradutils.utils.dist_stats_helpers`
489
+
490
+ **Statistical Distribution Validation**
491
+ - **`mean_hotelling_t2_test(sample_mean, true_mean, sample_cov, n, confidence_level=0.95)`**: One-sample Hotelling T² test for multivariate mean equality using confidence regions
492
+ - Tests whether hypothesized mean vector lies within confidence region around sample mean
493
+ - Uses F-distribution for threshold calculation with proper degrees of freedom
494
+ - Higher confidence levels create larger (more permissive) acceptance regions
495
+ - **`cov_nagao_test(emp_cov, ref_cov, n, confidence_level=0.95)`**: Nagao's test for covariance matrix equality using confidence regions
496
+ - Tests whether hypothesized covariance matrix is consistent with empirical covariance
497
+ - Uses χ² distribution with appropriate degrees of freedom
498
+ - Standardizes covariance matrices for improved numerical stability
499
+
500
+
501
+ ## 🤝 Contributing
502
+
503
+ We welcome contributions! Please see our contributing guidelines:
504
+
505
+ 1. **Issues**: Report bugs and request features via [GitHub Issues](https://github.com/cai4cai/torchsparsegradutils/issues)
506
+ 2. **Pull Requests**: Submit improvements via GitHub PRs
507
+ 3. **Testing**: Ensure all tests pass and add tests for new functionality
508
+ 4. **Documentation**: Update docstrings and examples for new features
509
+ 5. **Benchmarks**: Include performance benchmarks for new operations
510
+
511
+ ### Development Setup
512
+
513
+ #### Option 1: Local Development
514
+
515
+ ```bash
516
+ git clone https://github.com/cai4cai/torchsparsegradutils
517
+ cd torchsparsegradutils
518
+ pip install -e ".[dev]" # Install in development mode
519
+ pre-commit install # Install pre-commit hooks
520
+ ```
521
+
522
+ #### Option 2: Development Containers (Recommended)
523
+
524
+ For a consistent development environment with GPU support and all dependencies pre-installed, use VS Code Dev Containers:
525
+
526
+ **Prerequisites:**
527
+ - [Docker](https://docs.docker.com/get-docker/) with NVIDIA Container Toolkit (for GPU support)
528
+ - [VS Code](https://code.visualstudio.com/) with the [Dev Containers extension](https://marketplace.visualstudio.com/items?itemName=ms-vscode-remote.remote-containers)
529
+
530
+ **Quick Start:**
531
+ 1. Clone the repository and open in VS Code:
532
+ ```bash
533
+ git clone https://github.com/cai4cai/torchsparsegradutils
534
+ cd torchsparsegradutils
535
+ code .
536
+ ```
537
+
538
+ 2. When prompted, click **"Reopen in Container"** or use the Command Palette:
539
+ - Press `Ctrl+Shift+P` (or `Cmd+Shift+P` on macOS)
540
+ - Type "Dev Containers: Reopen in Container"
541
+
542
+ **Available Configurations:**
543
+
544
+ - **`.devcontainer/Dockerfile.stable`** (default): Uses stable PyTorch with CUDA 13.0 support
545
+ - **`.devcontainer/Dockerfile.nightly`**: Uses nightly PyTorch builds for latest features
546
+
547
+ To switch configurations, modify the `dockerfile` field in `.devcontainer/devcontainer.json`:
548
+ ```json
549
+ "build": {
550
+ "dockerfile": "./Dockerfile.nightly", // or "./Dockerfile.stable"
551
+ "context": "."
552
+ }
553
+ ```
554
+
555
+ **What's Included:**
556
+ - **CUDA 13.0**: Full GPU development support with NVIDIA drivers
557
+ - **Pre-installed Dependencies**: PyTorch, CuPy, JAX, SciPy, and all development tools
558
+ - **VS Code Extensions**: Python, Pylance, Jupyter, GitHub Copilot, and code formatting tools
559
+ - **Development Tools**: pytest, black, flake8, pre-commit hooks
560
+ - **Python Environment**: Python 3.10+ with all optional dependencies
561
+
562
+ **Benefits:**
563
+ - ✅ **Consistent Environment**: Same setup across different machines
564
+ - ✅ **GPU Support**: Pre-configured CUDA environment
565
+ - ✅ **Zero Setup**: All dependencies and tools pre-installed
566
+ - ✅ **Isolated**: No conflicts with host system packages
567
+ - ✅ **VS Code Integration**: Seamless debugging, IntelliSense, and testing
568
+
569
+ ## 📄 License
570
+
571
+ This project is licensed under the Apache License 2.0 - see the [LICENSE](LICENSE) file for details.
572
+
573
+ ## 🙏 Acknowledgments
574
+
575
+ - **PyTorch Team**: For the foundational sparse tensor implementations
576
+ - **SciPy/CuPy Teams**: For high-performance sparse linear algebra routines
577
+ - **JAX Team**: For cross-platform sparse operations and XLA compilation
578
+ - **Open Source Libraries**: We port and adapt algorithms from:
579
+ - [pykrylov](https://github.com/PythonOptimizers/pykrylov) (BICGSTAB)
580
+ - [cornellius-gp/linear_operator](https://github.com/cornellius-gp/linear_operator) (CG, MINRES)
581
+ - [pytorch-minimize](https://github.com/rfeinman/pytorch-minimize) (LSMR)
582
+
583
+ ## 📚 Citation
584
+
585
+ If you use this package in your research, please cite:
586
+
587
+ ```bibtex
588
+ @software{torchsparsegradutils,
589
+ title={torchsparsegradutils: Sparsity-preserving gradient utility tools for PyTorch},
590
+ author={Barfoot, Theodore and Glocker, Ben and Vercauteren, Tom},
591
+ url={https://github.com/cai4cai/torchsparsegradutils},
592
+ year={2024}
593
+ }
594
+ ```
595
+
596
+ ## ⚠️ Known Issues
597
+
598
+ ### PyTorch Sparse COO Index Dtype Conversion
599
+
600
+ **Issue**: PyTorch automatically converts `int32` indices to `int64` when creating sparse COO tensors, but preserves `int32` for sparse CSR tensors. This affects memory usage and performance for algorithms that benefit from `int32` indices (such as `sparse_mm`).
601
+
602
+ **Impact**:
603
+ - **Memory**: `int64` indices use 2× more memory than `int32`
604
+ - **Performance**: Some sparse operations may run faster with `int32` indices
605
+ - **Cross-format consistency**: Different behavior between COO and CSR formats
606
+
607
+ **Example**:
608
+ ```python
609
+ import torch
610
+
611
+ # Demonstrate the issue
612
+ indices_int32 = torch.tensor([[0, 1], [1, 0]], dtype=torch.int32)
613
+ values = torch.tensor([1.0, 2.0])
614
+
615
+ print(f"Original indices dtype: {indices_int32.dtype}") # torch.int32
616
+
617
+ # COO: int32 -> int64 conversion happens
618
+ coo_tensor = torch.sparse_coo_tensor(indices_int32, values, (2, 2)).coalesce()
619
+ print(f"COO indices dtype: {coo_tensor.indices().dtype}") # torch.int64 (converted!)
620
+
621
+ # CSR: int32 is preserved
622
+ crow_indices = torch.tensor([0, 1, 2], dtype=torch.int32)
623
+ col_indices = torch.tensor([1, 0], dtype=torch.int32)
624
+ csr_tensor = torch.sparse_csr_tensor(crow_indices, col_indices, values, (2, 2))
625
+ print(f"CSR crow_indices dtype: {csr_tensor.crow_indices().dtype}") # torch.int32 (preserved!)
626
+ print(f"CSR col_indices dtype: {csr_tensor.col_indices().dtype}") # torch.int32 (preserved!)
627
+ ```
628
+
629
+ **Workarounds**:
630
+ 1. **Use CSR format** when `int32` indices are important for performance
631
+ 2. **Account for extra memory** when using COO format with large sparse matrices
632
+ 3. **Test performance** with both dtypes to determine if the conversion impacts your use case
633
+
634
+ **Status**: This is a known PyTorch behavior. Our test suite documents and validates this behavior to catch any future changes in PyTorch's handling of sparse tensor index dtypes.
635
+
636
+ ### PairwiseEncoder CSR Memory Usage Issue
637
+
638
+ **Issue**: CSR sparse tensors generated by `PairwiseEncoder` consume significantly more memory during backward passes compared to COO format, particularly in integration tests with `SparseMultivariateNormal`.
639
+
640
+ **Impact**:
641
+ - **Memory Consumption**: CSR integration tests can use 2-3x more memory than equivalent COO tests during `.backward()`
642
+ - **Training Stability**: May cause out-of-memory errors during training with large spatial volumes
643
+ - **Development**: Affects integration testing with large tensor configurations
644
+
645
+ **Suspected Cause**: The issue may be related to CSR permutation operations within `PairwiseEncoder` that create additional intermediate tensors during gradient computation.
646
+
647
+ **Current Status**: Under investigation. The memory spike occurs specifically during backpropagation through the sparse matrix operations.
648
+
649
+ **Workarounds**:
650
+ 1. **Use COO format** for `PairwiseEncoder` when memory is constrained during training
651
+ 2. **Reduce batch sizes** or spatial dimensions when using CSR format
652
+ 3. **Monitor memory usage** carefully when integrating `PairwiseEncoder` with gradient-based optimization
653
+
654
+ **Example**:
655
+ ```python
656
+ # More memory-efficient approach for large tensors
657
+ encoder = PairwiseEncoder(
658
+ radius=2.0,
659
+ volume_shape=(4, 64, 64, 64),
660
+ layout=torch.sparse_coo # Use COO instead of CSR for memory efficiency
661
+ )
662
+ ```
663
+
664
+ ### SparseMultivariateNormal LL^T Precision Parameterization Gradient Issues
665
+
666
+ **Issue**: Large gradient magnitudes can occur when using LL^T parameterization with precision matrices in `SparseMultivariateNormal`, leading to training instability.
667
+
668
+ **Impact**:
669
+ - **Gradient Explosion**: Gradients can become extremely large (>1e6) during backpropagation
670
+ - **Training Instability**: May cause NaN values or divergent optimization
671
+ - **Numerical Issues**: Poor conditioning of the precision matrix can amplify gradient problems
672
+
673
+ **Affected Configurations**:
674
+ - LL^T parameterization (`scale_tril` parameter) combined with precision matrix formulation
675
+ - Both 2D and 3D spatial configurations show this behavior
676
+ - More pronounced with larger spatial dimensions and higher sparsity
677
+
678
+ **Root Cause**: The LL^T precision parameterization can lead to poor numerical conditioning, especially when the triangular matrix has small diagonal values or high condition number.
679
+
680
+ **Recommended Solution**: Use LDL^T parameterization instead, which provides better numerical stability:
681
+
682
+ ```python
683
+ # Problematic: LL^T precision parameterization
684
+ dist_unstable = SparseMultivariateNormal(
685
+ loc=loc,
686
+ precision_tril=scale_tril # LL^T with precision - can cause large gradients
687
+ )
688
+
689
+ # Better: LDL^T parameterization with separate diagonal
690
+ dist_stable = SparseMultivariateNormal(
691
+ loc=loc,
692
+ diagonal=diagonal, # Separate diagonal component for stability
693
+ precision_tril=unit_triangular_matrix # Unit triangular (LDL^T)
694
+ )
695
+ ```
696
+
697
+ **Benefits of LDL^T Parameterization**:
698
+ - **Numerical Stability**: Separates diagonal scaling from triangular structure
699
+ - **Gradient Stability**: More stable gradients during backpropagation
700
+ - **No SPD Constraints**: Doesn't require strict positive definiteness
701
+ - **Better Conditioning**: Diagonal component can be controlled independently
702
+
703
+ **Status**: This is a known limitation of the LL^T precision formulation. LDL^T parameterization is the recommended approach for precision matrices.