//! Example: load strip-v03 .pt through crucible, run a forward pass on a //! prepared arm, print the argmax digits. Validates the library end-to-end //! against the original Python sidecar (expect 3777 @ ~1.0). //! //! Usage: stripsolver //! arm_csv: 4400 comma-separated uint8 values (the 44x100 preprocessed input). const std = @import("std"); const crucible = @import("crucible"); pub fn main(init: std.process.Init) !void { const alloc = init.gpa; const io = init.io; var args_iter = std.process.Args.Iterator.init(init.minimal.args); defer args_iter.deinit(); _ = args_iter.skip(); const pt_path = args_iter.next() orelse return error.MissingPtPath; const arm_path = args_iter.next() orelse return error.MissingArmPath; // 1. load the state dict through crucible var sd = try crucible.loadStateDict(alloc, io, pt_path); defer sd.deinit(); std.debug.print("state dict entries: {d}\n", .{sd.entries.len}); for (sd.entries) |entry| std.debug.print(" {s}\n", .{entry.name}); const w0 = sd.get("trunk.0.weight") orelse return error.MissingTrunk0; std.debug.print("trunk.0.weight: {d} elems f32\n", .{w0.numel()}); // 2. read arm csv const csv = try std.Io.Dir.cwd().readFileAlloc(io, arm_path, alloc, .limited(1024 * 1024)); defer alloc.free(csv); var arm: [44 * 100]f32 = undefined; var fields = std.mem.splitScalar(u8, std.mem.trim(u8, csv, " \n\r\t"), ','); for (0..44 * 100) |i| { const f = fields.next() orelse return error.BadCsv; arm[i] = @as(f32, @floatFromInt(try std.fmt.parseInt(u8, f, 10))) / 255.0; } // 3. forward through layers var pose: [4]f32 = @splat(0); if (fields.next()) |f| pose[0] = @as(f32, @floatFromInt(try std.fmt.parseInt(u8, f, 10))) / 255.0; if (fields.next()) |f| pose[1] = @as(f32, @floatFromInt(try std.fmt.parseInt(u8, f, 10))) / 255.0; if (fields.next()) |f| pose[2] = @as(f32, @floatFromInt(try std.fmt.parseInt(u8, f, 10))) / 255.0; if (fields.next()) |f| pose[3] = @as(f32, @floatFromInt(try std.fmt.parseInt(u8, f, 10))) / 255.0; const layers = crucible.layers; var a1: [32 * 44 * 100]f32 = undefined; try layers.conv2d(alloc, &arm, 44, 100, 1, try (sd.get("trunk.0.weight") orelse return error.Missing).toF32(alloc), try tF32(alloc, &sd, "trunk.0.bias"), &a1, 32); layers.relu(&a1); var p1: [32 * 22 * 50]f32 = undefined; layers.maxpool2(&a1, 44, 100, 32, &p1); var a2: [64 * 22 * 50]f32 = undefined; try layers.conv2d(alloc, &p1, 22, 50, 32, try (sd.get("trunk.3.weight") orelse return error.Missing).toF32(alloc), try tF32(alloc, &sd, "trunk.3.bias"), &a2, 64); layers.relu(&a2); var p2: [64 * 11 * 25]f32 = undefined; layers.maxpool2(&a2, 22, 50, 64, &p2); var a3: [64 * 11 * 25]f32 = undefined; try layers.conv2d(alloc, &p2, 11, 25, 64, try (sd.get("trunk.6.weight") orelse return error.Missing).toF32(alloc), try tF32(alloc, &sd, "trunk.6.bias"), &a3, 64); layers.relu(&a3); var p3: [64 * 5 * 12]f32 = undefined; layers.maxpool2(&a3, 11, 25, 64, &p3); // 4. pose + per-strip heads var pose_p: [32]f32 = undefined; layers.linear(&pose, try tF32(alloc, &sd, "pose_fc.weight"), try tF32(alloc, &sd, "pose_fc.bias"), &pose_p); layers.relu(&pose_p); const sw = 12 / 4; const row_bins = [2][2]usize{ .{ 0, 3 }, .{ 2, 5 } }; const col_bins = [2][2]usize{ .{ 0, 2 }, .{ 1, 3 } }; var solution: [4]u8 = undefined; for (0..4) |h| { var pooled: [256]f32 = undefined; for (0..2) |oy| { for (0..2) |ox| { const rows = row_bins[oy]; const cols = col_bins[ox]; const n: f32 = @floatFromInt((rows[1] - rows[0]) * (cols[1] - cols[0])); const oi = oy * 2 + ox; for (0..64) |c| { var acc: f32 = 0; var y = rows[0]; while (y < rows[1]) : (y += 1) { var x = cols[0]; while (x < cols[1]) : (x += 1) { acc += p3[(c * 5 + y) * 12 + h * sw + x]; } } pooled[c * 4 + oi] = acc / n; } } } var fc_in: [288]f32 = undefined; @memcpy(fc_in[0..256], pooled[0..256]); @memcpy(fc_in[256..288], pose_p[0..32]); const fc_w_name = try std.fmt.allocPrint(alloc, "strip_fcs.{d}.0.weight", .{h}); const fc_b_name = try std.fmt.allocPrint(alloc, "strip_fcs.{d}.0.bias", .{h}); const h_w_name = try std.fmt.allocPrint(alloc, "heads.{d}.weight", .{h}); const h_b_name = try std.fmt.allocPrint(alloc, "heads.{d}.bias", .{h}); defer alloc.free(fc_w_name); defer alloc.free(fc_b_name); defer alloc.free(h_w_name); defer alloc.free(h_b_name); var fc: [64]f32 = undefined; layers.linear(&fc_in, try tF32(alloc, &sd, fc_w_name), try tF32(alloc, &sd, fc_b_name), &fc); layers.relu(&fc); var logits: [10]f32 = undefined; layers.linear(&fc, try tF32(alloc, &sd, h_w_name), try tF32(alloc, &sd, h_b_name), &logits); layers.softmax(&logits); var best: usize = 0; for (1..10) |d| { if (logits[d] > logits[best]) best = d; } solution[h] = @intCast(best); std.debug.print("head {d}: digit {d} conf {d:.4}\n", .{ h, best, logits[best] }); } std.debug.print("solution: {any}\n", .{solution}); } fn tF32(alloc: std.mem.Allocator, sd: *const crucible.StateDict, name: []const u8) ![]f32 { const t = sd.get(name) orelse { std.debug.print("MissingTensor: {s}\n", .{name}); return error.MissingTensor; }; return t.toF32(alloc); }