const std = @import("std"); pub fn Matrix(rows: comptime_int, cols: comptime_int, T: type) type { return struct { mat: [rows][cols]T, comptime rows: comptime_int = rows, comptime cols: comptime_int = cols, const Self = @This(); pub const zero: Self = .splat(0); pub const ident: Self = blk: { var m: [rows][cols]T = undefined; for (0..rows) |i| { for (0..cols) |j| { m[i][j] = if (i == j) 1 else 0; } } break :blk .{ .mat = m }; }; const vector = struct { fn get(self: Self, comptime s: Swizzle) T { return self.mat[s.index()][0]; } fn swizzle(self: Self, comptime select: anytype) Vector(select.len, T) { const sel: [select.len]Swizzle = select; var result: Vector(select.len, T) = .{ .mat = undefined }; inline for (&result.mat, sel) |*r, s| r[0] = self.get(s); return result; } fn dot(self: Self, other: Self) T { return self.t().mul(other).get(.x); } fn array(self: Self) [self.rows]T { return @bitCast(self.mat); } }; pub const get = if (cols == 1) vector.get else @compileError("Not a vector"); pub const swizzle = if (cols == 1) vector.swizzle else @compileError("Not a vector"); pub const dot = if (cols == 1) vector.dot else @compileError("Not a vector"); pub const array = if (cols == 1) vector.array else @compileError("Not a vector"); pub fn splat(s: T) Self { var result: Self = .{ .mat = undefined }; inline for (&result.mat) |*row| { inline for (row) |*r| r.* = s; } return result; } pub fn t(self: Self) Matrix(self.cols, self.rows, T) { var result: Matrix(self.cols, self.rows, T) = .{ .mat = undefined }; inline for (0..self.rows) |i| { inline for (0..self.cols) |j| result.mat[j][i] = self.mat[i][j]; } return result; } pub fn scale(self: Self, scalar: T) Self { var result: Self = .{ .mat = undefined }; inline for (&result.mat, self.mat) |*r_row, s_row| { inline for (r_row, s_row) |*r, s| { r.* = s * scalar; } } return result; } pub fn scaleInv(self: Self, scalar: T) Self { var result: Self = .{ .mat = undefined }; inline for (&result.mat, self.mat) |*r_row, s_row| { inline for (r_row, s_row) |*r, s| { r.* = s / scalar; } } return result; } pub fn neg(self: Self) Self { var result: Self = .{ .mat = undefined }; inline for (&result.mat, self.mat) |*r_row, s_row| { inline for (r_row, s_row) |*r, s| { r.* = -s; } } return result; } pub fn abs(self: Self) Self { var result: Self = .{ .mat = undefined }; inline for (&result.mat, self.mat) |*r_row, s_row| { inline for (r_row, s_row) |*r, s| { r.* = @abs(s); } } return result; } pub const BinOp = struct { f: @TypeOf(impl.add), const impl = struct { inline fn add(a: T, b: T) T { return a + b; } inline fn addWrap(a: T, b: T) T { return a +% b; } inline fn addSat(a: T, b: T) T { return a +| b; } inline fn sub(a: T, b: T) T { return a - b; } inline fn subWrap(a: T, b: T) T { return a -% b; } inline fn mul(a: T, b: T) T { return a * b; } inline fn mulWrap(a: T, b: T) T { return a *% b; } inline fn mulSat(a: T, b: T) T { return a *% b; } inline fn div(a: T, b: T) T { return a / b; } inline fn mod(a: T, b: T) T { return a % b; } inline fn bitAnd(a: T, b: T) T { return a & b; } inline fn bitOr(a: T, b: T) T { return a | b; } inline fn bitXOr(a: T, b: T) T { return a ^ b; } }; pub const add: BinOp = .{ .f = impl.add }; pub const addWrap: BinOp = .{ .f = impl.addWrap }; pub const addSat: BinOp = .{ .f = impl.addSat }; pub const sub: BinOp = .{ .f = impl.sub }; pub const subWrap: BinOp = .{ .f = impl.subWrap }; pub const mul: BinOp = .{ .f = impl.mul }; pub const mulWrap: BinOp = .{ .f = impl.mulWrap }; pub const mulSat: BinOp = .{ .f = impl.mulSat }; pub const div: BinOp = .{ .f = impl.div }; pub const mod: BinOp = .{ .f = impl.mod }; pub const bitAnd: BinOp = .{ .f = impl.bitAnd }; pub const bitOr: BinOp = .{ .f = impl.bitOr }; pub const bitXOr: BinOp = .{ .f = impl.bitXOr }; }; pub fn bin(self: Self, other: Self, comptime op: BinOp) Self { if (self.rows != other.rows or self.cols != other.cols) @compileError("Cannot operate on matrices of incompatible sizes"); var result: Self = .{ .mat = undefined }; inline for (&result.mat, self.mat, other.mat) |*r_row, s_row, o_row| { inline for (r_row, s_row, o_row) |*r, s, o| r.* = op.f(s, o); } return result; } pub fn add(self: Self, other: Self) Self { return self.bin(other, .add); } pub fn sub(self: Self, other: Self) Self { return self.bin(other, .sub); } pub fn mulElements(self: Self, other: Self) Self { return self.bin(other, .mul); } pub fn div(self: Self, other: Self) Self { return self.bin(other, .div); } pub const CompareOp = struct { f: @TypeOf(impl.eq), const impl = struct { inline fn eq(a: T, b: T) bool { return a == b; } inline fn ne(a: T, b: T) bool { return a != b; } inline fn gt(a: T, b: T) bool { return a > b; } inline fn gte(a: T, b: T) bool { return a >= b; } inline fn lt(a: T, b: T) bool { return a < b; } inline fn lte(a: T, b: T) bool { return a <= b; } }; pub const eq: CompareOp = .{ .f = impl.eq }; pub const ne: CompareOp = .{ .f = impl.ne }; pub const gt: CompareOp = .{ .f = impl.gt }; pub const gte: CompareOp = .{ .f = impl.gte }; pub const lt: CompareOp = .{ .f = impl.lt }; pub const lte: CompareOp = .{ .f = impl.lte }; }; pub fn all(self: Self, other: Self, comptime op: CompareOp) bool { if (self.rows != other.rows or self.cols != other.cols) @compileError("Cannot compare matrices of incompatible sizes"); inline for (self.mat, other.mat) |s_row, o_row| { inline for (s_row, o_row) |s, o| { if (!op.f(s, o)) return false; } } return true; } pub fn any(self: Self, other: Self, comptime op: CompareOp) bool { if (self.rows != other.rows or self.cols != other.cols) @compileError("Cannot compare matrices of incompatible sizes"); inline for (self.mat, other.mat) |s_row, o_row| { inline for (s_row, o_row) |s, o| { if (op.f(s, o)) return true; } } return false; } pub fn eq(self: Self, other: Self) bool { return self.all(other, .eq); } pub fn mul(self: Self, other: anytype) Matrix(self.rows, other.cols, T) { if (other.rows != self.cols) @compileError(std.fmt.comptimePrint( "Cannot multiply {}x{} matrix by {}x{} matrix", .{ self.rows, self.cols, other.rows, other.cols }, )); var result: Matrix(self.rows, other.cols, T) = .{ .mat = undefined }; inline for (0..result.rows) |i| { inline for (0..result.cols) |j| { var sum: T = 0; inline for (0..self.cols) |k| { switch (@typeInfo(T)) { .float => sum = @mulAdd(T, self.mat[i][k], other.mat[k][j], sum), else => sum += self.mat[i][k] * other.mat[k][j], } } result.mat[i][j] = sum; } } return result; } pub fn normSquared(self: Self) T { var n: T = 0; inline for (self.mat) |row| { inline for (row) |e| { switch (@typeInfo(T)) { .float => n = @mulAdd(T, e, e, n), else => n += e * e, } } } return n; } pub const norm = switch (@typeInfo(T)) { .float => struct { fn norm(self: Self) T { return @sqrt(self.normSquared()); } }.norm, else => @compileError("Connot take a square root on an integer matrix"), }; pub fn floatFromInt(self: Self, F: type) Matrix(rows, cols, F) { var result: Matrix(rows, cols, F) = .{ .mat = undefined }; inline for (&result.mat, self.mat) |*r_row, s_row| { inline for (r_row, s_row) |*r, s| { r.* = @floatFromInt(s); } } return result; } pub fn round(self: Self, I: type) Matrix(rows, cols, I) { var result: Matrix(rows, cols, I) = .{ .mat = undefined }; inline for (&result.mat, self.mat) |*r_row, s_row| { inline for (r_row, s_row) |*r, s| { r.* = @round(s); } } return result; } pub fn floor(self: Self, I: type) Matrix(rows, cols, I) { var result: Matrix(rows, cols, I) = .{ .mat = undefined }; inline for (&result.mat, self.mat) |*r_row, s_row| { inline for (r_row, s_row) |*r, s| { r.* = @floor(s); } } return result; } pub fn ceil(self: Self, I: type) Matrix(rows, cols, I) { var result: Matrix(rows, cols, I) = .{ .mat = undefined }; inline for (&result.mat, self.mat) |*r_row, s_row| { inline for (r_row, s_row) |*r, s| { r.* = @ceil(s); } } return result; } pub fn trunc(self: Self, I: type) Matrix(rows, cols, I) { var result: Matrix(rows, cols, I) = .{ .mat = undefined }; inline for (&result.mat, self.mat) |*r_row, s_row| { inline for (r_row, s_row) |*r, s| { r.* = @trunc(s); } } return result; } }; } pub fn Vector(rows: comptime_int, T: type) type { return Matrix(rows, 1, T); } pub const Swizzle = enum { x, y, z, w, r, g, b, a, fn index(comptime s: Swizzle) comptime_int { return switch (s) { .x, .r => 0, .y, .g => 1, .z, .b => 2, .w, .a => 3, }; } }; pub const StackDirection = enum { rows, cols }; fn stackItemInfo(Item: type, comptime dir: StackDirection) struct { Elem: type, rows: comptime_int, cols: comptime_int, kind: enum { mat, tuple, scalar }, } { switch (@typeInfo(Item)) { .@"struct" => |info| { if (info.is_tuple) switch (dir) { .rows => return .{ .Elem = info.fields[0].type, .rows = 1, .cols = info.fields.len, .kind = .tuple, }, .cols => return .{ .Elem = info.fields[0].type, .rows = info.fields.len, .cols = 1, .kind = .tuple, }, }; const mat_info = @typeInfo(@FieldType(Item, "mat")).array; const mat_row_info = @typeInfo(mat_info.child).array; return .{ .Elem = mat_row_info.child, .rows = mat_info.len, .cols = mat_row_info.len, .kind = .mat, }; }, else => return .{ .Elem = Item, .rows = 1, .cols = 1, .kind = .scalar }, } } pub fn StackOf(Elem: type, T: type, comptime dir: StackDirection) type { var rows = 0; var cols = 0; for (@typeInfo(T).@"struct".fields) |f| { const info = stackItemInfo(f.type, dir); switch (dir) { .rows => { if (cols == 0) cols = info.cols; rows += info.rows; }, .cols => { if (rows == 0) rows = info.rows; cols += info.cols; }, } } return Matrix(rows, cols, Elem); } pub fn stackRows( Elem: type, items: anytype, ) StackOf(Elem, @TypeOf(items), .rows) { var result: StackOf(Elem, @TypeOf(items), .rows) = .{ .mat = undefined }; comptime var offset = 0; inline for (items) |item| { const item_info = stackItemInfo(@TypeOf(item), .rows); if (item_info.cols != result.cols) @compileError(std.fmt.comptimePrint( "Cannot append a {}-column item to matrix with {} columns", .{ item_info.cols, result.cols }, )); switch (item_info.kind) { .mat => @memcpy(result.mat[offset..][0..item_info.rows], &item.mat), .tuple => result.mat[offset] = item, .scalar => result.mat[offset][0] = item, } offset += item_info.rows; } comptime std.debug.assert(offset == result.rows); return result; } pub fn stackCols( Elem: type, items: anytype, ) StackOf(Elem, @TypeOf(items), .cols) { var result: StackOf(Elem, @TypeOf(items), .cols) = .{ .mat = undefined }; comptime var offset = 0; inline for (items) |item| { const item_info = stackItemInfo(@TypeOf(item), .cols); if (item_info.rows != result.rows) @compileError(std.fmt.comptimePrint( "Cannot append a {}-row item to matrix with {} rows", .{ item_info.cols, result.cols }, )); switch (item_info.kind) { .mat => inline for (&result.mat, item.mat) |*row, item_row| { @memcpy(row[offset..][0..item_info.cols], &item_row); }, .tuple => inline for (&result.mat, item) |*row, e| { row[offset] = e; }, .scalar => result.mat[0][offset] = item, } offset += item_info.cols; } comptime std.debug.assert(offset == result.cols); return result; } test "add" { try std.testing.expect(stackRows(u32, .{ 1, 2, 3 }) .add(.splat(1)) .eq(stackRows(u32, .{ 2, 3, 4 }))); } test "swizzle" { const a: Vector(3, u32) = .{ .mat = .{ .{1}, .{2}, .{3} } }; try std.testing.expectEqual(2, a.get(.y)); try std.testing.expectEqual(Vector(2, u32){ .mat = .{ .{3}, .{2} } }, a.swizzle(.{ .z, .y })); try std.testing.expectEqual(Vector(4, u32).splat(1), a.swizzle(.{ .r, .r, .r, .r })); } test "mul" { const m: Matrix(3, 3, u32) = .ident; const a: Vector(3, u32) = .{ .mat = .{ .{1}, .{2}, .{3} } }; try std.testing.expectEqual(a, m.mul(a)); } test "dot" { try std.testing.expectEqual(12, stackRows(u32, .{ 1, 2, 3 }).dot(.splat(2))); } test "stack" { const m1 = stackRows(f32, .{ .{ 1, 2, 3 }, .{ 4, 5, 6 }, .{ 7, 8, 9 }, }); const m2 = stackRows(f32, .{ m1, Matrix(2, 3, f32).zero, }); const m3 = stackCols(f32, .{ m2, Matrix(5, 2, f32).zero, .{ 1, 2, 3, 4, 5 }, }); const expected: [5][6]f32 = .{ .{ 1, 2, 3, 0, 0, 1 }, .{ 4, 5, 6, 0, 0, 2 }, .{ 7, 8, 9, 0, 0, 3 }, .{ 0, 0, 0, 0, 0, 4 }, .{ 0, 0, 0, 0, 0, 5 }, }; try std.testing.expectEqual(expected, m3.mat); } test "stack scalar rows" { const m1 = stackRows(f32, .{ 1, 2, 3 }); const m2 = stackRows(f32, .{ m1, 4, 5, 6 }); const expected: [6][1]f32 = .{ .{1}, .{2}, .{3}, .{4}, .{5}, .{6} }; try std.testing.expectEqual(expected, m2.mat); } test "stack scalar cols" { const m1 = stackCols(f32, .{ 1, 2, 3 }); const m2 = stackCols(f32, .{ m1, 4, 5, 6 }); const expected: [1][6]f32 = .{.{ 1, 2, 3, 4, 5, 6 }}; try std.testing.expectEqual(expected, m2.mat); } test "norm" { const int = stackRows(u32, .{ 3, 4 }); try std.testing.expectEqual(25, int.normSquared()); const float = stackRows(f64, .{ 3, 4 }); try std.testing.expectEqual(5, float.norm()); } test "cast" { try std.testing.expectEqual(@as(f64, 3), stackRows(u32, .{3}).floatFromInt(f64).get(.x)); try std.testing.expectEqual(@as(i32, -4), stackRows(f64, .{-3.5}).round(i32).get(.x)); try std.testing.expectEqual(@as(i32, 4), stackRows(f64, .{3.5}).ceil(i32).get(.x)); try std.testing.expectEqual(@as(i32, 3), stackRows(f64, .{3.5}).floor(i32).get(.x)); }