diff --git a/build.zig b/build.zig index 94831d6..4449b31 100644 --- a/build.zig +++ b/build.zig @@ -5,7 +5,7 @@ pub fn build(b: *std.Build) void { const optimize = b.standardOptimizeOption(.{}); // The library module — what downstream consumers import. - const crucible_mod = b.addModule("crucible", .{ + const cruc_mod = b.addModule("crucible", .{ .root_source_file = b.path("src/crucible.zig"), .target = target, .optimize = optimize, @@ -14,19 +14,18 @@ pub fn build(b: *std.Build) void { // Static lib artifact (optional install). const lib = b.addLibrary(.{ .name = "crucible", - .root_module = crucible_mod, + .root_module = cruc_mod, }); b.installArtifact(lib); - // Example executable: stripsolver (needs a tests/ png decoder? no — - // example is self-contained: it uses crucible only for state dict + - // layers; arm input comes from CSV). + // Example executable: stripsolver (self-contained: uses crucible only + // for state dict + layers; arm input comes from CSV). const example_mod = b.createModule(.{ .root_source_file = b.path("examples/stripsolver.zig"), .target = target, .optimize = optimize, .imports = &.{ - .{ .name = "crucible", .module = crucible_mod }, + .{ .name = "crucible", .module = cruc_mod }, }, }); const example = b.addExecutable(.{ @@ -41,8 +40,20 @@ pub fn build(b: *std.Build) void { const run_step = b.step("run", "Run the stripsolver example"); run_step.dependOn(&run_example.step); + // linear selftest + const tl_mod = b.createModule(.{ + .root_source_file = b.path("examples/test_linear.zig"), + .target = target, + .optimize = optimize, + .imports = &.{ + .{ .name = "crucible", .module = cruc_mod }, + }, + }); + const tl = b.addExecutable(.{ .name = "test_linear", .root_module = tl_mod }); + b.installArtifact(tl); + // Unit tests. - const tests = b.addTest(.{ .root_module = crucible_mod }); + const tests = b.addTest(.{ .root_module = cruc_mod }); const run_tests = b.addRunArtifact(tests); const test_step = b.step("test", "Run unit tests"); test_step.dependOn(&run_tests.step); diff --git a/examples/test_linear.zig b/examples/test_linear.zig new file mode 100644 index 0000000..b8086a2 --- /dev/null +++ b/examples/test_linear.zig @@ -0,0 +1,29 @@ +//! Minimal known-answer test for layers.linear. +const std = @import("std"); +const crucible = @import("crucible"); + +pub fn main(init: std.process.Init) !void { + const alloc = init.gpa; + + // in_n = 8, out_n = 1, bias 0. w = 1..8, x = 1..8 -> expect 204. + var w: [8]f32 = .{ 1, 2, 3, 4, 5, 6, 7, 8 }; + var x: [8]f32 = .{ 1, 2, 3, 4, 5, 6, 7, 8 }; + var b: [1]f32 = .{0}; + var out: [1]f32 = undefined; + crucible.layers.linear(&x, &w, &b, &out); + std.debug.print("case1 (expect 204): {e}\n", .{out[0]}); + + // bias 100 + b[0] = 100; + crucible.layers.linear(&x, &w, &b, &out); + std.debug.print("case2 (expect 304): {e}\n", .{out[0]}); + + // in_n = 4 (the pose_fc case!) w=1..4 x=1..4 -> 30 + var w4: [4]f32 = .{ 1, 2, 3, 4 }; + var x4: [4]f32 = .{ 1, 2, 3, 4 }; + var b1: [1]f32 = .{0}; + var out1: [1]f32 = undefined; + crucible.layers.linear(&x4, &w4, &b1, &out1); + std.debug.print("case3 (expect 30): {e}\n", .{out1[0]}); + _ = alloc; +} diff --git a/src/layers.zig b/src/layers.zig index ac3ec7d..d9ef4d3 100644 --- a/src/layers.zig +++ b/src/layers.zig @@ -79,7 +79,8 @@ pub fn conv2d( pub fn linear(input: []const f32, weight: []const f32, bias: ?[]const f32, output: []f32) void { const in_n = weight.len / output.len; for (0..output.len) |o| { - var acc: V = @splat(if (bias) |b| b[o] else 0); + var acc: V = @splat(0); + var tail: f32 = 0; const w_row = weight[o * in_n ..][0..in_n]; var x: usize = 0; while (x + 8 <= in_n) : (x += 8) { @@ -87,7 +88,11 @@ pub fn linear(input: []const f32, weight: []const f32, bias: ?[]const f32, outpu const wv: V = w_row[x..][0..8].*; acc = @mulAdd(V, iv, wv, acc); } - output[o] = @reduce(.Add, acc); + while (x < in_n) : (x += 1) { + tail += input[x] * w_row[x]; + } + // bias added ONCE — splatting it into acc lanes would multiply it by 8 + output[o] = @reduce(.Add, acc) + tail + (if (bias) |b| b[o] else 0); } }