lib/accy/src/tensor/dsl/parameter.zig

daab053ee43316e1809a84551d573ddd1e5bf3d2

  1 const std = @import("std");
  2 const tensor = @import("../root.zig");
  3 const autodiff = tensor.autodiff;
  4 const batch = tensor.batch;
  5 const trace = tensor.trace;
  6 
  7 pub fn Standard(comptime parameters: anytype) type {
  8     return struct {
  9         pub const count = arity(parameters);
 10 
 11         pub fn index(comptime name_value: anytype) usize {
 12             return parameterIndex(parameters, name_value);
 13         }
 14 
 15         pub fn indices(comptime names: anytype) *const [nameCount(names)]usize {
 16             return parameterIndices(parameters, names);
 17         }
 18 
 19         pub fn axes(comptime axis_spec: anytype) *const [count]batch.Axis {
 20             return parameterAxes(parameters, axis_spec);
 21         }
 22     };
 23 }
 24 
 25 pub fn Jvp(comptime SourceLayout: type, comptime parameters: anytype, comptime options: autodiff.JvpOptions) type {
 26     return struct {
 27         pub const count = SourceLayout.count + options.wrt.len;
 28 
 29         pub fn index(comptime name_value: anytype) usize {
 30             return SourceLayout.index(name_value);
 31         }
 32 
 33         pub fn indices(comptime names: anytype) *const [nameCount(names)]usize {
 34             return SourceLayout.indices(names);
 35         }
 36 
 37         pub fn axes(comptime axis_spec: anytype) *const [count]batch.Axis {
 38             return jvpParameterAxes(SourceLayout, parameters, options, axis_spec);
 39         }
 40     };
 41 }
 42 
 43 pub fn Args(comptime parameters: anytype) type {
 44     return struct {
 45         values: []const trace.Value,
 46 
 47         pub fn param(self: @This(), comptime name_value: anytype) trace.Value {
 48             return self.values[parameterIndex(parameters, name_value)];
 49         }
 50     };
 51 }
 52 
 53 pub fn named(comptime parameters: anytype) bool {
 54     return switch (@typeInfo(@TypeOf(parameters))) {
 55         .@"struct" => |info| !info.is_tuple,
 56         else => false,
 57     };
 58 }
 59 
 60 pub fn arity(comptime parameters: anytype) usize {
 61     return if (comptime named(parameters))
 62         @typeInfo(@TypeOf(parameters)).@"struct".field_names.len
 63     else
 64         parameters.len;
 65 }
 66 
 67 pub fn nameCount(comptime names: anytype) usize {
 68     return switch (@typeInfo(@TypeOf(names))) {
 69         .@"struct" => |info| info.field_names.len,
 70         .enum_literal => 1,
 71         else => @compileError("tensor parameter names must be enum literals or a tuple of enum literals"),
 72     };
 73 }
 74 
 75 pub fn wrtOffset(comptime wrt: []const usize, comptime parameter_index: usize) usize {
 76     inline for (wrt, 0..) |item, index_value| {
 77         if (item == parameter_index) return index_value;
 78     }
 79     @compileError("tensor jvp tangent axis references a parameter outside wrt");
 80 }
 81 
 82 fn parameterIndices(comptime parameters: anytype, comptime names: anytype) *const [nameCount(names)]usize {
 83     comptime var indices_result: [nameCount(names)]usize = undefined;
 84     switch (@typeInfo(@TypeOf(names))) {
 85         .@"struct" => |info| {
 86             inline for (info.field_names, 0..) |field_name, index_value| {
 87                 indices_result[index_value] = parameterIndex(parameters, @field(names, field_name));
 88             }
 89         },
 90         .enum_literal => indices_result[0] = parameterIndex(parameters, names),
 91         else => unreachable,
 92     }
 93     const final = indices_result;
 94     return &final;
 95 }
 96 
 97 fn parameterAxes(comptime parameters: anytype, comptime axes: anytype) *const [arity(parameters)]batch.Axis {
 98     comptime var result = @as([arity(parameters)]batch.Axis, @splat(.none));
 99     const info = @typeInfo(@TypeOf(axes));
100     if (info != .@"struct" or info.@"struct".is_tuple) {
101         @compileError("tensor named axes must be a struct keyed by parameter name");
102     }
103     inline for (info.@"struct".field_names) |field_name| {
104         result[parameterIndex(parameters, field_name)] = @field(axes, field_name);
105     }
106     const final = result;
107     return &final;
108 }
109 
110 fn jvpParameterAxes(
111     comptime SourceLayout: type,
112     comptime parameters: anytype,
113     comptime options: autodiff.JvpOptions,
114     comptime axes: anytype,
115 ) *const [SourceLayout.count + options.wrt.len]batch.Axis {
116     comptime var result = @as([(SourceLayout.count + options.wrt.len)]batch.Axis, @splat(.none));
117     const Axes = @TypeOf(axes);
118     const info = @typeInfo(Axes);
119     if (info != .@"struct" or info.@"struct".is_tuple) {
120         @compileError("tensor jvp axes must be a struct with primals and tangents fields");
121     }
122 
123     if (comptime @hasField(Axes, "primals")) {
124         const primal_axes = SourceLayout.axes(axes.primals);
125         inline for (0..SourceLayout.count) |index_value| {
126             result[index_value] = primal_axes[index_value];
127         }
128     }
129 
130     if (comptime @hasField(Axes, "tangents")) {
131         const tangent_axes = jvpTangentAxes(parameters, options, axes.tangents);
132         inline for (0..options.wrt.len) |index_value| {
133             result[SourceLayout.count + index_value] = tangent_axes[index_value];
134         }
135     }
136 
137     const final = result;
138     return &final;
139 }
140 
141 fn jvpTangentAxes(
142     comptime parameters: anytype,
143     comptime options: autodiff.JvpOptions,
144     comptime axes: anytype,
145 ) *const [options.wrt.len]batch.Axis {
146     if (comptime !named(parameters)) {
147         @compileError("tensor jvp tangent axes require named Program parameters");
148     }
149     comptime var result = @as([options.wrt.len]batch.Axis, @splat(.none));
150     const Axes = @TypeOf(axes);
151     const info = @typeInfo(Axes);
152     if (info != .@"struct" or info.@"struct".is_tuple) {
153         @compileError("tensor jvp tangent axes must be a struct keyed by differentiated parameter name");
154     }
155     inline for (info.@"struct".field_names) |field_name| {
156         const parameter_index = parameterIndex(parameters, field_name);
157         const tangent_index = wrtOffset(options.wrt, parameter_index);
158         result[tangent_index] = @field(axes, field_name);
159     }
160     const final = result;
161     return &final;
162 }
163 
164 fn parameterIndex(comptime parameters: anytype, comptime name_value: anytype) usize {
165     if (comptime !named(parameters)) {
166         @compileError("named tensor parameter lookup requires named Program parameters");
167     }
168     const parameter_name = comptime name(name_value);
169     comptime var found: ?usize = null;
170     inline for (@typeInfo(@TypeOf(parameters)).@"struct".field_names, 0..) |field_name, index_value| {
171         if (comptime std.mem.eql(u8, field_name, parameter_name)) {
172             found = index_value;
173         }
174     }
175     return found orelse @compileError("unknown tensor parameter: " ++ parameter_name);
176 }
177 
178 fn name(comptime name_value: anytype) []const u8 {
179     return switch (@typeInfo(@TypeOf(name_value))) {
180         .enum_literal => @tagName(name_value),
181         .pointer => |pointer| blk: {
182             if (pointer.child == u8) break :blk name_value[0..];
183             if (@typeInfo(pointer.child) == .array) {
184                 const array = @typeInfo(pointer.child).array;
185                 if (array.child == u8) break :blk name_value[0..array.len];
186             }
187             @compileError("tensor parameter name pointers must point at string literals");
188         },
189         .array => |array| blk: {
190             if (array.child == u8) break :blk name_value[0..];
191             @compileError("tensor parameter name arrays must contain bytes");
192         },
193         else => @compileError("tensor parameter names must be enum literals or string literals"),
194     };
195 }