lib/choir/src/profiling/versus/core/abi.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2
3 const oracle = @import("oracle.zig");
4 const workload_mod = @import("workload.zig");
5
6 pub const SaxpyFn = *const fn (i64, [*]const f32, [*]const f32, [*]const f32, [*]f32) callconv(.c) void;
7 pub const DotFn = *const fn (i64, [*]const f64, [*]const f64, [*]f64) callconv(.c) void;
8 pub const SumFn = *const fn (i64, [*]const i64, [*]i64) callconv(.c) void;
9 pub const MatmulFn = *const fn (i64, [*]const f64, [*]const f64, [*]f64) callconv(.c) void;
10 pub const PolybenchGemmFn = *const fn (i32, i32, i32, f64, f64, [*]f64, [*]f64, [*]f64) callconv(.c) void;
11 pub const Stencil3Fn = *const fn (i64, [*]const f64, [*]f64) callconv(.c) void;
12 pub const ClampsumFn = *const fn (i64, [*]const i64, i64, i64, [*]i64) callconv(.c) void;
13
14 pub fn FnType(comptime kind: workload_mod.Kind) type {
15 return switch (kind) {
16 .saxpy => SaxpyFn,
17 .dot => DotFn,
18 .sum => SumFn,
19 .matmul => MatmulFn,
20 .polybench_gemm => PolybenchGemmFn,
21 .stencil3 => Stencil3Fn,
22 .clampsum => ClampsumFn,
23 };
24 }
25
26 pub const Kernel = union(workload_mod.Kind) {
27 saxpy: SaxpyFn,
28 dot: DotFn,
29 sum: SumFn,
30 matmul: MatmulFn,
31 polybench_gemm: PolybenchGemmFn,
32 stencil3: Stencil3Fn,
33 clampsum: ClampsumFn,
34
35 pub fn call(self: Kernel, workload: workload_mod.Workload, buffers: *const oracle.Buffers) void {
36 const n: i64 = @intCast(workload.n);
37 switch (self) {
38 .saxpy => |kernel| kernel(n, buffers.saxpy.a.ptr, buffers.saxpy.x.ptr, buffers.saxpy.y.ptr, buffers.saxpy.out.ptr),
39 .dot => |kernel| kernel(n, buffers.dot.x.ptr, buffers.dot.y.ptr, buffers.dot.out.ptr),
40 .sum => |kernel| kernel(n, buffers.sum.x.ptr, buffers.sum.out.ptr),
41 .matmul => |kernel| kernel(n, buffers.matmul.a.ptr, buffers.matmul.b.ptr, buffers.matmul.c.ptr),
42 .polybench_gemm => |kernel| {
43 const dim: i32 = @intCast(workload.n);
44 kernel(dim, dim, dim, workload_mod.gemm_alpha, workload_mod.gemm_beta, buffers.polybench_gemm.c.ptr, buffers.polybench_gemm.a.ptr, buffers.polybench_gemm.b.ptr);
45 },
46 .stencil3 => |kernel| kernel(n, buffers.stencil3.in.ptr, buffers.stencil3.out.ptr),
47 .clampsum => |kernel| kernel(n, buffers.clampsum.x.ptr, workload_mod.clamp_lo, workload_mod.clamp_hi, buffers.clampsum.out.ptr),
48 }
49 }
50 };
51
52 fn referenceSaxpy(n: i64, a: [*]const f32, x: [*]const f32, y: [*]const f32, out: [*]f32) callconv(.c) void {
53 const s = a[0];
54 var i: usize = 0;
55 while (i < @as(usize, @intCast(n))) : (i += 1) {
56 out[i] = s * x[i] + y[i];
57 }
58 }
59
60 test "kernel dispatch calls through the C ABI" {
61 const workload = workload_mod.Workload{ .name = "saxpy_test", .kind = .saxpy, .n = 64 };
62 var buffers = try oracle.Buffers.alloc(std.testing.allocator, workload);
63 defer buffers.deinit(std.testing.allocator);
64
65 const kernel = Kernel{ .saxpy = referenceSaxpy };
66 kernel.call(workload, &buffers);
67 try oracle.verify(workload, &buffers);
68 }