torchsparsegradutils 0.2.5__tar.gz → 0.2.6__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 (89) hide show
  1. {torchsparsegradutils-0.2.5/torchsparsegradutils.egg-info → torchsparsegradutils-0.2.6}/PKG-INFO +1 -1
  2. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/docs/source/conf.py +2 -2
  3. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/pyproject.toml +1 -1
  4. torchsparsegradutils-0.2.6/torchsparsegradutils/benchmarks/linear_cg_convergence.py +224 -0
  5. torchsparsegradutils-0.2.6/torchsparsegradutils/tests/test_linear_cg.py +366 -0
  6. torchsparsegradutils-0.2.6/torchsparsegradutils/tests/test_release_version.py +94 -0
  7. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/utils/__init__.py +2 -1
  8. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/utils/linear_cg.py +218 -55
  9. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6/torchsparsegradutils.egg-info}/PKG-INFO +1 -1
  10. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils.egg-info/SOURCES.txt +2 -0
  11. torchsparsegradutils-0.2.5/torchsparsegradutils/tests/test_linear_cg.py +0 -104
  12. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/LICENSE +0 -0
  13. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/MANIFEST.in +0 -0
  14. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/README.md +0 -0
  15. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/setup.cfg +0 -0
  16. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/setup.py +0 -0
  17. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/__init__.py +0 -0
  18. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/_compat.py +0 -0
  19. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/benchmarks/__init__.py +0 -0
  20. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/benchmarks/batched_sparse_mm_rand.py +0 -0
  21. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/benchmarks/benchmark_suite.py +0 -0
  22. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/benchmarks/benchmark_utils.py +0 -0
  23. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/benchmarks/sparse_bidir_logsumexp_rand.py +0 -0
  24. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/benchmarks/sparse_bidir_logsumexp_suitesparse.py +0 -0
  25. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/benchmarks/sparse_generic_solve_rand.py +0 -0
  26. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/benchmarks/sparse_generic_solve_suite.py +0 -0
  27. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/benchmarks/sparse_logsumexp_rand.py +0 -0
  28. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/benchmarks/sparse_logsumexp_suitesparse.py +0 -0
  29. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/benchmarks/sparse_mm_rand.py +0 -0
  30. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/benchmarks/sparse_mm_suite.py +0 -0
  31. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/benchmarks/sparse_triangular_solve_rand.py +0 -0
  32. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/benchmarks/sparse_triangular_solve_suitesparse.py +0 -0
  33. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/benchmarks/visualize_benchmark_results.py +0 -0
  34. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/cupy/__init__.py +0 -0
  35. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/cupy/cupy_bindings.py +0 -0
  36. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/cupy/cupy_sparse_solve.py +0 -0
  37. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/distributions/__init__.py +0 -0
  38. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/distributions/constraints.py +0 -0
  39. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/distributions/sparse_multivariate_normal.py +0 -0
  40. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/encoders/__init__.py +0 -0
  41. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/encoders/pairwise_encoder.py +0 -0
  42. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/encoders/pairwise_voxel_encoder.py +0 -0
  43. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/indexed_matmul.py +0 -0
  44. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/jax/__init__.py +0 -0
  45. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/jax/jax_bindings.py +0 -0
  46. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/jax/jax_sparse_solve.py +0 -0
  47. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/sparse_logsumexp.py +0 -0
  48. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/sparse_lstsq.py +0 -0
  49. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/sparse_matmul.py +0 -0
  50. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/sparse_solve.py +0 -0
  51. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/__init__.py +0 -0
  52. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/conftest.py +0 -0
  53. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_bicgstab.py +0 -0
  54. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_config.py +0 -0
  55. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_cupy_bindings.py +0 -0
  56. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_cupy_sparse_solve.py +0 -0
  57. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_deprecated_torch_apis.py +0 -0
  58. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_dist_stats_helpers.py +0 -0
  59. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_distributions.py +0 -0
  60. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_doctests.py +0 -0
  61. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_encoders.py +0 -0
  62. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_indexed_matmul.py +0 -0
  63. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_integration_pairwise_sparse_mvn.py +0 -0
  64. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_jax_bindings.py +0 -0
  65. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_jax_sparse_solve.py +0 -0
  66. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_lsmr.py +0 -0
  67. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_minres.py +0 -0
  68. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_params/czyx_shifts.yaml +0 -0
  69. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_params/pairwise_coo_indices.yaml +0 -0
  70. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_params/xyz_coords.yaml +0 -0
  71. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_quickstart_guide.py +0 -0
  72. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_random.py +0 -0
  73. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_sparse_bidir_logsumexp.py +0 -0
  74. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_sparse_logsumexp.py +0 -0
  75. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_sparse_lstsq.py +0 -0
  76. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_sparse_matmul.py +0 -0
  77. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_sparse_solve.py +0 -0
  78. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_sparse_triangular_solve.py +0 -0
  79. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_utils.py +0 -0
  80. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/utils/bicgstab.py +0 -0
  81. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/utils/dist_stats_helpers.py +0 -0
  82. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/utils/lsmr.py +0 -0
  83. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/utils/minres.py +0 -0
  84. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/utils/random_sparse.py +0 -0
  85. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/utils/utils.py +0 -0
  86. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils.egg-info/dependency_links.txt +0 -0
  87. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils.egg-info/not-zip-safe +0 -0
  88. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils.egg-info/requires.txt +0 -0
  89. {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils.egg-info/top_level.txt +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: torchsparsegradutils
3
- Version: 0.2.5
3
+ Version: 0.2.6
4
4
  Summary: A collection of utility functions to work with PyTorch sparse tensors
5
5
  Author-email: CAI4CAI research group <contact@cai4cai.uk>
6
6
  License-Expression: Apache-2.0
@@ -15,8 +15,8 @@ sys.path.insert(0, os.path.abspath("../../"))
15
15
  project = "torchsparsegradutils"
16
16
  copyright = "2026, CAI4CAI research group"
17
17
  author = "CAI4CAI research group"
18
- release = "0.2.5"
19
- version = "0.2.5"
18
+ release = "0.2.6"
19
+ version = "0.2.6"
20
20
 
21
21
  # -- General configuration ---------------------------------------------------
22
22
  # https://www.sphinx-doc.org/en/master/usage/configuration.html#general-configuration
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
4
4
 
5
5
  [project]
6
6
  name = "torchsparsegradutils"
7
- version = "0.2.5"
7
+ version = "0.2.6"
8
8
  description = "A collection of utility functions to work with PyTorch sparse tensors"
9
9
  readme = "README.md"
10
10
  requires-python = ">=3.10"
@@ -0,0 +1,224 @@
1
+ #!/usr/bin/env python3
2
+ """Compare linear-CG convergence before and after the diagnostics fix.
3
+
4
+ This benchmark focuses on numerical behavior rather than throughput. It loads
5
+ the historical implementation from Git so both solvers run in one Python
6
+ environment on identical deterministic inputs.
7
+
8
+ Example
9
+ -------
10
+ python -m torchsparsegradutils.benchmarks.linear_cg_convergence \
11
+ --baseline-ref ea7b8f0 --repeats 100
12
+ """
13
+
14
+ from __future__ import annotations
15
+
16
+ import argparse
17
+ import json
18
+ import statistics
19
+ import subprocess
20
+ import time
21
+ import types
22
+ import warnings
23
+ from dataclasses import dataclass
24
+ from pathlib import Path
25
+ from typing import Callable
26
+
27
+ import torch
28
+
29
+ from torchsparsegradutils.utils import linear_cg
30
+
31
+ REPOSITORY_ROOT = Path(__file__).resolve().parents[2]
32
+ LINEAR_CG_PATH = "torchsparsegradutils/utils/linear_cg.py"
33
+
34
+
35
+ @dataclass(frozen=True)
36
+ class Problem:
37
+ name: str
38
+ matrix: torch.Tensor
39
+ rhs: torch.Tensor
40
+ tolerance: float
41
+ max_iter: int
42
+ initial_guess: torch.Tensor | None = None
43
+
44
+
45
+ class CountedMatmul:
46
+ def __init__(self, matrix: torch.Tensor):
47
+ self.matrix = matrix
48
+ self.calls = 0
49
+
50
+ def __call__(self, value: torch.Tensor) -> torch.Tensor:
51
+ self.calls += 1
52
+ return self.matrix @ value
53
+
54
+
55
+ def load_historical_linear_cg(revision: str) -> Callable[..., torch.Tensor]:
56
+ """Load ``linear_cg`` from a Git revision without changing the worktree."""
57
+ command = ["git", "-C", str(REPOSITORY_ROOT), "show", f"{revision}:{LINEAR_CG_PATH}"]
58
+ try:
59
+ completed = subprocess.run(command, check=True, capture_output=True, text=True)
60
+ except FileNotFoundError as error:
61
+ raise RuntimeError("Cannot load the historical solver because the Git executable was not found") from error
62
+ except subprocess.CalledProcessError as error:
63
+ detail = error.stderr.strip() or "Git could not resolve the requested revision and file"
64
+ raise RuntimeError(
65
+ f"Cannot load the historical solver from revision {revision!r}. "
66
+ f"Run this benchmark from a Git checkout containing that revision. Git reported: {detail}"
67
+ ) from error
68
+ module = types.ModuleType("historical_linear_cg")
69
+ exec(compile(completed.stdout, f"{revision}:{LINEAR_CG_PATH}", "exec"), module.__dict__)
70
+ return module.linear_cg
71
+
72
+
73
+ def make_problems(device: torch.device) -> list[Problem]:
74
+ dtype = torch.float64
75
+
76
+ size = 36
77
+ diagonal = torch.full((size,), 2.0, dtype=dtype, device=device)
78
+ off_diagonal = torch.full((size - 1,), -0.25, dtype=dtype, device=device)
79
+ tight_matrix = torch.diag(diagonal) + torch.diag(off_diagonal, diagonal=1) + torch.diag(off_diagonal, diagonal=-1)
80
+ tight_rhs = torch.zeros(size, dtype=dtype, device=device)
81
+ tight_rhs[size // 2] = 1
82
+
83
+ size = 40
84
+ multiple_matrix = torch.diag(torch.linspace(1.0, 100.0, size, dtype=dtype, device=device))
85
+ multiple_rhs = torch.zeros((size, 101), dtype=dtype, device=device)
86
+ multiple_rhs[:, -1] = 1
87
+
88
+ zero_matrix = torch.diag(torch.tensor([1.0, 2.0, 3.0], dtype=dtype, device=device))
89
+ zero_rhs = torch.zeros(3, dtype=dtype, device=device)
90
+
91
+ return [
92
+ Problem("tight_tolerance", tight_matrix, tight_rhs, tolerance=1e-12, max_iter=200),
93
+ Problem("multiple_rhs", multiple_matrix, multiple_rhs, tolerance=1e-4, max_iter=size),
94
+ Problem(
95
+ "zero_rhs_nonzero_guess",
96
+ zero_matrix,
97
+ zero_rhs,
98
+ tolerance=1e-5,
99
+ max_iter=20,
100
+ initial_guess=torch.ones_like(zero_rhs),
101
+ ),
102
+ ]
103
+
104
+
105
+ def synchronize(device: torch.device) -> None:
106
+ if device.type == "cuda":
107
+ torch.cuda.synchronize(device)
108
+
109
+
110
+ def relative_residual_per_rhs(matrix: torch.Tensor, solution: torch.Tensor, rhs: torch.Tensor) -> torch.Tensor:
111
+ if rhs.ndim == 1:
112
+ rhs = rhs.unsqueeze(-1)
113
+ solution = solution.unsqueeze(-1)
114
+ residual_norm = torch.linalg.vector_norm(rhs - matrix @ solution, dim=-2)
115
+ rhs_norm = torch.linalg.vector_norm(rhs, dim=-2)
116
+ return residual_norm / rhs_norm.masked_fill(rhs_norm.eq(0), 1)
117
+
118
+
119
+ def run_once(
120
+ solver: Callable[..., torch.Tensor],
121
+ problem: Problem,
122
+ *,
123
+ fixed: bool,
124
+ ) -> tuple[torch.Tensor, int, int, str]:
125
+ matmul = CountedMatmul(problem.matrix)
126
+ arguments = {
127
+ "tolerance": problem.tolerance,
128
+ "max_iter": problem.max_iter,
129
+ "initial_guess": problem.initial_guess,
130
+ }
131
+ with warnings.catch_warnings():
132
+ warnings.simplefilter("ignore", UserWarning)
133
+ if fixed:
134
+ solution, info = solver(matmul, problem.rhs, return_info=True, **arguments)
135
+ return solution, info.iterations, matmul.calls, info.reason
136
+ solution = solver(matmul, problem.rhs, **arguments)
137
+ # The historical solver performs one initial matvec followed by one per
138
+ # iteration and does not recompute the true residual before returning.
139
+ return solution, max(0, matmul.calls - 1), matmul.calls, "not_reported"
140
+
141
+
142
+ def measure(
143
+ solver: Callable[..., torch.Tensor],
144
+ problem: Problem,
145
+ *,
146
+ fixed: bool,
147
+ warmup: int,
148
+ repeats: int,
149
+ ) -> dict[str, object]:
150
+ for _ in range(warmup):
151
+ run_once(solver, problem, fixed=fixed)
152
+ synchronize(problem.rhs.device)
153
+
154
+ durations = []
155
+ for _ in range(repeats):
156
+ start = time.perf_counter()
157
+ run_once(solver, problem, fixed=fixed)
158
+ synchronize(problem.rhs.device)
159
+ durations.append(time.perf_counter() - start)
160
+
161
+ solution, iterations, matvecs, reason = run_once(solver, problem, fixed=fixed)
162
+ true_residual = relative_residual_per_rhs(problem.matrix, solution, problem.rhs)
163
+ reference = torch.linalg.solve(problem.matrix, problem.rhs)
164
+ error = torch.linalg.vector_norm(solution - reference) / torch.linalg.vector_norm(reference).clamp_min(1)
165
+
166
+ return {
167
+ "iterations": iterations,
168
+ "matvecs": matvecs,
169
+ "reason": reason,
170
+ "true_relative_residual_max": float(true_residual.max()),
171
+ "unconverged_rhs": int((true_residual > problem.tolerance).sum()),
172
+ "relative_solution_error": float(error),
173
+ "median_time_ms": statistics.median(durations) * 1e3,
174
+ }
175
+
176
+
177
+ def parse_args() -> argparse.Namespace:
178
+ parser = argparse.ArgumentParser(description=__doc__)
179
+ parser.add_argument("--baseline-ref", default="ea7b8f0", help="Git revision containing the historical solver")
180
+ parser.add_argument("--device", default="cpu", help="PyTorch device (default: cpu)")
181
+ parser.add_argument("--warmup", type=int, default=5)
182
+ parser.add_argument("--repeats", type=int, default=50)
183
+ arguments = parser.parse_args()
184
+ if arguments.warmup < 0 or arguments.repeats < 1:
185
+ parser.error("--warmup must be nonnegative and --repeats must be positive")
186
+ return arguments
187
+
188
+
189
+ def main() -> None:
190
+ arguments = parse_args()
191
+ device = torch.device(arguments.device)
192
+ historical_linear_cg = load_historical_linear_cg(arguments.baseline_ref)
193
+ results = {}
194
+ for problem in make_problems(device):
195
+ results[problem.name] = {
196
+ "baseline": measure(
197
+ historical_linear_cg,
198
+ problem,
199
+ fixed=False,
200
+ warmup=arguments.warmup,
201
+ repeats=arguments.repeats,
202
+ ),
203
+ "fixed": measure(
204
+ linear_cg,
205
+ problem,
206
+ fixed=True,
207
+ warmup=arguments.warmup,
208
+ repeats=arguments.repeats,
209
+ ),
210
+ }
211
+
212
+ payload = {
213
+ "baseline_ref": arguments.baseline_ref,
214
+ "device": str(device),
215
+ "torch_version": torch.__version__,
216
+ "warmup": arguments.warmup,
217
+ "repeats": arguments.repeats,
218
+ "results": results,
219
+ }
220
+ print(json.dumps(payload, indent=2, sort_keys=True))
221
+
222
+
223
+ if __name__ == "__main__":
224
+ main()
@@ -0,0 +1,366 @@
1
+ import pytest
2
+ import torch
3
+
4
+ from torchsparsegradutils.utils.linear_cg import linear_cg
5
+
6
+
7
+ # Test basic CG solve for vectors and matrices
8
+ def test_cg():
9
+ size = 100
10
+ # SPD matrix
11
+ matrix = torch.randn(size, size, dtype=torch.float64)
12
+ matrix = matrix.matmul(matrix.mT)
13
+ matrix.div_(torch.linalg.vector_norm(matrix)).add_(torch.eye(size, dtype=torch.float64) * 1e-1)
14
+ # single RHS
15
+ rhs = torch.randn(size, dtype=torch.float64)
16
+ solves = linear_cg(matrix.matmul, rhs=rhs, max_iter=size)
17
+ init = torch.randn(size, dtype=torch.float64)
18
+ solves_init = linear_cg(matrix.matmul, rhs=rhs, max_iter=size, initial_guess=init)
19
+ # reference solve
20
+ chol = torch.linalg.cholesky(matrix)
21
+ actual = torch.cholesky_solve(rhs.unsqueeze(1), chol).squeeze()
22
+ assert torch.allclose(solves, actual, atol=1e-3, rtol=1e-4)
23
+ assert torch.allclose(solves_init, actual, atol=1e-3, rtol=1e-4)
24
+ # multiple RHS
25
+ rhs_mat = torch.randn(size, 50, dtype=torch.float64)
26
+ solves = linear_cg(matrix.matmul, rhs=rhs_mat, max_iter=size)
27
+ init_mat = torch.randn(size, 50, dtype=torch.float64)
28
+ solves_init = linear_cg(matrix.matmul, rhs=rhs_mat, max_iter=size, initial_guess=init_mat)
29
+ actual_mat = torch.cholesky_solve(rhs_mat, chol)
30
+ assert torch.allclose(solves, actual_mat, atol=1e-3, rtol=1e-4)
31
+ assert torch.allclose(solves_init, actual_mat, atol=1e-3, rtol=1e-4)
32
+
33
+
34
+ # Test CG with tridiagonal outputs
35
+ def test_cg_with_tridiag():
36
+ size = 10
37
+ matrix = torch.randn(size, size, dtype=torch.float64)
38
+ matrix = matrix.matmul(matrix.mT)
39
+ matrix.div_(torch.linalg.vector_norm(matrix)).add_(torch.eye(size, dtype=torch.float64) * 1e-1)
40
+ rhs = torch.randn(size, 50, dtype=torch.float64)
41
+ solves, t_mats = linear_cg(
42
+ matrix.matmul, rhs=rhs, n_tridiag=5, max_tridiag_iter=10, max_iter=size, tolerance=0, eps=1e-15
43
+ )
44
+ chol = torch.linalg.cholesky(matrix)
45
+ actual = torch.cholesky_solve(rhs, chol)
46
+ assert torch.allclose(solves, actual, atol=1e-3, rtol=1e-4)
47
+ eigs = torch.linalg.eigvalsh(matrix)
48
+ for i in range(5):
49
+ approx = torch.linalg.eigvalsh(t_mats[i])
50
+ assert torch.allclose(eigs, approx, atol=1e-3, rtol=1e-4)
51
+
52
+
53
+ # Device parameterized CG tests
54
+ @pytest.mark.parametrize("batch", [None, 5])
55
+ def test_batch_cg(batch):
56
+ size = 100
57
+ shape = (batch, size, size) if batch else (size, size)
58
+ matrix = torch.randn(*shape, dtype=torch.float64)
59
+ matrix = matrix.matmul(matrix.mT)
60
+ matrix.div_(torch.linalg.vector_norm(matrix)).add_(torch.eye(size, dtype=torch.float64) * 1e-1)
61
+ b_shape = (batch, size, 50) if batch else (size, 50)
62
+ rhs = torch.randn(*b_shape, dtype=torch.float64)
63
+ solves = linear_cg(matrix.matmul, rhs=rhs, max_iter=size)
64
+ chol = torch.linalg.cholesky(matrix)
65
+ actual = torch.cholesky_solve(rhs, chol)
66
+ assert torch.allclose(solves, actual, atol=1e-3, rtol=1e-4)
67
+
68
+
69
+ @pytest.mark.parametrize("batch", [None, 5])
70
+ def test_batch_cg_with_tridiag(batch):
71
+ size = 10
72
+ shape = (batch, size, size) if batch else (size, size)
73
+ matrix = torch.randn(*shape, dtype=torch.float64)
74
+ matrix = matrix.matmul(matrix.mT)
75
+ matrix.div_(torch.linalg.vector_norm(matrix)).add_(torch.eye(size, dtype=torch.float64) * 1e-1)
76
+ b_shape = (batch, size, 10) if batch else (size, 10)
77
+ rhs = torch.randn(*b_shape, dtype=torch.float64)
78
+ solves, t_mats = linear_cg(
79
+ matrix.matmul, rhs=rhs, n_tridiag=8, max_iter=size, max_tridiag_iter=10, tolerance=0, eps=1e-30
80
+ )
81
+ chol = torch.linalg.cholesky(matrix)
82
+ actual = torch.cholesky_solve(rhs, chol)
83
+ assert torch.allclose(solves, actual, atol=1e-3, rtol=1e-4)
84
+ batch_dim = 5 if batch else 1
85
+ for i in range(batch_dim):
86
+ eigs = torch.linalg.eigvalsh(matrix[i] if batch else matrix)
87
+ for j in range(8):
88
+ approx = torch.linalg.eigvalsh(t_mats[j, i] if batch else t_mats[j])
89
+ assert torch.allclose(eigs, approx, atol=1e-3, rtol=1e-4)
90
+
91
+
92
+ # Test CG initialization reuse
93
+ def test_batch_cg_init():
94
+ batch = 5
95
+ size = 100
96
+ matrix = torch.randn(batch, size, size, dtype=torch.float64)
97
+ matrix = matrix.matmul(matrix.mT)
98
+ matrix.div_(torch.linalg.vector_norm(matrix)).add_(torch.eye(size, dtype=torch.float64) * 1e-1)
99
+ rhs = torch.randn(batch, size, 50, dtype=torch.float64)
100
+ solves = linear_cg(matrix.matmul, rhs=rhs, max_iter=size, max_tridiag_iter=0)
101
+ solves_init = linear_cg(matrix.matmul, rhs=rhs, max_iter=1, initial_guess=solves, max_tridiag_iter=0)
102
+ chol = torch.linalg.cholesky(matrix)
103
+ actual = torch.cholesky_solve(rhs, chol)
104
+ assert torch.allclose(solves_init, actual, atol=1e-3, rtol=1e-4)
105
+
106
+
107
+ def test_tight_tolerance_does_not_stall_at_historical_absolute_epsilon():
108
+ size = 36
109
+ diagonal = torch.full((size,), 2.0, dtype=torch.float64)
110
+ off_diagonal = torch.full((size - 1,), -0.25, dtype=torch.float64)
111
+ matrix = torch.diag(diagonal) + torch.diag(off_diagonal, 1) + torch.diag(off_diagonal, -1)
112
+ rhs = torch.zeros(size, dtype=torch.float64)
113
+ rhs[size // 2] = 1
114
+
115
+ solution, info = linear_cg(matrix, rhs, tolerance=1e-12, max_iter=200, return_info=True)
116
+ expected = torch.linalg.solve(matrix, rhs)
117
+
118
+ torch.testing.assert_close(solution, expected, rtol=1e-10, atol=1e-12)
119
+ assert info.reason == "converged"
120
+ assert info.converged.all()
121
+ assert info.true_relative_residual.max() <= 1e-12
122
+ assert info.iterations <= size
123
+
124
+
125
+ def test_all_rhs_convergence_does_not_hide_one_hard_column():
126
+ size = 40
127
+ matrix = torch.diag(torch.linspace(1.0, 100.0, size, dtype=torch.float64))
128
+ rhs = torch.zeros((size, 101), dtype=torch.float64)
129
+ rhs[:, -1] = 1
130
+
131
+ with pytest.warns(UserWarning):
132
+ _, mean_info = linear_cg(
133
+ matrix,
134
+ rhs,
135
+ tolerance=1e-4,
136
+ max_iter=size,
137
+ convergence_reduction="mean",
138
+ return_info=True,
139
+ )
140
+ all_solution, all_info = linear_cg(
141
+ matrix,
142
+ rhs,
143
+ tolerance=1e-4,
144
+ max_iter=size,
145
+ convergence_reduction="all",
146
+ return_info=True,
147
+ )
148
+
149
+ assert mean_info.reason == "mean_converged"
150
+ assert not mean_info.converged[..., -1].all()
151
+ assert all_info.reason == "converged"
152
+ assert all_info.converged.all()
153
+ assert all_info.true_relative_residual.max() <= 1e-4
154
+ torch.testing.assert_close(all_solution, torch.linalg.solve(matrix, rhs), rtol=2e-4, atol=1e-6)
155
+
156
+
157
+ def test_zero_rhs_returns_zero_even_with_nonzero_initial_guess():
158
+ matrix = torch.diag(torch.tensor([1.0, 2.0, 3.0], dtype=torch.float64))
159
+ rhs = torch.zeros(3, dtype=torch.float64)
160
+ initial_guess = torch.ones(3, dtype=torch.float64)
161
+
162
+ solution, info = linear_cg(matrix, rhs, initial_guess=initial_guess, tolerance=0, return_info=True)
163
+
164
+ torch.testing.assert_close(solution, torch.zeros_like(rhs), rtol=0, atol=0)
165
+ assert info.iterations == 0
166
+ assert info.converged.all()
167
+ assert info.true_relative_residual.max() == 0
168
+ assert info.tolerance == 0
169
+
170
+
171
+ def test_mixed_zero_and_nonzero_rhs_columns_are_supported():
172
+ matrix = torch.diag(torch.tensor([1.0, 2.0, 4.0], dtype=torch.float64))
173
+ rhs = torch.stack((torch.zeros(3, dtype=torch.float64), torch.ones(3, dtype=torch.float64)), dim=-1)
174
+
175
+ solution, info = linear_cg(matrix, rhs, tolerance=1e-12, return_info=True)
176
+
177
+ expected = torch.linalg.solve(matrix, rhs)
178
+ torch.testing.assert_close(solution, expected, rtol=1e-12, atol=1e-12)
179
+ torch.testing.assert_close(solution[:, 0], torch.zeros(3, dtype=torch.float64), rtol=0, atol=0)
180
+ assert info.converged.all()
181
+
182
+
183
+ def test_preconditioner_and_breakdown_validation():
184
+ diagonal = torch.tensor([1.0, 2.0, 4.0], dtype=torch.float64)
185
+ matrix = torch.diag(diagonal)
186
+ rhs = torch.ones(3, dtype=torch.float64)
187
+
188
+ solution, info = linear_cg(
189
+ matrix,
190
+ rhs,
191
+ preconditioner=lambda value: value / diagonal.unsqueeze(-1),
192
+ return_info=True,
193
+ )
194
+ torch.testing.assert_close(solution, torch.linalg.solve(matrix, rhs), rtol=1e-12, atol=1e-12)
195
+ assert info.reason == "converged"
196
+
197
+ indefinite = torch.diag(torch.tensor([1.0, -1.0, 2.0], dtype=torch.float64))
198
+ with pytest.raises(RuntimeError, match=r"p\^T A p"):
199
+ linear_cg(indefinite, torch.tensor([0.0, 1.0, 0.0], dtype=torch.float64))
200
+
201
+ with pytest.raises(RuntimeError, match="preconditioned residual inner product"):
202
+ linear_cg(matrix, rhs, preconditioner=lambda value: -value)
203
+
204
+
205
+ def test_user_eps_raises_instead_of_freezing_active_column():
206
+ matrix = torch.eye(3, dtype=torch.float64)
207
+ rhs = torch.ones(3, dtype=torch.float64)
208
+
209
+ with pytest.raises(RuntimeError, match="at least eps"):
210
+ linear_cg(
211
+ matrix,
212
+ rhs,
213
+ initial_guess=(1 - 1e-6) * rhs,
214
+ tolerance=1e-8,
215
+ eps=1e-10,
216
+ max_tridiag_iter=0,
217
+ )
218
+
219
+
220
+ def test_stopped_initial_updates_warn_and_report_reason():
221
+ matrix = torch.eye(3, dtype=torch.float64)
222
+ rhs = torch.ones(3, dtype=torch.float64)
223
+
224
+ with pytest.warns(UserWarning, match="did not converge"):
225
+ _, info = linear_cg(
226
+ matrix,
227
+ rhs,
228
+ initial_guess=0.5 * rhs,
229
+ tolerance=0.1,
230
+ stop_updating_after=0.6,
231
+ return_info=True,
232
+ )
233
+
234
+ assert info.iterations == 0
235
+ assert info.reason == "stopped_updating"
236
+ assert info.tolerance == 0.1
237
+ assert not info.converged.any()
238
+
239
+
240
+ def test_recursive_convergence_is_distinguished_from_true_convergence():
241
+ calls = 0
242
+
243
+ def matmul_with_final_residual_drift(value):
244
+ nonlocal calls
245
+ calls += 1
246
+ return 1.1 * value if calls == 3 else value
247
+
248
+ with pytest.warns(UserWarning, match="did not converge"):
249
+ _, info = linear_cg(
250
+ matmul_with_final_residual_drift,
251
+ torch.ones(3, dtype=torch.float64),
252
+ tolerance=1e-12,
253
+ max_tridiag_iter=0,
254
+ return_info=True,
255
+ )
256
+
257
+ assert info.iterations == 1
258
+ assert info.reason == "recursive_converged"
259
+ assert info.recursive_relative_residual.max() == 0
260
+ assert info.true_relative_residual.max() > info.tolerance
261
+
262
+
263
+ def test_tridiagonalization_can_return_info():
264
+ matrix = torch.diag(torch.tensor([1.0, 2.0, 3.0], dtype=torch.float64))
265
+ rhs = torch.ones(3, dtype=torch.float64)
266
+
267
+ solution, tridiagonal, info = linear_cg(
268
+ matrix,
269
+ rhs,
270
+ n_tridiag=1,
271
+ tolerance=0,
272
+ max_iter=3,
273
+ max_tridiag_iter=3,
274
+ return_info=True,
275
+ )
276
+
277
+ torch.testing.assert_close(solution, torch.linalg.solve(matrix, rhs), rtol=1e-12, atol=1e-12)
278
+ assert tridiagonal.shape == (1, 3, 3)
279
+ assert info.iterations == 3
280
+ assert info.matvecs == 5
281
+
282
+
283
+ @pytest.mark.parametrize(
284
+ ("kwargs", "match"),
285
+ [
286
+ ({"tolerance": float("nan")}, "tolerance"),
287
+ ({"tolerance": float("inf")}, "tolerance"),
288
+ ({"tolerance": -1.0}, "tolerance"),
289
+ ({"eps": float("nan")}, "eps"),
290
+ ({"eps": float("inf")}, "eps"),
291
+ ({"eps": 0.0}, "eps"),
292
+ ({"stop_updating_after": float("nan")}, "stop_updating_after"),
293
+ ({"stop_updating_after": float("inf")}, "stop_updating_after"),
294
+ ({"stop_updating_after": -1.0}, "stop_updating_after"),
295
+ ({"convergence_reduction": "median"}, "convergence_reduction"),
296
+ ({"min_iter": -1}, "min_iter"),
297
+ ],
298
+ )
299
+ def test_invalid_solver_settings_raise(kwargs, match):
300
+ with pytest.raises(ValueError, match=match):
301
+ linear_cg(torch.eye(2, dtype=torch.float64), torch.ones(2, dtype=torch.float64), **kwargs)
302
+
303
+
304
+ def test_eps_must_be_representable_in_rhs_dtype():
305
+ with pytest.raises(ValueError, match="representable"):
306
+ linear_cg(torch.eye(2), torch.ones(2), eps=1e-100)
307
+
308
+
309
+ def test_rhs_and_operator_output_validation():
310
+ matrix = torch.eye(2, dtype=torch.float64)
311
+ rhs = torch.ones(2, dtype=torch.float64)
312
+
313
+ with pytest.raises(TypeError, match="floating-point"):
314
+ linear_cg(torch.eye(2, dtype=torch.int64), torch.ones(2, dtype=torch.int64))
315
+ with pytest.raises(RuntimeError, match="matmul_closure output"):
316
+ linear_cg(lambda value: value[:-1], rhs)
317
+ with pytest.raises(RuntimeError, match="preconditioner output"):
318
+ linear_cg(matrix, rhs, preconditioner=lambda value: value.to(torch.float32))
319
+
320
+
321
+ def test_vector_rank_is_preserved_with_vector_initial_guess():
322
+ matrix = torch.diag(torch.tensor([1.0, 2.0, 3.0], dtype=torch.float64))
323
+ rhs = torch.ones(3, dtype=torch.float64)
324
+ solution = linear_cg(matrix, rhs, initial_guess=torch.zeros_like(rhs))
325
+ assert solution.shape == rhs.shape
326
+
327
+ column_guess_solution = linear_cg(matrix, rhs, initial_guess=torch.zeros((3, 1), dtype=torch.float64))
328
+ assert column_guess_solution.shape == rhs.shape
329
+
330
+ with pytest.raises(ValueError, match="initial_guess must have shape"):
331
+ linear_cg(matrix, rhs, initial_guess=torch.zeros((3, 2), dtype=torch.float64))
332
+
333
+
334
+ def test_min_iter_is_honored_for_converged_initial_guess():
335
+ matrix = torch.eye(3, dtype=torch.float64)
336
+ rhs = torch.ones(3, dtype=torch.float64)
337
+
338
+ solution, info = linear_cg(
339
+ matrix,
340
+ rhs,
341
+ initial_guess=rhs,
342
+ min_iter=3,
343
+ max_iter=5,
344
+ max_tridiag_iter=0,
345
+ return_info=True,
346
+ )
347
+
348
+ torch.testing.assert_close(solution, rhs, rtol=0, atol=0)
349
+ assert info.iterations == 3
350
+ assert info.matvecs == 5
351
+ assert info.reason == "converged"
352
+ assert info.converged.all()
353
+
354
+
355
+ def test_multiple_batch_dimensions_are_supported():
356
+ with pytest.raises(ValueError, match="at least one dimension"):
357
+ linear_cg(torch.ones((1, 1), dtype=torch.float64), torch.tensor(1.0, dtype=torch.float64))
358
+
359
+ matrix = torch.diag(torch.tensor([1.0, 2.0, 3.0], dtype=torch.float64))
360
+ rhs = torch.ones((2, 4, 3, 2), dtype=torch.float64)
361
+
362
+ solution = linear_cg(matrix, rhs, tolerance=1e-12)
363
+ expected = torch.linalg.solve(matrix, rhs)
364
+
365
+ assert solution.shape == rhs.shape
366
+ torch.testing.assert_close(solution, expected, rtol=1e-12, atol=1e-12)