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];
}
}
Comments