exact-solve / tests /test_exact_solve.py
phanerozoic's picture
Upload folder using huggingface_hub
91e51cc verified
Raw
History Blame
4.31 kB
import math
import random
from fractions import Fraction
import pytest
import torch
import kernels
es = kernels.get_kernel("phanerozoic/exact-solve", version=1, trust_remote_code=True)
requires_cuda = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required")
def randmat(rows, cols, bits, rng):
hi = 1 << bits
return [[rng.randint(-hi, hi) for _ in range(cols)] for _ in range(rows)]
def frac_solve(A, B):
"""Exact Gaussian elimination over Q (oracle)."""
n, c = len(A), len(B[0])
M = [[Fraction(A[i][j]) for j in range(n)] + [Fraction(B[i][k]) for k in range(c)]
for i in range(n)]
for col in range(n):
piv = next((r for r in range(col, n) if M[r][col] != 0), None)
if piv is None:
return None
M[col], M[piv] = M[piv], M[col]
inv = M[col][col]
M[col] = [v / inv for v in M[col]]
for r in range(n):
if r != col and M[r][col] != 0:
f = M[r][col]
M[r] = [M[r][k] - f * M[col][k] for k in range(n + c)]
return [[M[i][n + k] for k in range(c)] for i in range(n)]
def certifies(A, B, X):
"""Exact integer identity A @ N == d * B, d the common denominator."""
n, c = len(A), len(B[0])
d = 1
for i in range(n):
for j in range(c):
d = d * X[i][j].denominator // math.gcd(d, X[i][j].denominator)
N = [[int(X[i][j] * d) for j in range(c)] for i in range(n)]
for i in range(n):
for j in range(c):
if sum(A[i][k] * N[k][j] for k in range(n)) != d * B[i][j]:
return False
return True
def _cuda(A, B):
return (torch.tensor(A, dtype=torch.int64, device="cuda"),
torch.tensor(B, dtype=torch.int64, device="cuda"))
@requires_cuda
@pytest.mark.kernels_ci
@pytest.mark.parametrize("n,bits,c", [(8, 12, 1), (16, 20, 1), (24, 24, 2), (32, 16, 3)])
def test_matches_fraction_oracle(n, bits, c):
rng = random.Random(n * 100 + c)
A = randmat(n, n, bits, rng)
B = randmat(n, c, bits, rng)
ref = frac_solve(A, B)
if ref is None:
pytest.skip("singular draw")
At, Bt = _cuda(A, B)
X = es.solve(At, Bt)
assert X is not None
assert all(X[i][j] == ref[i][j] for i in range(n) for j in range(c))
@requires_cuda
@pytest.mark.kernels_ci
def test_certificate_larger():
"""Exact A@N == d*B at a larger size / wider entries, no external oracle."""
rng = random.Random(7)
n, c = 64, 2
A = randmat(n, n, 40, rng)
B = randmat(n, c, 40, rng)
At, Bt = _cuda(A, B)
X = es.solve(At, Bt)
assert X is not None and certifies(A, B, X)
@requires_cuda
@pytest.mark.kernels_ci
def test_planted_integer_solution():
"""b = A @ x for integer x: the exact solution must recover x."""
rng = random.Random(3)
n = 48
A = randmat(n, n, 20, rng)
x = [rng.randint(-500, 500) for _ in range(n)]
b = [[sum(A[i][k] * x[k] for k in range(n))] for i in range(n)]
At, Bt = _cuda(A, b)
X = es.solve(At, Bt)
assert X is not None
assert all(X[i][0] == Fraction(x[i]) for i in range(n))
@requires_cuda
@pytest.mark.kernels_ci
def test_fractional_solution_dense_denominators():
"""Scaled Hilbert-like system with genuinely fractional answers."""
rng = random.Random(5)
n = 20
L = 1
for k in range(1, 2 * n + 1):
L = L * k // math.gcd(L, k)
A = [[L // (i + j + 1) for j in range(n)] for i in range(n)]
B = [[rng.randint(-100, 100)] for _ in range(n)]
ref = frac_solve(A, B)
At, Bt = _cuda(A, B)
X = es.solve(At, Bt)
assert X is not None and all(X[i][0] == ref[i][0] for i in range(n))
assert certifies(A, B, X)
@requires_cuda
@pytest.mark.kernels_ci
def test_singular_returns_none():
rng = random.Random(9)
n = 16
A = randmat(n, n, 16, rng)
A[n - 1] = [2 * A[0][j] - A[1][j] for j in range(n)] # dependent row
B = randmat(n, 1, 16, rng)
At, Bt = _cuda(A, B)
assert es.solve(At, Bt) is None
@requires_cuda
@pytest.mark.kernels_ci
def test_deterministic():
rng = random.Random(11)
A = randmat(32, 32, 24, rng)
B = randmat(32, 4, 24, rng)
At, Bt = _cuda(A, B)
X1, X2 = es.solve(At, Bt), es.solve(At, Bt)
assert all(X1[i][j] == X2[i][j] for i in range(32) for j in range(4))