lib/choir/src/profiling/versus/core/oracle.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2
3 const workload_mod = @import("workload.zig");
4
5 const Allocator = std.mem.Allocator;
6 const Workload = workload_mod.Workload;
7
8 pub const float32_tolerance = 1e-4;
9 pub const float64_tolerance = 1e-9;
10
11 pub const Buffers = union(workload_mod.Kind) {
12 saxpy: struct { a: []f32, x: []f32, y: []f32, out: []f32 },
13 dot: struct { x: []f64, y: []f64, out: []f64 },
14 sum: struct { x: []i64, out: []i64 },
15 matmul: struct { a: []f64, b: []f64, c: []f64 },
16 polybench_gemm: struct { a: []f64, b: []f64, c_initial: []f64, c: []f64 },
17 stencil3: struct { in: []f64, out: []f64 },
18 clampsum: struct { x: []i64, out: []i64 },
19
20 pub fn alloc(allocator: Allocator, workload: Workload) !Buffers {
21 const n: usize = @intCast(workload.n);
22 switch (workload.kind) {
23 .saxpy => {
24 const a = try allocator.alloc(f32, 1);
25 const x = try allocator.alloc(f32, n);
26 const y = try allocator.alloc(f32, n);
27 const out = try allocator.alloc(f32, n);
28 workload_mod.fillUnit(f32, a, 0);
29 workload_mod.fillUnit(f32, x, 1);
30 workload_mod.fillUnit(f32, y, 1 + workload.n);
31 @memset(out, 0);
32 return .{ .saxpy = .{ .a = a, .x = x, .y = y, .out = out } };
33 },
34 .dot => {
35 const x = try allocator.alloc(f64, n);
36 const y = try allocator.alloc(f64, n);
37 const out = try allocator.alloc(f64, 1);
38 workload_mod.fillUnit(f64, x, 0);
39 workload_mod.fillUnit(f64, y, workload.n);
40 @memset(out, 0);
41 return .{ .dot = .{ .x = x, .y = y, .out = out } };
42 },
43 .sum => {
44 const x = try allocator.alloc(i64, n);
45 const out = try allocator.alloc(i64, 1);
46 workload_mod.fillInt(x, 0);
47 @memset(out, 0);
48 return .{ .sum = .{ .x = x, .out = out } };
49 },
50 .matmul => {
51 const cells = n * n;
52 const a = try allocator.alloc(f64, cells);
53 const b = try allocator.alloc(f64, cells);
54 const c = try allocator.alloc(f64, cells);
55 workload_mod.fillUnit(f64, a, 0);
56 workload_mod.fillUnit(f64, b, workload.n * workload.n);
57 @memset(c, 0);
58 return .{ .matmul = .{ .a = a, .b = b, .c = c } };
59 },
60 .polybench_gemm => {
61 const cells = n * n;
62 const a = try allocator.alloc(f64, cells);
63 const b = try allocator.alloc(f64, cells);
64 const c_initial = try allocator.alloc(f64, cells);
65 const c = try allocator.alloc(f64, cells);
66 workload_mod.fillUnit(f64, a, 0);
67 workload_mod.fillUnit(f64, b, workload.n * workload.n);
68 workload_mod.fillUnit(f64, c_initial, 2 * workload.n * workload.n);
69 @memcpy(c, c_initial);
70 return .{ .polybench_gemm = .{ .a = a, .b = b, .c_initial = c_initial, .c = c } };
71 },
72 .stencil3 => {
73 const in = try allocator.alloc(f64, n);
74 const out = try allocator.alloc(f64, n);
75 workload_mod.fillUnit(f64, in, 0);
76 @memset(out, 0);
77 return .{ .stencil3 = .{ .in = in, .out = out } };
78 },
79 .clampsum => {
80 const x = try allocator.alloc(i64, n);
81 const out = try allocator.alloc(i64, 1);
82 workload_mod.fillInt(x, 0);
83 @memset(out, 0);
84 return .{ .clampsum = .{ .x = x, .out = out } };
85 },
86 }
87 }
88
89 pub fn reset(self: *Buffers) void {
90 switch (self.*) {
91 .saxpy => |buffers| @memset(buffers.out, 0),
92 .dot => |buffers| @memset(buffers.out, 0),
93 .sum => |buffers| @memset(buffers.out, 0),
94 .matmul => |buffers| @memset(buffers.c, 0),
95 .polybench_gemm => |buffers| @memcpy(buffers.c, buffers.c_initial),
96 .stencil3 => |buffers| @memset(buffers.out, 0),
97 .clampsum => |buffers| @memset(buffers.out, 0),
98 }
99 }
100
101 pub fn deinit(self: *Buffers, allocator: Allocator) void {
102 switch (self.*) {
103 .saxpy => |buffers| {
104 allocator.free(buffers.a);
105 allocator.free(buffers.x);
106 allocator.free(buffers.y);
107 allocator.free(buffers.out);
108 },
109 .dot => |buffers| {
110 allocator.free(buffers.x);
111 allocator.free(buffers.y);
112 allocator.free(buffers.out);
113 },
114 .sum => |buffers| {
115 allocator.free(buffers.x);
116 allocator.free(buffers.out);
117 },
118 .matmul => |buffers| {
119 allocator.free(buffers.a);
120 allocator.free(buffers.b);
121 allocator.free(buffers.c);
122 },
123 .polybench_gemm => |buffers| {
124 allocator.free(buffers.a);
125 allocator.free(buffers.b);
126 allocator.free(buffers.c_initial);
127 allocator.free(buffers.c);
128 },
129 .stencil3 => |buffers| {
130 allocator.free(buffers.in);
131 allocator.free(buffers.out);
132 },
133 .clampsum => |buffers| {
134 allocator.free(buffers.x);
135 allocator.free(buffers.out);
136 },
137 }
138 }
139 };
140
141 fn expectClose(expected: f64, actual: f64, tolerance: f64) !void {
142 if (!std.math.isFinite(expected) or !std.math.isFinite(actual)) return error.OracleMismatch;
143 const magnitude = @max(@abs(expected), 1.0);
144 if (@abs(expected - actual) > tolerance * magnitude) return error.OracleMismatch;
145 }
146
147 pub fn verify(workload: Workload, buffers: *const Buffers) !void {
148 switch (buffers.*) {
149 .saxpy => |b| {
150 for (b.x, b.y, b.out) |x, y, actual| {
151 const expected = @as(f64, b.a[0]) * @as(f64, x) + @as(f64, y);
152 try expectClose(expected, @as(f64, actual), float32_tolerance);
153 }
154 },
155 .dot => |b| {
156 var acc: f64 = 0;
157 for (b.x, b.y) |x, y| acc += x * y;
158 try expectClose(acc, b.out[0], float64_tolerance);
159 },
160 .sum => |b| {
161 var acc: i64 = 0;
162 for (b.x) |x| acc +%= x;
163 if (acc != b.out[0]) return error.OracleMismatch;
164 },
165 .matmul => |b| {
166 const n: usize = @intCast(workload.n);
167 for (b.c, 0..) |actual, cell| {
168 const row = cell / n;
169 const col = cell % n;
170 var acc: f64 = 0;
171 var k: usize = 0;
172 while (k < n) : (k += 1) {
173 acc += b.a[row * n + k] * b.b[k * n + col];
174 }
175 try expectClose(acc, actual, float64_tolerance);
176 }
177 },
178 .polybench_gemm => |b| {
179 const n: usize = @intCast(workload.n);
180 for (b.c, 0..) |actual, cell| {
181 const row = cell / n;
182 const col = cell % n;
183 var acc = workload_mod.gemm_beta * b.c_initial[cell];
184 var k: usize = 0;
185 while (k < n) : (k += 1) {
186 acc += workload_mod.gemm_alpha * b.a[row * n + k] * b.b[k * n + col];
187 }
188 try expectClose(acc, actual, float64_tolerance);
189 }
190 },
191 .stencil3 => |b| {
192 for (b.out, 0..) |actual, cell| {
193 const expected = if (cell == 0 or cell + 1 == b.out.len)
194 0
195 else
196 0.25 * b.in[cell - 1] + 0.5 * b.in[cell] + 0.25 * b.in[cell + 1];
197 try expectClose(expected, actual, float64_tolerance);
198 }
199 },
200 .clampsum => |b| {
201 var acc: i64 = 0;
202 for (b.x) |x| {
203 acc +%= std.math.clamp(x, workload_mod.clamp_lo, workload_mod.clamp_hi);
204 }
205 if (acc != b.out[0]) return error.OracleMismatch;
206 },
207 }
208 }
209
210 fn expectEveryCorruptionRejected(comptime T: type, workload: Workload, buffers: *const Buffers, outputs: []T) !void {
211 for (outputs) |*actual| {
212 actual.* += 1;
213 try std.testing.expectError(error.OracleMismatch, verify(workload, buffers));
214 actual.* -= 1;
215 }
216 }
217
218 test "verify accepts reference outputs and rejects corrupted ones" {
219 const allocator = std.testing.allocator;
220 const workload = workload_mod.Workload{ .name = "sum_test", .kind = .sum, .n = 128 };
221 var buffers = try Buffers.alloc(allocator, workload);
222 defer buffers.deinit(allocator);
223
224 var acc: i64 = 0;
225 for (buffers.sum.x) |x| acc +%= x;
226 buffers.sum.out[0] = acc;
227 try verify(workload, &buffers);
228
229 buffers.sum.out[0] += 1;
230 try std.testing.expectError(error.OracleMismatch, verify(workload, &buffers));
231 }
232
233 test "verify rejects corruption at every stencil cell" {
234 const allocator = std.testing.allocator;
235 const workload = workload_mod.Workload{ .name = "stencil_test", .kind = .stencil3, .n = 9 };
236 var buffers = try Buffers.alloc(allocator, workload);
237 defer buffers.deinit(allocator);
238
239 const b = buffers.stencil3;
240 var i: usize = 1;
241 while (i < b.in.len - 1) : (i += 1) {
242 b.out[i] = 0.25 * b.in[i - 1] + 0.5 * b.in[i] + 0.25 * b.in[i + 1];
243 }
244 try verify(workload, &buffers);
245
246 try expectEveryCorruptionRejected(f64, workload, &buffers, b.out);
247 }
248
249 test "verify rejects corruption at every saxpy cell" {
250 const allocator = std.testing.allocator;
251 const workload = workload_mod.Workload{ .name = "saxpy_test", .kind = .saxpy, .n = 9 };
252 var buffers = try Buffers.alloc(allocator, workload);
253 defer buffers.deinit(allocator);
254
255 const b = buffers.saxpy;
256 for (b.x, b.y, b.out) |x, y, *out| out.* = b.a[0] * x + y;
257 try verify(workload, &buffers);
258
259 try expectEveryCorruptionRejected(f32, workload, &buffers, b.out);
260 }
261
262 test "verify rejects corruption at every matrix cell" {
263 const allocator = std.testing.allocator;
264 const workload = workload_mod.Workload{ .name = "matmul_test", .kind = .matmul, .n = 4 };
265 var buffers = try Buffers.alloc(allocator, workload);
266 defer buffers.deinit(allocator);
267
268 const b = buffers.matmul;
269 const n: usize = @intCast(workload.n);
270 for (b.c, 0..) |*out, cell| {
271 const row = cell / n;
272 const col = cell % n;
273 for (0..n) |k| out.* += b.a[row * n + k] * b.b[k * n + col];
274 }
275 try verify(workload, &buffers);
276
277 try expectEveryCorruptionRejected(f64, workload, &buffers, b.c);
278 }
279
280 test "verify rejects corruption at every PolyBench GEMM cell" {
281 const allocator = std.testing.allocator;
282 const workload = workload_mod.Workload{ .name = "gemm_test", .kind = .polybench_gemm, .n = 4 };
283 var buffers = try Buffers.alloc(allocator, workload);
284 defer buffers.deinit(allocator);
285
286 const b = buffers.polybench_gemm;
287 const n: usize = @intCast(workload.n);
288 for (0..n) |row| {
289 for (0..n) |col| b.c[row * n + col] *= workload_mod.gemm_beta;
290 for (0..n) |k| {
291 for (0..n) |col| {
292 b.c[row * n + col] += workload_mod.gemm_alpha * b.a[row * n + k] * b.b[k * n + col];
293 }
294 }
295 }
296 try verify(workload, &buffers);
297
298 try expectEveryCorruptionRejected(f64, workload, &buffers, b.c);
299 }
300
301 test "verify rejects non-finite floating results" {
302 const allocator = std.testing.allocator;
303 const workload = workload_mod.Workload{ .name = "dot_test", .kind = .dot, .n = 128 };
304 var buffers = try Buffers.alloc(allocator, workload);
305 defer buffers.deinit(allocator);
306
307 for ([_]f64{ std.math.nan(f64), std.math.inf(f64), -std.math.inf(f64) }) |non_finite| {
308 buffers.dot.out[0] = non_finite;
309 try std.testing.expectError(error.OracleMismatch, verify(workload, &buffers));
310 }
311 }