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