568 lines
18 KiB
Zig
568 lines
18 KiB
Zig
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));
|
|
}
|