File size: 4,305 Bytes
91e51cc | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 | 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))
|