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 }