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.
- {torchsparsegradutils-0.2.5/torchsparsegradutils.egg-info → torchsparsegradutils-0.2.6}/PKG-INFO +1 -1
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/docs/source/conf.py +2 -2
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/pyproject.toml +1 -1
- torchsparsegradutils-0.2.6/torchsparsegradutils/benchmarks/linear_cg_convergence.py +224 -0
- torchsparsegradutils-0.2.6/torchsparsegradutils/tests/test_linear_cg.py +366 -0
- torchsparsegradutils-0.2.6/torchsparsegradutils/tests/test_release_version.py +94 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/utils/__init__.py +2 -1
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/utils/linear_cg.py +218 -55
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6/torchsparsegradutils.egg-info}/PKG-INFO +1 -1
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils.egg-info/SOURCES.txt +2 -0
- torchsparsegradutils-0.2.5/torchsparsegradutils/tests/test_linear_cg.py +0 -104
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/LICENSE +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/MANIFEST.in +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/README.md +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/setup.cfg +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/setup.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/__init__.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/_compat.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/benchmarks/__init__.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/benchmarks/batched_sparse_mm_rand.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/benchmarks/benchmark_suite.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/benchmarks/benchmark_utils.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/benchmarks/sparse_bidir_logsumexp_rand.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/benchmarks/sparse_bidir_logsumexp_suitesparse.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/benchmarks/sparse_generic_solve_rand.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/benchmarks/sparse_generic_solve_suite.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/benchmarks/sparse_logsumexp_rand.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/benchmarks/sparse_logsumexp_suitesparse.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/benchmarks/sparse_mm_rand.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/benchmarks/sparse_mm_suite.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/benchmarks/sparse_triangular_solve_rand.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/benchmarks/sparse_triangular_solve_suitesparse.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/benchmarks/visualize_benchmark_results.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/cupy/__init__.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/cupy/cupy_bindings.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/cupy/cupy_sparse_solve.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/distributions/__init__.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/distributions/constraints.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/distributions/sparse_multivariate_normal.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/encoders/__init__.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/encoders/pairwise_encoder.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/encoders/pairwise_voxel_encoder.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/indexed_matmul.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/jax/__init__.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/jax/jax_bindings.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/jax/jax_sparse_solve.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/sparse_logsumexp.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/sparse_lstsq.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/sparse_matmul.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/sparse_solve.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/__init__.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/conftest.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_bicgstab.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_config.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_cupy_bindings.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_cupy_sparse_solve.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_deprecated_torch_apis.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_dist_stats_helpers.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_distributions.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_doctests.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_encoders.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_indexed_matmul.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_integration_pairwise_sparse_mvn.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_jax_bindings.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_jax_sparse_solve.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_lsmr.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_minres.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_params/czyx_shifts.yaml +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_params/pairwise_coo_indices.yaml +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_params/xyz_coords.yaml +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_quickstart_guide.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_random.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_sparse_bidir_logsumexp.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_sparse_logsumexp.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_sparse_lstsq.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_sparse_matmul.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_sparse_solve.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_sparse_triangular_solve.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/tests/test_utils.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/utils/bicgstab.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/utils/dist_stats_helpers.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/utils/lsmr.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/utils/minres.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/utils/random_sparse.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils/utils/utils.py +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils.egg-info/dependency_links.txt +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils.egg-info/not-zip-safe +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils.egg-info/requires.txt +0 -0
- {torchsparsegradutils-0.2.5 → torchsparsegradutils-0.2.6}/torchsparsegradutils.egg-info/top_level.txt +0 -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.
|
|
19
|
-
version = "0.2.
|
|
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
|
|
@@ -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)
|