crucible v0.1: CPU inference library for PyTorch .pt state dicts, in pure Zig
- restricted pickle VM (whitelisted globals — cannot execute arbitrary Python, unlike torch.load) - zip container reader (store + raw-deflate) - Tensor views: dtype/offset/sizes/strides, f32 materialization - layers: conv2d f32x8 FMA, linear, relu, maxpool2, adaptiveAvgPool2d, softmax, padInput - examples/stripsolver: real .pt forward, 3777 @ 1.0 👾 Generated with [Letta Code](https://letta.com) Co-Authored-By: Letta Code <noreply@letta.com>
This commit is contained in:
commit
d401f6b07e
11 files changed
+1228
No files matched your search
@@ -0,0 +1,134 @@
|
||||
//! 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 <model.pt> <arm_csv>
|
||||
//! 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);
|
||||
}
|
||||
Reference in new issue
Block a user