ctgPi icon

Toeplitz rank

ctgPi | PRO | 04/17/21 01:10:03 PM UTC (Edited) | 0 ⭐ | 207 👁️ | Never ⏰ | []
text |

14.6 KB

|

None

|

0 👍

/

0 👎

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