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 }