lib/accy/src/kernel/model/program/parameter.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const kernel = @import("../core/root.zig");
3
4 pub fn Standard(comptime parameters: anytype) type {
5 return struct {
6 pub const count = arity(parameters);
7
8 pub fn index(comptime name_value: anytype) usize {
9 return parameterIndex(parameters, name_value);
10 }
11 };
12 }
13
14 pub fn Args(comptime parameters: anytype, comptime Builder: type) type {
15 return struct {
16 builder: *Builder,
17
18 pub fn param(self: @This(), comptime name_value: anytype) Value(parameterSpec(parameters, name_value)) {
19 return argument(self.builder, parameterSpec(parameters, name_value), parameterIndex(parameters, name_value));
20 }
21 };
22 }
23
24 pub fn Value(comptime param: kernel.Param) type {
25 return switch (param) {
26 .scalar => |dtype| kernel.TypedValue(dtype),
27 .buffer => |buffer| kernel.BufferView(buffer.dtype),
28 };
29 }
30
31 pub fn schema(comptime parameters: anytype) *const [arity(parameters)]kernel.Param {
32 const count = comptime arity(parameters);
33 comptime var result: [count]kernel.Param = undefined;
34 if (comptime named(parameters)) {
35 inline for (@typeInfo(@TypeOf(parameters)).@"struct".field_names, 0..) |field_name, index_value| {
36 result[index_value] = @field(parameters, field_name);
37 }
38 } else {
39 inline for (0..count) |index_value| {
40 result[index_value] = parameters[index_value];
41 }
42 }
43 const final = result;
44 return &final;
45 }
46
47 pub fn named(comptime parameters: anytype) bool {
48 return switch (@typeInfo(@TypeOf(parameters))) {
49 .@"struct" => |info| !info.is_tuple,
50 else => false,
51 };
52 }
53
54 pub fn arity(comptime parameters: anytype) usize {
55 return if (comptime named(parameters))
56 @typeInfo(@TypeOf(parameters)).@"struct".field_names.len
57 else
58 parameters.len;
59 }
60
61 fn argument(builder: anytype, comptime param: kernel.Param, comptime index_value: usize) Value(param) {
62 return switch (param) {
63 .scalar => |dtype| builder.typedArgument(dtype, index_value),
64 .buffer => |buffer| builder.bufferArgument(buffer.dtype, index_value),
65 };
66 }
67
68 fn parameterSpec(comptime parameters: anytype, comptime name_value: anytype) kernel.Param {
69 if (comptime !named(parameters)) {
70 @compileError("named kernel parameter lookup requires named Program parameters");
71 }
72 const parameter_name = comptime name(name_value);
73 inline for (@typeInfo(@TypeOf(parameters)).@"struct".field_names) |field_name| {
74 if (comptime std.mem.eql(u8, field_name, parameter_name)) {
75 return @field(parameters, field_name);
76 }
77 }
78 @compileError("unknown kernel parameter: " ++ parameter_name);
79 }
80
81 fn parameterIndex(comptime parameters: anytype, comptime name_value: anytype) usize {
82 if (comptime !named(parameters)) {
83 @compileError("named kernel parameter lookup requires named Program parameters");
84 }
85 const parameter_name = comptime name(name_value);
86 comptime var found: ?usize = null;
87 inline for (@typeInfo(@TypeOf(parameters)).@"struct".field_names, 0..) |field_name, index_value| {
88 if (comptime std.mem.eql(u8, field_name, parameter_name)) {
89 found = index_value;
90 }
91 }
92 return found orelse @compileError("unknown kernel parameter: " ++ parameter_name);
93 }
94
95 fn name(comptime name_value: anytype) []const u8 {
96 return switch (@typeInfo(@TypeOf(name_value))) {
97 .enum_literal => @tagName(name_value),
98 else => @compileError("kernel parameter names must be enum literals"),
99 };
100 }