| 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)] |
| 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)) |
|
|