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 }