fix(linear): add bias once, not once per SIMD lane
acc was seeded with @splat(bias) and reduced with .Add, so every linear output carried 7 extra copies of its bias. conv2d was unaffected (bias via memset), which is why the trunk matched zig-solver to 5e-6 while fc/head activations diverged — the 81% labels regression. Also adds a scalar tail for in_n < 8 (pose_fc has in_n=4) and a known-answer selftest (test_linear). 👾 Generated with [Letta Code](https://letta.com) Co-Authored-By: Letta Code <noreply@letta.com>
This commit is contained in:
1 parent
d401f6b07e
commit
a2ff07234c
3 files changed
+54
-9
No files matched your search
@@ -5,7 +5,7 @@ pub fn build(b: *std.Build) void {
|
|||||||
const optimize = b.standardOptimizeOption(.{});
|
const optimize = b.standardOptimizeOption(.{});
|
||||||
|
|
||||||
// The library module — what downstream consumers import.
|
// 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"),
|
.root_source_file = b.path("src/crucible.zig"),
|
||||||
.target = target,
|
.target = target,
|
||||||
.optimize = optimize,
|
.optimize = optimize,
|
||||||
@@ -14,19 +14,18 @@ pub fn build(b: *std.Build) void {
|
|||||||
// Static lib artifact (optional install).
|
// Static lib artifact (optional install).
|
||||||
const lib = b.addLibrary(.{
|
const lib = b.addLibrary(.{
|
||||||
.name = "crucible",
|
.name = "crucible",
|
||||||
.root_module = crucible_mod,
|
.root_module = cruc_mod,
|
||||||
});
|
});
|
||||||
b.installArtifact(lib);
|
b.installArtifact(lib);
|
||||||
|
|
||||||
// Example executable: stripsolver (needs a tests/ png decoder? no —
|
// Example executable: stripsolver (self-contained: uses crucible only
|
||||||
// example is self-contained: it uses crucible only for state dict +
|
// for state dict + layers; arm input comes from CSV).
|
||||||
// layers; arm input comes from CSV).
|
|
||||||
const example_mod = b.createModule(.{
|
const example_mod = b.createModule(.{
|
||||||
.root_source_file = b.path("examples/stripsolver.zig"),
|
.root_source_file = b.path("examples/stripsolver.zig"),
|
||||||
.target = target,
|
.target = target,
|
||||||
.optimize = optimize,
|
.optimize = optimize,
|
||||||
.imports = &.{
|
.imports = &.{
|
||||||
.{ .name = "crucible", .module = crucible_mod },
|
.{ .name = "crucible", .module = cruc_mod },
|
||||||
},
|
},
|
||||||
});
|
});
|
||||||
const example = b.addExecutable(.{
|
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");
|
const run_step = b.step("run", "Run the stripsolver example");
|
||||||
run_step.dependOn(&run_example.step);
|
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.
|
// 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 run_tests = b.addRunArtifact(tests);
|
||||||
const test_step = b.step("test", "Run unit tests");
|
const test_step = b.step("test", "Run unit tests");
|
||||||
test_step.dependOn(&run_tests.step);
|
test_step.dependOn(&run_tests.step);
|
||||||
|
|||||||
@@ -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;
|
||||||
|
}
|
||||||
+7
-2
@@ -79,7 +79,8 @@ pub fn conv2d(
|
|||||||
pub fn linear(input: []const f32, weight: []const f32, bias: ?[]const f32, output: []f32) void {
|
pub fn linear(input: []const f32, weight: []const f32, bias: ?[]const f32, output: []f32) void {
|
||||||
const in_n = weight.len / output.len;
|
const in_n = weight.len / output.len;
|
||||||
for (0..output.len) |o| {
|
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];
|
const w_row = weight[o * in_n ..][0..in_n];
|
||||||
var x: usize = 0;
|
var x: usize = 0;
|
||||||
while (x + 8 <= in_n) : (x += 8) {
|
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].*;
|
const wv: V = w_row[x..][0..8].*;
|
||||||
acc = @mulAdd(V, iv, wv, acc);
|
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);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
Reference in new issue
Block a user