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 }