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 }