lib/accy/src/tensor/session/test.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const tensor = @import("../root.zig");
3 const fixture = @import("../../fixture/root.zig");
4 const session_mod = @import("root.zig");
5
6 const Builder = tensor.Builder;
7 const Descriptor = session_mod.Descriptor;
8 const Program = tensor.program.Program;
9 const Session = session_mod.Session;
10 const Status = session_mod.Status;
11
12 const ProgramFixture = *const fn (std.mem.Allocator) anyerror!Program;
13 const max_caller_buffers = 16;
14
15 test "tensor session matches cpu execution bit for bit on caller memory" {
16 try fixture.requireNativeCpuArtifacts();
17 const allocator = std.testing.allocator;
18 var session = try Session.init(allocator);
19 defer session.deinit();
20
21 const programs = [_]ProgramFixture{ blendProgram, aliasProgram, scanProgram };
22 for (programs) |build| {
23 var source = try build(allocator);
24 defer source.deinit();
25 try expectSessionMatchesCpu(allocator, &session, &source);
26 }
27 }
28
29 test "tensor session reports stable status codes for rejected calls" {
30 try fixture.requireNativeCpuArtifacts();
31 const allocator = std.testing.allocator;
32 var session = try Session.init(allocator);
33 defer session.deinit();
34
35 var source = try blendProgram(allocator);
36 defer source.deinit();
37 const bytes = try tensor.wire.encode(allocator, &source);
38 defer allocator.free(bytes);
39 try expectStatus(.invalid_program, session.compile(bytes[0 .. bytes.len - 1]));
40 const bumped = try allocator.dupe(u8, bytes);
41 defer allocator.free(bumped);
42 std.mem.writeInt(u32, bumped[4..8], tensor.wire.version + 1, .little);
43 try expectStatus(.unsupported_version, session.compile(bumped));
44
45 const handle = try session.compile(bytes);
46 var memory: CallerMemory = .{};
47 defer memory.deinit(allocator);
48 const inputs = try memory.buffers(allocator, try session.inputDescriptors(handle));
49 const outputs = try memory.buffers(allocator, try session.outputDescriptors(handle));
50 try expectStatus(.invalid_buffer, session.invoke(handle, allocator, inputs[0..1], outputs));
51 const overlapping = [_][]u8{ inputs[0], outputs[1] };
52 try expectStatus(.invalid_buffer, session.invoke(handle, allocator, inputs, &overlapping));
53 try session.invoke(handle, allocator, inputs, outputs);
54
55 try session.release(handle);
56 try expectStatus(.invalid_handle, session.invoke(handle, allocator, inputs, outputs));
57 try expectStatus(.invalid_handle, session.release(handle));
58 try expectStatus(.invalid_handle, session.inputDescriptors(handle));
59 const reused = try session.compile(bytes);
60 try std.testing.expectEqual(handle.index, reused.index);
61 try std.testing.expect(handle.generation != reused.generation);
62
63 const pinned = [_]Status{
64 .ok,
65 .invalid_program,
66 .unsupported_version,
67 .invalid_handle,
68 .invalid_buffer,
69 .too_many_programs,
70 .out_of_memory,
71 .compile_failed,
72 .launch_failed,
73 };
74 for (pinned, 0..) |status, value| {
75 try std.testing.expectEqual(@as(u32, @intCast(value)), @backingInt(status));
76 }
77 }
78
79 fn expectSessionMatchesCpu(
80 allocator: std.mem.Allocator,
81 session: *Session,
82 source: *const Program,
83 ) !void {
84 const bytes = try tensor.wire.encode(allocator, source);
85 defer allocator.free(bytes);
86 const handle = try session.compile(bytes);
87 defer session.release(handle) catch unreachable;
88 const input_descriptors = try session.inputDescriptors(handle);
89 const output_descriptors = try session.outputDescriptors(handle);
90 try std.testing.expectEqual(source.parameters.len, input_descriptors.len);
91 try std.testing.expectEqual(source.outputs.len, output_descriptors.len);
92
93 var memory: CallerMemory = .{};
94 defer memory.deinit(allocator);
95 const inputs = try memory.buffers(allocator, input_descriptors);
96 for (inputs, 0..) |input, index| fillFloats(input, index);
97 const snapshots = try memory.buffers(allocator, input_descriptors);
98 for (snapshots, inputs) |snapshot, input| @memcpy(snapshot, input);
99 const outputs = try memory.buffers(allocator, output_descriptors);
100 for (outputs) |output| @memset(output, 0xa5);
101 const expected = try memory.buffers(allocator, output_descriptors);
102
103 try session.invoke(handle, allocator, inputs, outputs);
104 try tensor.runCpu(allocator, source, inputs, expected);
105 for (expected, outputs, output_descriptors) |want, got, descriptor| {
106 try std.testing.expectEqual(descriptor.byte_size, got.len);
107 try std.testing.expectEqualSlices(u8, want, got);
108 }
109 for (snapshots, inputs) |snapshot, input| {
110 try std.testing.expectEqualSlices(u8, snapshot, input);
111 }
112 }
113
114 fn expectStatus(expected: Status, result: anytype) !void {
115 if (result) |_| {
116 return error.TestExpectedError;
117 } else |err| {
118 try std.testing.expectEqual(expected, Status.fromError(err));
119 }
120 }
121
122 fn fillFloats(bytes: []u8, seed: usize) void {
123 const values = std.mem.bytesAsSlice(f32, bytes);
124 for (values, 0..) |*value, index| {
125 const step: f32 = @floatFromInt(seed * 7 + index);
126 value.* = 0.375 * step - 1.5;
127 }
128 }
129
130 const CallerMemory = struct {
131 slots: [max_caller_buffers][]u8 = undefined,
132 alignments: [max_caller_buffers]std.mem.Alignment = undefined,
133 count: usize = 0,
134
135 fn buffers(
136 self: *CallerMemory,
137 allocator: std.mem.Allocator,
138 descriptors: []const Descriptor,
139 ) ![]const []u8 {
140 const start = self.count;
141 for (descriptors) |descriptor| {
142 if (self.count == max_caller_buffers) return error.OutOfMemory;
143 const alignment = std.mem.Alignment.fromByteUnits(descriptor.alignment);
144 const size = descriptor.byte_size;
145 const memory = allocator.rawAlloc(size, alignment, @returnAddress()) orelse {
146 return error.OutOfMemory;
147 };
148 self.slots[self.count] = memory[0..size];
149 self.alignments[self.count] = alignment;
150 self.count += 1;
151 }
152 return self.slots[start..self.count];
153 }
154
155 fn deinit(self: *CallerMemory, allocator: std.mem.Allocator) void {
156 for (self.slots[0..self.count], self.alignments[0..self.count]) |bytes, alignment| {
157 allocator.rawFree(bytes, alignment, @returnAddress());
158 }
159 }
160 };
161
162 fn blendProgram(allocator: std.mem.Allocator) anyerror!Program {
163 var builder = try Builder.init(allocator, "session_blend");
164 errdefer builder.deinit();
165 const x = try builder.input(.f32, .{ .lane = 8 });
166 const y = try builder.input(.f32, .{ .lane = 8 });
167 const bias = try builder.full(.f32, .{ .lane = 8 }, 0.5);
168 const scaled = try builder.binary(.mul, try builder.unary(.tanh, x), y);
169 const mixed = try builder.binary(.add, scaled, bias);
170 const larger = try builder.binary(.max, x, y);
171 return builder.finish(&.{ mixed, larger });
172 }
173
174 fn aliasProgram(allocator: std.mem.Allocator) anyerror!Program {
175 var builder = try Builder.init(allocator, "session_alias");
176 errdefer builder.deinit();
177 const x = try builder.input(.f32, .{ .lane = 8 });
178 const y = try builder.input(.f32, .{ .lane = 8 });
179 const sum = try builder.binary(.add, x, y);
180 return builder.finish(&.{ x, sum, sum });
181 }
182
183 fn scanProgram(allocator: std.mem.Allocator) anyerror!Program {
184 var builder = try Builder.init(allocator, "session_scan");
185 errdefer builder.deinit();
186 const x = try builder.input(.f32, .{ .lane = 8 });
187 const c = try builder.input(.f32, .{ .lane = 8 });
188 const acc = try builder.full(.f32, .{ .lane = 8 }, 0.0);
189 const walked = try builder.scan(.{
190 .length = 5,
191 .init = .{ .x = x, .acc = acc, .c = c },
192 .body = scanStep,
193 });
194 return builder.finish(&.{ walked.x, walked.acc });
195 }
196
197 fn scanStep(builder: *Builder, carry: anytype) !@TypeOf(carry) {
198 const half = try builder.scalar(.f32, 0.5);
199 const next_x = try (try (try carry.x.mul(carry.x)).mul(half)).add(carry.c);
200 return .{ .x = next_x, .acc = try carry.acc.add(next_x), .c = carry.c };
201 }