const std = @import("std"); const ROUND_MAGIC: f64 = 6755399441055744.0; const z = 5; fn pow(b: usize, m: usize) usize { var result: usize = 1; var i: usize = 0; while (i < m) : (i += 1) { result *= b; } return result; } const r: usize = (pow(3, z - 1) + 1) / 2; const VECTOR_LENGTH = 2; fn vec(comptime T: type) type { return @Vector(VECTOR_LENGTH, T); } fn splat(scalar: anytype) vec(@TypeOf(scalar)) { return @splat(VECTOR_LENGTH, scalar); } // FIXME: switch to @intToFloat once Zig ships it fn intVecToFloatVec(comptime R: type, comptime T: type, x: @Vector(VECTOR_LENGTH, T)) vec(R) { var result: [VECTOR_LENGTH]R = undefined; for (result) |*it, i| { it.* = @intToFloat(R, x[i]); } return @as(vec(R), result); } // FIXME: switch to @intCast once Zig ships it fn intVecCast(comptime R: type, comptime T: type, x: @Vector(VECTOR_LENGTH, T)) vec(R) { var result: [VECTOR_LENGTH]R = undefined; for (result) |*it, i| { it.* = @intCast(R, x[i]); } return @as(vec(R), result); } // FIXME: switch to @blend once Zig ships it fn blend(a: vec(f64), b: vec(f64), mask: vec(i64)) vec(f64) { return @bitCast(vec(f64), (mask & @bitCast(vec(i64), a)) | (~mask & @bitCast(vec(i64), b))); } // FIXME: switch to @boolToInt once Zig ships it fn boolVecToIntVec(comptime R: type, x: @Vector(VECTOR_LENGTH, bool)) vec(R) { var result: [VECTOR_LENGTH]R = undefined; for (result) |*it, i| { if (x[i]) { it.* = ~@as(R, 0); } else { it.* = @as(R, 0); } } return @as(vec(R), result); } const RingElement = struct { value: vec(f64), }; const Ring = struct { inverse_modulus: vec(f64), modulus: vec(f64), inverse_exponent: vec(i64), pub fn init(modulus: vec(u32)) Ring { return Ring{ .modulus = intVecToFloatVec(f64, u32, modulus), .inverse_exponent = intVecCast(i64, u32, modulus) - splat(@intCast(i64, 2)), // TODO: assert modulus is prime .inverse_modulus = splat(@floatCast(f64, 1.0)) / intVecToFloatVec(f64, u32, modulus), }; } // TODO: support more int types? pub fn reduce(self: *const Ring, x: i32) RingElement { var result = intVecToFloatVec(f64, i32, splat(x)); var q = result * self.inverse_modulus; q += splat(ROUND_MAGIC); q -= splat(ROUND_MAGIC); result -= q * self.modulus; return RingElement{ .value = result, }; } pub fn mul(self: *const Ring, a: RingElement, b: RingElement) RingElement { var result = a.value * b.value; var q = result * self.inverse_modulus; q += splat(ROUND_MAGIC); q -= splat(ROUND_MAGIC); result -= q * self.modulus; return RingElement{ .value = result, }; } pub fn fma(self: *const Ring, a: RingElement, b: RingElement, c: RingElement) RingElement { var result = a.value + b.value * c.value; var q = result * self.inverse_modulus; q += splat(ROUND_MAGIC); q -= splat(ROUND_MAGIC); result -= q * self.modulus; return RingElement{ .value = result, }; } pub fn fms(self: *const Ring, a: RingElement, b: RingElement, c: RingElement) RingElement { var result = a.value - b.value * c.value; var q = result * self.inverse_modulus; q += splat(ROUND_MAGIC); q -= splat(ROUND_MAGIC); result -= q * self.modulus; return RingElement{ .value = result, }; } pub fn inv(self: *const Ring, a: RingElement) RingElement { var result = self.reduce(1); // TODO: inline? var i: u32 = 0; var e = self.inverse_exponent; // TODO: skip zeros at beginning while (i < 64) : ({ i += 1; e <<= splat(@as(u6, 1)); }) { result = self.mul(result, result); const b = RingElement{ .value = blend(a.value, splat(@as(f64, 1.0)), e >> splat(@as(u6, 63))), }; result = self.mul(result, b); } return result; } }; fn runTests() !void { const p: u32 = 67108859; const q: u32 = 67108837; const ring = Ring.init([_]u32{ p, q }); std.debug.warn("modulus = {}\n", .{ring.modulus}); std.debug.warn("inverse_modulus = {}\n", .{ring.inverse_modulus}); std.debug.warn("inverse_exponent = {}\n", .{ring.inverse_exponent}); { const testCases = [_]i32{ -1, 0, 1, (p - 1) / 2, (p + 1) / 2, p - 1, p, p + 1 }; for (testCases) |item| { std.debug.warn("reduce({}) = {}\n", .{ item, ring.reduce(item) }); } } { const testCases = [_][2]i32{ [_]i32{ -1, -1 }, [_]i32{ -1, 0 }, [_]i32{ 0, 0 }, [_]i32{ 1, 0 }, [_]i32{ 1, 1 }, [_]i32{ 2, 2 }, [_]i32{ p, 1 }, [_]i32{ p, 0 }, [_]i32{ p, p }, [_]i32{ (p - 1) / 2, 2 }, [_]i32{ (p + 1) / 2, 2 }, }; for (testCases) |item| { std.debug.warn("mul({}, {}) = {}\n", .{ item[0], item[1], ring.mul(ring.reduce(item[0]), ring.reduce(item[1])) }); } } { const testCases = [_]i32{ -1, 0, 1, (p - 1) / 2, (p + 1) / 2, p - 1, p, p + 1, (q - 1) / 2, (q + 1) / 2, q - 1, q, q + 1 }; for (testCases) |item| { std.debug.warn("inv({}) = {}\n", .{ item, ring.fms(ring.reduce(0), ring.reduce(item), ring.inv(ring.reduce(item))) }); } } } fn doGaussianEliminationStep(ring: Ring, M: [][2 * z + 1]RingElement, result: *DeterminantResult) void { const pivot = M[0][z]; const inverse_pivot = ring.inv(pivot); // TODO: document why pivot == 0 still works var d: usize = 1; while (d <= z) : (d += 1) { if (d >= M.len) { break; } const t = ring.mul(M[d][z - d], inverse_pivot); var j: usize = 0; while (j <= z) : (j += 1) { M[d][z - d + j] = ring.fms(M[d][z - d + j], M[0][z + j], t); } } const mask = boolVecToIntVec(i64, pivot.value == splat(@as(f64, 0.0))); result.determinant = ring.mul(result.determinant, RingElement{ .value = blend(splat(@as(f64, 1.0)), pivot.value, mask) }); result.corank += @bitCast(vec(usize), -mask); } const DeterminantResult = struct { determinant: RingElement, corank: vec(usize), // TODO: move to Ring pub fn mul(self: *const DeterminantResult, ring: Ring, other: DeterminantResult) DeterminantResult { return DeterminantResult{ .determinant = ring.mul(self.determinant, other.determinant), .corank = self.corank + other.corank, }; } }; fn findDeterminant(ring: Ring, M: [][2 * z + 1]RingElement) DeterminantResult { // TODO: document (e.g. k-diagonal layout, skips last row) var result = DeterminantResult{ .determinant = ring.reduce(1), // TODO: inline .corank = splat(@as(usize, 0)), }; var k: usize = 0; while (k < M.len) : (k += 1) { doGaussianEliminationStep(ring, M[k..], &result); } return result; } var L: [2 * r + 1]RingElement = undefined; // Forward/backward vectors // TODO: keep only two values? var u: [3][r + 1]RingElement = undefined; var v: [3][r + 1]RingElement = undefined; pub fn main() !void { const ring = Ring.init([_]u32{ 67108859, 67108837 }); { // Compute determinant for small cases. var M: [2 * z][2 * z + 1]RingElement = undefined; var n: usize = 0; while (n < 2 * z) : (n += 1) { // N[i][z+k] == M[i][i+k] var i: usize = 0; while (i < n) : (i += 1) { // FIXME: avoid double initialization? var j: usize = 0; while (j < 2 * z + 1) : (j += 1) { M[i][j] = ring.reduce(0); } M[i][z] = ring.reduce(@intCast(i32, (if (i < z) i else z) + (if (n - 1 - i < z) n - 1 - i else z))); var d: usize = 1; while (d <= z) : (d += 1) { if (i >= d) { M[i][z - d] = ring.reduce(-1); } if (i + d < n) { M[i][z + d] = ring.reduce(-1); } } } const det = findDeterminant(ring, M[0..n]); if (n >= z) { L[n - z] = det.determinant; // TODO: .corank? } } } var head: [z + 1][2 * z + 1]RingElement = undefined; { var i: usize = 0; while (i < z) : (i += 1) { // TODO: avoid double initialization var j: usize = 0; while (j < 2 * z + 1) : (j += 1) { head[i][j] = ring.reduce(0); } head[i][z] = ring.reduce(@intCast(i32, (if (i < z) i else z) + z)); var d: usize = 1; while (d <= z) : (d += 1) { if (i >= d) { head[i][z - d] = ring.reduce(-1); } head[i][z + d] = ring.reduce(-1); } } // Last row is a temporary buffer. } var tail: [z][2 * z + 1]RingElement = undefined; { var i: usize = 0; while (i < z) : (i += 1) { // TODO: avoid double initialization var j: usize = 0; while (j < 2 * z + 1) : (j += 1) { tail[i][j] = ring.reduce(0); } tail[i][z] = ring.reduce(@intCast(i32, 2 * z - 1 - (if (i < z) i else z))); var d: usize = 1; while (d <= z) : (d += 1) { tail[i][z - d] = ring.reduce(-1); if (i + d < z) { tail[i][z + d] = ring.reduce(-1); } } } } { var prefix: DeterminantResult = DeterminantResult{ .determinant = ring.reduce(1), // TODO: inline? .corank = splat(@as(usize, 0)), }; var n: usize = 2 * z; while (n < z + 1 + 2 * r) : (n += 1) { var M: [2 * z][2 * z + 1]RingElement = undefined; { var i: usize = 0; while (i < z) : (i += 1) { var j: usize = 0; while (j < 2 * z + 1) : (j += 1) { M[i][j] = head[i][j]; } } } { var i: usize = 0; while (i < z) : (i += 1) { var j: usize = 0; while (j < 2 * z + 1) : (j += 1) { M[i + z][j] = tail[i][j]; } } } const det = prefix.mul(ring, findDeterminant(ring, M[0..])); L[n - z] = det.determinant; // TODO: .corank? { // TODO: avoid double initialization var j: usize = 0; while (j < 2 * z + 1) : (j += 1) { head[z][j] = ring.reduce(-1); } head[z][z] = ring.reduce(@intCast(i32, 2 * z)); } doGaussianEliminationStep(ring, head[0..], &prefix); { var i: usize = 0; while (i < z) : (i += 1) { head[i] = head[i + 1]; } } } } std.debug.warn("r = {}\n", .{r}); if (false) { var i: usize = 0; while (i < 2 * r + 1) : (i += 1) { std.debug.warn("det = {}\n", .{L[i]}); } } // TODO: explain how M is defined in terms of L var corank: vec(usize) = splat(@as(usize, 0)); u[0][0] = ring.inv(L[r]); v[0][0] = ring.inv(L[r]); var i: usize = 1; while (i < r + 1) : (i += 1) { // TODO: indices // M[][] * u = [1 0 ... 0 0] // M[][] * v = [0 0 ... 0 1] // w = a * [0 u] + b * [v 0] // L[r-i] // L[r+i] // M[][] * [0 u] = [1 0 ... 0 p] var p = ring.reduce(0); { var k: usize = 0; while (k < i) : (k += 1) { p = ring.fma(p, L[r + 1 + k], u[0][k]); } } // M[][] * [v 0] = [q 0 ... 0 1] var q = ring.reduce(0); { var k: usize = 0; while (k < i) : (k += 1) { q = ring.fma(q, L[r - i + k], v[0][k]); } } const D = ring.fms(ring.reduce(1), p, q); const mask = boolVecToIntVec(i64, D.value == splat(@as(f64, 0.0))); corank += @bitCast(vec(usize), -mask); const iD = ring.inv(D); { const a = ring.mul(ring.reduce(1), iD); const b = ring.fms(ring.reduce(0), p, iD); { var k: usize = 0; while (k <= i) : (k += 1) { u[1][k] = ring.reduce(0); } } { var k: usize = 0; while (k <= i) : (k += 1) { if (k > 0) { u[1][k] = ring.fma(u[1][k], a, u[0][k - 1]); } if (k < i) { u[1][k] = ring.fma(u[1][k], b, v[0][k]); } } } } { const a = ring.fms(ring.reduce(0), q, iD); const b = ring.mul(ring.reduce(1), iD); { var k: usize = 0; while (k <= i) : (k += 1) { v[1][k] = ring.reduce(0); } } { var k: usize = 0; while (k <= i) : (k += 1) { if (k > 0) { v[1][k] = ring.fma(v[1][k], a, u[0][k - 1]); } if (k < i) { v[1][k] = ring.fma(v[1][k], b, v[0][k]); } } } } if (i == r - 1 or i == r) { var j: usize = 0; std.debug.warn("corank[{}] = {}\n", .{ i, corank }); } // TODO: swap pointers u[2] = u[0]; u[0] = u[1]; u[1] = u[2]; v[2] = v[0]; v[0] = v[1]; v[1] = v[2]; } }