lib/accy/src/tensor/grad.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const gpu = @import("gpu");
3 const accy = @import("../root.zig");
4 const autodiff = @import("autodiff.zig");
5 const emit = @import("emit.zig");
6 const execute = @import("execute.zig");
7 const gating = @import("accy_validation_gating");
8 const hook = @import("hook.zig");
9 const interpret_mod = @import("interpret/root.zig");
10 const lower = @import("lower.zig");
11 const program_mod = @import("program.zig");
12 const reverse = @import("reverse.zig");
13 const trace = @import("trace/root.zig");
14 const unroll = @import("unroll.zig");
15 const types = @import("type/root.zig");
16
17 pub const Options = autodiff.LinearizeOptions;
18
19 pub fn grad(
20 allocator: std.mem.Allocator,
21 source: *const program_mod.Program,
22 options: Options,
23 ) !program_mod.Program {
24 try validate(source, options);
25
26 var linearized = try autodiff.linearize(allocator, source, options);
27 defer linearized.deinit();
28
29 var transposed = try reverse.pullback(allocator, &linearized, .{});
30 defer transposed.deinit();
31
32 return seedGradient(allocator, &transposed);
33 }
34
35 pub fn valueAndGrad(
36 allocator: std.mem.Allocator,
37 source: *const program_mod.Program,
38 options: Options,
39 ) !program_mod.Program {
40 try validate(source, options);
41
42 var linearized = try autodiff.linearize(allocator, source, options);
43 defer linearized.deinit();
44
45 var transposed = try reverse.pullback(allocator, &linearized, .{ .keep_primal_outputs = true });
46 defer transposed.deinit();
47
48 return seedGradient(allocator, &transposed);
49 }
50
51 pub fn gradWith(
52 allocator: std.mem.Allocator,
53 source: *const program_mod.Program,
54 options: Options,
55 hooks: anytype,
56 ) !program_mod.Program {
57 try validate(source, options);
58
59 var linearized = try autodiff.linearizeWith(allocator, source, options, hooks);
60 defer linearized.deinit();
61
62 var builder = try trace.Builder.init(allocator, linearized.program.name);
63 errdefer builder.deinit();
64
65 const graph = interpret_mod.Graph{ .builder = &builder };
66 const pullback_initial = hook.attach("pullback", hooks, graph);
67 var transposed = if (comptime hook.has(@TypeOf(hooks), "vjp"))
68 try reverse.pullbackWithRules(allocator, &linearized, .{}, pullback_initial, hooks.vjp)
69 else
70 try reverse.pullbackWith(allocator, &linearized, .{}, pullback_initial);
71 defer transposed.deinit();
72
73 return seedGradientWith(allocator, &transposed, hooks);
74 }
75
76 pub fn gradWithRules(
77 allocator: std.mem.Allocator,
78 source: *const program_mod.Program,
79 options: Options,
80 jvp_rules: anytype,
81 vjp_rules: anytype,
82 ) !program_mod.Program {
83 if (source.containsScan()) {
84 var expanded = try unroll.apply(allocator, source);
85 defer expanded.deinit();
86 return gradWithRules(allocator, &expanded, options, jvp_rules, vjp_rules);
87 }
88
89 try validate(source, options);
90
91 var linearize_builder = try trace.Builder.init(allocator, source.name);
92 errdefer linearize_builder.deinit();
93 const linearize_graph = interpret_mod.Graph{ .builder = &linearize_builder };
94 const linear = autodiff.semantics(source, linearize_graph, options);
95 var linearized = try interpret_mod.run(
96 allocator,
97 source,
98 interpret_mod.layer(autodiff.Dual, linear, jvp_rules),
99 );
100 defer linearized.deinit();
101
102 var pullback_builder = try trace.Builder.init(allocator, linearized.program.name);
103 errdefer pullback_builder.deinit();
104 const pullback_graph = interpret_mod.Graph{ .builder = &pullback_builder };
105 var transposed = try reverse.pullbackWithRules(allocator, &linearized, .{}, pullback_graph, vjp_rules);
106 defer transposed.deinit();
107
108 return seedGradient(allocator, &transposed);
109 }
110
111 pub fn interpret(
112 allocator: std.mem.Allocator,
113 source: *const program_mod.Program,
114 options: Options,
115 initial: anytype,
116 ) !interpret_mod.result(@TypeOf(initial)) {
117 return interpretWith(allocator, source, options, .{}, initial);
118 }
119
120 pub fn interpretWith(
121 allocator: std.mem.Allocator,
122 source: *const program_mod.Program,
123 options: Options,
124 hooks: anytype,
125 initial: anytype,
126 ) !interpret_mod.result(@TypeOf(initial)) {
127 if (comptime @hasDecl(@TypeOf(initial), "attach")) {
128 try validate(source, options);
129
130 var linearized = try autodiff.linearizeWith(allocator, source, options, hooks);
131 defer linearized.deinit();
132
133 var pullback_builder = try trace.Builder.init(allocator, linearized.program.name);
134 errdefer pullback_builder.deinit();
135
136 const pullback_graph = interpret_mod.Graph{ .builder = &pullback_builder };
137 var transposed = try reverse.pullbackWith(allocator, &linearized, .{}, hook.attach("pullback", hooks, pullback_graph));
138 defer transposed.deinit();
139
140 var builder = try trace.Builder.init(allocator, transposed.program.name);
141 errdefer builder.deinit();
142
143 const graph = interpret_mod.Graph{ .builder = &builder };
144 return seedGradientInto(allocator, &transposed, hook.attach("seed", hooks, initial.attach(graph)));
145 }
146
147 var differentiated = try gradWith(allocator, source, options, hooks);
148 defer differentiated.deinit();
149
150 return interpret_mod.run(allocator, &differentiated, initial);
151 }
152
153 pub fn validate(source: *const program_mod.Program, options: Options) !void {
154 try validateScalarOutput(source);
155 try autodiff.validate(source, options);
156 }
157
158 fn seedGradient(allocator: std.mem.Allocator, transposed: *const reverse.Pullback) !program_mod.Program {
159 var builder = try trace.Builder.init(allocator, transposed.program.name);
160 errdefer builder.deinit();
161
162 const graph = interpret_mod.Graph{ .builder = &builder };
163 return seedGradientInto(allocator, transposed, graph);
164 }
165
166 fn seedGradientWith(allocator: std.mem.Allocator, transposed: *const reverse.Pullback, hooks: anytype) !program_mod.Program {
167 var builder = try trace.Builder.init(allocator, transposed.program.name);
168 errdefer builder.deinit();
169
170 const graph = interpret_mod.Graph{ .builder = &builder };
171 return seedGradientInto(allocator, transposed, hook.attach("seed", hooks, graph));
172 }
173
174 fn seedGradientInto(allocator: std.mem.Allocator, transposed: *const reverse.Pullback, initial: anytype) !SeedSemantics(@TypeOf(initial)).Result {
175 if (transposed.seed_parameter_count != 1) return error.GradientRequiresScalarOutput;
176
177 return interpret_mod.run(allocator, &transposed.program, SeedSemantics(@TypeOf(initial)){
178 .next = initial,
179 .seed_parameter = transposed.primal_parameter_count,
180 });
181 }
182
183 fn SeedSemantics(comptime Next: type) type {
184 return struct {
185 next: Next,
186 seed_parameter: usize,
187
188 pub const Value: type = trace.Value;
189 pub const Result: type = Next.Result;
190
191 pub fn operation(self: *@This(), step: *interpret_mod.Step(Value)) !Value {
192 var buffer: [program_mod.max_operation_operands]Value = undefined;
193 return self.bind(step.op, interpret_mod.arguments(Value, step.op, step.values, &buffer));
194 }
195
196 pub fn bind(self: *@This(), op: *const program_mod.Operation, args: []const Value) !Value {
197 switch (op.kind) {
198 .parameter => |parameter| {
199 if (parameter.index == self.seed_parameter) return emit.fullFloat(self, op.result, 1.0);
200 },
201 else => {},
202 }
203 return self.next.bind(op, args);
204 }
205
206 pub fn finish(self: *@This(), outputs: []const Value) !Result {
207 return self.next.finish(outputs);
208 }
209
210 pub fn builderHandle(self: *@This()) *trace.Builder {
211 return self.next.builderHandle();
212 }
213 };
214 }
215
216 fn validateScalarOutput(source: *const program_mod.Program) !void {
217 if (source.outputs.len != 1) return error.GradientRequiresScalarOutput;
218 if (source.typeOf(source.outputs[0]).rank() != 0) return error.GradientRequiresScalarOutput;
219 }
220
221 fn elementwiseLoss(_: *trace.Builder, args: []const trace.Value) !trace.Value {
222 const product = try args[0].mul(args[1]);
223 return try product.sum(.lane);
224 }
225
226 fn embeddingLookupLoss(_: *trace.Builder, args: []const trace.Value) !trace.Value {
227 const gathered = try args[0].gather(args[1], .vocab);
228 return try gathered.sum(.{ .token, .channel });
229 }
230
231 test "tensor grad derives scalar-output gradients through pullback" {
232 var source = try trace.define(std.testing.allocator, "grad_elementwise", &.{
233 types.spec(.f32, .{ .lane = 4 }),
234 types.spec(.f32, .{ .lane = 4 }),
235 }, elementwiseLoss);
236 defer source.deinit();
237
238 var differentiated = try grad(std.testing.allocator, &source, .{ .wrt = &.{ 0, 1 } });
239 defer differentiated.deinit();
240
241 try std.testing.expectEqual(@as(usize, 2), differentiated.parameters.len);
242 try std.testing.expectEqual(@as(usize, 2), differentiated.outputs.len);
243 try types.expectExtents(&.{4}, differentiated.typeOf(differentiated.outputs[0]));
244 try types.expectExtents(&.{4}, differentiated.typeOf(differentiated.outputs[1]));
245 }
246
247 test "tensor grad transposes gather with scatter add" {
248 var source = try trace.define(std.testing.allocator, "grad_embedding_lookup", &.{
249 types.spec(.f32, .{ .vocab = 16, .channel = 4 }),
250 types.spec(.i32, .{ .token = 3 }),
251 }, embeddingLookupLoss);
252 defer source.deinit();
253
254 var differentiated = try grad(std.testing.allocator, &source, .{ .wrt = &.{0} });
255 defer differentiated.deinit();
256
257 try std.testing.expectEqual(@as(usize, 2), differentiated.parameters.len);
258 try types.expectExtents(&.{ 16, 4 }, differentiated.typeOf(differentiated.parameters[0]));
259 try types.expectExtents(&.{3}, differentiated.typeOf(differentiated.parameters[1]));
260 try std.testing.expectEqual(@as(usize, 1), differentiated.outputs.len);
261 try types.expectExtents(&.{ 16, 4 }, differentiated.typeOf(differentiated.outputs[0]));
262
263 var scatter_adds: usize = 0;
264 for (differentiated.operations) |op| {
265 switch (op.kind) {
266 .scatter_add => scatter_adds += 1,
267 else => {},
268 }
269 }
270 try std.testing.expect(scatter_adds >= 1);
271 }
272
273 const GeneratedGradCounts = struct {
274 constant: usize = 0,
275 mul: usize = 0,
276 broadcast_in_dim: usize = 0,
277 };
278
279 const GeneratedGradCounter = struct {
280 counts: *GeneratedGradCounts,
281
282 pub fn constant(self: *@This(), ctx: anytype) !trace.Value {
283 if (ctx.op.id.index == program_mod.synthetic_id.index) self.counts.constant += 1;
284 return ctx.default();
285 }
286
287 pub fn mul(self: *@This(), ctx: anytype) !trace.Value {
288 if (ctx.op.id.index == program_mod.synthetic_id.index) self.counts.mul += 1;
289 return ctx.default();
290 }
291
292 pub fn broadcastInDim(self: *@This(), ctx: anytype) !trace.Value {
293 if (ctx.op.id.index == program_mod.synthetic_id.index) self.counts.broadcast_in_dim += 1;
294 return ctx.default();
295 }
296 };
297
298 test "tensor gradWith routes generated ops through user semantics" {
299 var source = try trace.define(std.testing.allocator, "grad_with_elementwise", &.{
300 types.spec(.f32, .{ .lane = 4 }),
301 types.spec(.f32, .{ .lane = 4 }),
302 }, elementwiseLoss);
303 defer source.deinit();
304
305 var counts = GeneratedGradCounts{};
306 var seed_counts = GeneratedGradCounts{};
307 var differentiated = try gradWith(
308 std.testing.allocator,
309 &source,
310 .{ .wrt = &.{ 0, 1 } },
311 .{
312 .linearize = interpret_mod.bind(GeneratedGradCounter{ .counts = &counts }),
313 .pullback = interpret_mod.bind(GeneratedGradCounter{ .counts = &counts }),
314 .seed = interpret_mod.bind(GeneratedGradCounter{ .counts = &seed_counts }),
315 },
316 );
317 defer differentiated.deinit();
318
319 try std.testing.expect(counts.mul >= 4);
320 try std.testing.expect(counts.broadcast_in_dim >= 1);
321 try std.testing.expectEqual(@as(usize, 1), seed_counts.constant);
322 try std.testing.expectEqual(@as(usize, 2), differentiated.parameters.len);
323 try std.testing.expectEqual(@as(usize, 2), differentiated.outputs.len);
324 try types.expectExtents(&.{4}, differentiated.typeOf(differentiated.outputs[0]));
325 try types.expectExtents(&.{4}, differentiated.typeOf(differentiated.outputs[1]));
326 }
327
328 fn reluLoss(builder: *trace.Builder, args: []const trace.Value) !trace.Value {
329 const zero = try builder.scalar(.f32, 0.0);
330 const rectified = try args[0].max(zero);
331 return try rectified.sum(.lane);
332 }
333
334 test "tensor grad differentiates relu through select" {
335 var source = try trace.define(std.testing.allocator, "grad_relu", &.{
336 types.spec(.f32, .{ .lane = 4 }),
337 }, reluLoss);
338 defer source.deinit();
339
340 var differentiated = try grad(std.testing.allocator, &source, .{ .wrt = &.{0} });
341 defer differentiated.deinit();
342
343 try std.testing.expectEqual(@as(usize, 1), differentiated.outputs.len);
344 try types.expectExtents(&.{4}, differentiated.typeOf(differentiated.outputs[0]));
345 }
346
347 fn cubeLoss(builder: *trace.Builder, args: []const trace.Value) !trace.Value {
348 const three = try builder.scalar(.f32, 3.0);
349 const cubed = try args[0].pow(three);
350 return try cubed.sum(.lane);
351 }
352
353 test "tensor grad differentiates pow with inactive exponent" {
354 var source = try trace.define(std.testing.allocator, "grad_cube", &.{
355 types.spec(.f32, .{ .lane = 4 }),
356 }, cubeLoss);
357 defer source.deinit();
358
359 var differentiated = try grad(std.testing.allocator, &source, .{ .wrt = &.{0} });
360 defer differentiated.deinit();
361
362 try std.testing.expectEqual(@as(usize, 1), differentiated.outputs.len);
363 try types.expectExtents(&.{4}, differentiated.typeOf(differentiated.outputs[0]));
364
365 var logs: usize = 0;
366 for (differentiated.operations) |op| {
367 switch (op.kind) {
368 .unary => |unary_op| {
369 if (unary_op.op == .log) logs += 1;
370 },
371 else => {},
372 }
373 }
374 try std.testing.expectEqual(@as(usize, 0), logs);
375 }
376
377 fn plainSumLoss(_: *trace.Builder, args: []const trace.Value) !trace.Value {
378 return try args[0].sum(.lane);
379 }
380
381 test "tensor grad preserves source parameters when no residuals are needed" {
382 var source = try trace.define(std.testing.allocator, "grad_plain_sum", &.{
383 types.spec(.f32, .{ .lane = 4 }),
384 }, plainSumLoss);
385 defer source.deinit();
386
387 var differentiated = try grad(std.testing.allocator, &source, .{ .wrt = &.{0} });
388 defer differentiated.deinit();
389
390 try std.testing.expectEqual(@as(usize, 1), differentiated.parameters.len);
391 try types.expectExtents(&.{4}, differentiated.typeOf(differentiated.parameters[0]));
392 try std.testing.expectEqual(@as(usize, 1), differentiated.outputs.len);
393 try types.expectExtents(&.{4}, differentiated.typeOf(differentiated.outputs[0]));
394
395 const x = [_]f32{ 1.0, 2.0, 3.0, 4.0 };
396 var gradient_out = @as([4]f32, @splat(0));
397 const outputs = [_][]u8{std.mem.sliceAsBytes(gradient_out[0..])};
398 try execute.runCpu(std.testing.allocator, &differentiated, &.{std.mem.sliceAsBytes(x[0..])}, outputs[0..]);
399 try std.testing.expectEqualSlices(f32, &.{ 1.0, 1.0, 1.0, 1.0 }, gradient_out[0..]);
400 }
401
402 test "tensor grad executes with an unused parameter on native CPU" {
403 try @import("../fixture/root.zig").requireNativeCpuArtifacts();
404 var source = try trace.define(std.testing.allocator, "grad_unused_parameter_cpu", &.{
405 types.spec(.f32, .{ .vocab = 4, .channel = 2 }),
406 types.spec(.i32, .{ .token = 3 }),
407 }, embeddingLookupLoss);
408 defer source.deinit();
409
410 var differentiated = try grad(std.testing.allocator, &source, .{ .wrt = &.{0} });
411 defer differentiated.deinit();
412
413 try std.testing.expectEqual(@as(usize, 2), differentiated.parameters.len);
414 try types.expectExtents(&.{ 4, 2 }, differentiated.typeOf(differentiated.parameters[0]));
415 try types.expectExtents(&.{3}, differentiated.typeOf(differentiated.parameters[1]));
416
417 const table = [_]f32{ 0.5, -1.0, 2.0, 0.25, 3.0, -2.0, 1.5, 0.75 };
418 const indices = [_]i32{ 2, 0, 2 };
419 var table_grad = @as([8]f32, @splat(-1.0));
420 const outputs = [_][]u8{std.mem.sliceAsBytes(table_grad[0..])};
421 try execute.runCpu(std.testing.allocator, &differentiated, &.{
422 std.mem.sliceAsBytes(table[0..]),
423 std.mem.sliceAsBytes(indices[0..]),
424 }, outputs[0..]);
425 try std.testing.expectEqualSlices(f32, &.{ 1.0, 1.0, 0.0, 0.0, 2.0, 2.0, 0.0, 0.0 }, table_grad[0..]);
426 }
427
428 test "tensor valueAndGrad emits the loss ahead of the gradients" {
429 var source = try trace.define(std.testing.allocator, "value_and_grad_elementwise", &.{
430 types.spec(.f32, .{ .lane = 4 }),
431 types.spec(.f32, .{ .lane = 4 }),
432 }, elementwiseLoss);
433 defer source.deinit();
434
435 var combined = try valueAndGrad(std.testing.allocator, &source, .{ .wrt = &.{ 0, 1 } });
436 defer combined.deinit();
437
438 try std.testing.expectEqual(@as(usize, 2), combined.parameters.len);
439 try types.expectExtents(&.{4}, combined.typeOf(combined.parameters[0]));
440 try types.expectExtents(&.{4}, combined.typeOf(combined.parameters[1]));
441 try std.testing.expectEqual(@as(usize, 3), combined.outputs.len);
442 try types.expectExtents(&.{}, combined.typeOf(combined.outputs[0]));
443 try types.expectExtents(&.{4}, combined.typeOf(combined.outputs[1]));
444 try types.expectExtents(&.{4}, combined.typeOf(combined.outputs[2]));
445 }
446
447 test "tensor valueAndGrad computes the loss and gradients in one launch" {
448 try @import("../fixture/root.zig").requireNativeCpuArtifacts();
449 var source = try trace.define(std.testing.allocator, "value_and_grad_numeric", &.{
450 types.spec(.f32, .{ .lane = 4 }),
451 types.spec(.f32, .{ .lane = 4 }),
452 }, elementwiseLoss);
453 defer source.deinit();
454
455 var combined = try valueAndGrad(std.testing.allocator, &source, .{ .wrt = &.{ 0, 1 } });
456 defer combined.deinit();
457
458 const x = [_]f32{ 1.0, 2.0, 3.0, 4.0 };
459 const y = [_]f32{ 0.5, -1.0, 2.0, 0.25 };
460 var loss_value: f32 = 0;
461 var x_grad = @as([4]f32, @splat(0));
462 var y_grad = @as([4]f32, @splat(0));
463 const outputs = [_][]u8{
464 std.mem.asBytes(&loss_value),
465 std.mem.sliceAsBytes(x_grad[0..]),
466 std.mem.sliceAsBytes(y_grad[0..]),
467 };
468 try execute.runCpu(std.testing.allocator, &combined, &.{
469 std.mem.sliceAsBytes(x[0..]),
470 std.mem.sliceAsBytes(y[0..]),
471 }, outputs[0..]);
472
473 try std.testing.expectApproxEqAbs(@as(f32, 5.5), loss_value, 0.000001);
474 try std.testing.expectEqualSlices(f32, y[0..], x_grad[0..]);
475 try std.testing.expectEqualSlices(f32, x[0..], y_grad[0..]);
476 }
477
478 fn maxLoss(builder: *trace.Builder, args: []const trace.Value) !trace.Value {
479 const init = try builder.scalar(.f32, -3.4e38);
480 return try args[0].reduce(init, .max, .lane);
481 }
482
483 fn maxLossWithInitParam(_: *trace.Builder, args: []const trace.Value) !trace.Value {
484 return try args[0].reduce(args[1], .max, .lane);
485 }
486
487 test "tensor grad differentiates reduce max" {
488 var source = try trace.define(std.testing.allocator, "grad_reduce_max", &.{
489 types.spec(.f32, .{ .lane = 4 }),
490 }, maxLoss);
491 defer source.deinit();
492
493 var differentiated = try grad(std.testing.allocator, &source, .{ .wrt = &.{0} });
494 defer differentiated.deinit();
495
496 try std.testing.expectEqual(@as(usize, 1), differentiated.outputs.len);
497 try types.expectExtents(&.{4}, differentiated.typeOf(differentiated.outputs[0]));
498 }
499
500 test "tensor grad differentiates reduce max init parameter" {
501 var source = try trace.define(std.testing.allocator, "grad_reduce_max_init", &.{
502 types.spec(.f32, .{ .lane = 4 }),
503 types.spec(.f32, .{}),
504 }, maxLossWithInitParam);
505 defer source.deinit();
506
507 var differentiated = try grad(std.testing.allocator, &source, .{ .wrt = &.{ 0, 1 } });
508 defer differentiated.deinit();
509
510 try std.testing.expectEqual(@as(usize, 2), differentiated.outputs.len);
511 try types.expectExtents(&.{4}, differentiated.typeOf(differentiated.outputs[0]));
512 try types.expectExtents(&.{}, differentiated.typeOf(differentiated.outputs[1]));
513 }
514
515 fn absLoss(builder: *trace.Builder, args: []const trace.Value) !trace.Value {
516 _ = builder;
517 const magnitude = try args[0].abs();
518 return try magnitude.sum(.lane);
519 }
520
521 test "tensor grad differentiates abs through the sign mask" {
522 var source = try trace.define(std.testing.allocator, "grad_abs", &.{
523 types.spec(.f32, .{ .lane = 4 }),
524 }, absLoss);
525 defer source.deinit();
526
527 var differentiated = try grad(std.testing.allocator, &source, .{ .wrt = &.{0} });
528 defer differentiated.deinit();
529 try types.expectExtents(&.{4}, differentiated.typeOf(differentiated.outputs[0]));
530 }
531
532 fn tanLoss(builder: *trace.Builder, args: []const trace.Value) !trace.Value {
533 _ = builder;
534 const slope = try args[0].tan();
535 return try slope.sum(.lane);
536 }
537
538 test "tensor grad differentiates tan" {
539 var source = try trace.define(std.testing.allocator, "grad_tan", &.{
540 types.spec(.f32, .{ .lane = 4 }),
541 }, tanLoss);
542 defer source.deinit();
543
544 var differentiated = try grad(std.testing.allocator, &source, .{ .wrt = &.{0} });
545 defer differentiated.deinit();
546 try types.expectExtents(&.{4}, differentiated.typeOf(differentiated.outputs[0]));
547 }
548
549 test "tensor grad requires scalar output" {
550 var source = try trace.define(std.testing.allocator, "grad_vector", &.{
551 types.spec(.f32, .{ .lane = 4 }),
552 types.spec(.f32, .{ .lane = 4 }),
553 }, struct {
554 fn body(_: *trace.Builder, args: []const trace.Value) !trace.Value {
555 return try args[0].mul(args[1]);
556 }
557 }.body);
558 defer source.deinit();
559
560 try std.testing.expectError(error.GradientRequiresScalarOutput, grad(std.testing.allocator, &source, .{ .wrt = &.{0} }));
561 }
562
563 fn batchedDotGeneralLoss(_: *trace.Builder, args: []const trace.Value) !trace.Value {
564 const product = try args[0].contract(args[1], .k);
565 const weighted = try product.mul(args[2]);
566 return try weighted.sum(.{ .b, .m, .n });
567 }
568
569 test "tensor grad executes batched dot_general gradients on native CPU" {
570 try @import("../fixture/root.zig").requireNativeCpuArtifacts();
571 const allocator = std.testing.allocator;
572 const batch = 2;
573 const m = 2;
574 const k = 3;
575 const n = 2;
576
577 var source = try trace.define(allocator, "grad_batched_dot_general_cpu", &.{
578 types.spec(.f32, .{ .b = batch, .m = m, .k = k }),
579 types.spec(.f32, .{ .b = batch, .k = k, .n = n }),
580 types.spec(.f32, .{ .b = batch, .m = m, .n = n }),
581 }, batchedDotGeneralLoss);
582 defer source.deinit();
583
584 var differentiated = try grad(allocator, &source, .{ .wrt = &.{ 0, 1 } });
585 defer differentiated.deinit();
586
587 var state = gpu.cpu.State.init(allocator);
588 defer state.deinit();
589 const options = lower.FragmentCompilerOptions{
590 .artifact_format = .cpu_object,
591 };
592 const compiled = try lower.compileFragment(allocator, state.handle(), &differentiated, options);
593 var fragment = try accy.executable.loadFragment(allocator, state.handle(), compiled, options);
594 defer fragment.deinit();
595
596 var lhs = [_]f32{
597 1.0, 2.0, 3.0,
598 4.0, 5.0, 6.0,
599 -1.0, 0.5, 2.0,
600 3.0, -2.0, 1.0,
601 };
602 var rhs = [_]f32{
603 7.0, 8.0,
604 9.0, 10.0,
605 11.0, 12.0,
606 0.25, -1.0,
607 2.0, -3.0,
608 4.0, 0.5,
609 };
610 var weights = [_]f32{
611 1.0, -0.5,
612 0.25, 2.0,
613 -1.5, 0.75,
614 3.0, -2.0,
615 };
616 var lhs_grad = @as([(batch * m * k)]f32, @splat(0.0));
617 var rhs_grad = @as([(batch * k * n)]f32, @splat(0.0));
618 var outputs = [_][]u8{
619 std.mem.sliceAsBytes(lhs_grad[0..]),
620 std.mem.sliceAsBytes(rhs_grad[0..]),
621 };
622 try accy.executable.invoke(fragment, allocator, allocator, &.{
623 std.mem.sliceAsBytes(lhs[0..]),
624 std.mem.sliceAsBytes(rhs[0..]),
625 std.mem.sliceAsBytes(weights[0..]),
626 }, &outputs);
627
628 var expected_lhs_grad = @as([(batch * m * k)]f32, @splat(0.0));
629 var expected_rhs_grad = @as([(batch * k * n)]f32, @splat(0.0));
630 for (0..batch) |b| {
631 for (0..m) |row| {
632 for (0..k) |inner| {
633 var sum: f32 = 0.0;
634 for (0..n) |col| {
635 sum += weights[(b * m + row) * n + col] * rhs[(b * k + inner) * n + col];
636 }
637 expected_lhs_grad[(b * m + row) * k + inner] = sum;
638 }
639 }
640 for (0..k) |inner| {
641 for (0..n) |col| {
642 var sum: f32 = 0.0;
643 for (0..m) |row| {
644 sum += lhs[(b * m + row) * k + inner] * weights[(b * m + row) * n + col];
645 }
646 expected_rhs_grad[(b * k + inner) * n + col] = sum;
647 }
648 }
649 }
650
651 try std.testing.expectEqual(@as(usize, 2), fragment.outputCount());
652 for (expected_lhs_grad, lhs_grad) |want, got| {
653 try std.testing.expectApproxEqAbs(want, got, 0.000001);
654 }
655 for (expected_rhs_grad, rhs_grad) |want, got| {
656 try std.testing.expectApproxEqAbs(want, got, 0.000001);
657 }
658 }
659
660 fn transposedProjectionLoss(_: *trace.Builder, args: []const trace.Value) !trace.Value {
661 const product = try args[0].builder.dotGeneralOp(args[0], args[1], &.{0}, &.{0}, &.{}, &.{});
662 return try product.sum(.{ .m, .n });
663 }
664
665 test "tensor grad executes noncanonical dot_general gradients on native CPU" {
666 try @import("../fixture/root.zig").requireNativeCpuArtifacts();
667 const allocator = std.testing.allocator;
668
669 var source = try trace.define(allocator, "grad_noncanonical_dot_cpu", &.{
670 types.spec(.f32, .{ .k = 3, .m = 2 }),
671 types.spec(.f32, .{ .k = 3, .n = 2 }),
672 }, transposedProjectionLoss);
673 defer source.deinit();
674
675 var differentiated = try grad(allocator, &source, .{ .wrt = &.{ 0, 1 } });
676 defer differentiated.deinit();
677
678 const lhs = [_]f32{ 1.0, 2.0, 3.0, 4.0, 5.0, 6.0 };
679 const rhs = [_]f32{ 0.5, -1.0, 2.0, 0.25, -0.75, 1.5 };
680 var lhs_grad = @as([6]f32, @splat(0.0));
681 var rhs_grad = @as([6]f32, @splat(0.0));
682 var outputs = [_][]u8{
683 std.mem.sliceAsBytes(lhs_grad[0..]),
684 std.mem.sliceAsBytes(rhs_grad[0..]),
685 };
686 try execute.runCpu(allocator, &differentiated, &.{
687 std.mem.sliceAsBytes(lhs[0..]),
688 std.mem.sliceAsBytes(rhs[0..]),
689 }, outputs[0..]);
690
691 for (0..3) |k| {
692 const rhs_row_sum = rhs[k * 2] + rhs[k * 2 + 1];
693 try std.testing.expectApproxEqAbs(rhs_row_sum, lhs_grad[k * 2], 0.000001);
694 try std.testing.expectApproxEqAbs(rhs_row_sum, lhs_grad[k * 2 + 1], 0.000001);
695 const lhs_row_sum = lhs[k * 2] + lhs[k * 2 + 1];
696 try std.testing.expectApproxEqAbs(lhs_row_sum, rhs_grad[k * 2], 0.000001);
697 try std.testing.expectApproxEqAbs(lhs_row_sum, rhs_grad[k * 2 + 1], 0.000001);
698 }
699 }
700
701 const attention_tokens = 4;
702 const attention_channels = 3;
703 const attention_mask_penalty: f32 = -30.0;
704
705 fn causalAttentionLoss(_: *trace.Builder, args: []const trace.Value) !trace.Value {
706 const queries = args[0];
707 const keys = args[1];
708 const values = args[2];
709 const mask = args[3];
710
711 const scores = try queries.contract(try keys.rename(.token, .key), .channel);
712 const masked = try scores.add(mask);
713 const centered = try masked.sub(try masked.max(.key));
714 const weights = try centered.exp();
715 const probs = try weights.div(try weights.sum(.key));
716 const mixed = try probs.contract(try values.rename(.token, .key), .key);
717 const squared = try mixed.mul(mixed);
718 return try squared.sum(.{ .token, .channel });
719 }
720
721 fn attentionNoise(seed: usize) f32 {
722 var state: u64 = @as(u64, @intCast(seed)) +% 0x9e3779b97f4a7c15;
723 state = (state ^ (state >> 30)) *% 0xbf58476d1ce4e5b9;
724 state = (state ^ (state >> 27)) *% 0x94d049bb133111eb;
725 state ^= state >> 31;
726 const unit = @as(f32, @floatFromInt(state & 0xffff)) / 65535.0;
727 return 2.0 * unit - 1.0;
728 }
729
730 test "tensor valueAndGrad matches finite differences through causal attention" {
731 try @import("../fixture/root.zig").requireNativeCpuArtifacts();
732 const allocator = std.testing.allocator;
733 const element_count = attention_tokens * attention_channels;
734
735 var source = try trace.define(allocator, "grad_causal_attention_fd", &.{
736 types.spec(.f32, .{ .token = attention_tokens, .channel = attention_channels }),
737 types.spec(.f32, .{ .token = attention_tokens, .channel = attention_channels }),
738 types.spec(.f32, .{ .token = attention_tokens, .channel = attention_channels }),
739 types.spec(.f32, .{ .token = attention_tokens, .key = attention_tokens }),
740 }, causalAttentionLoss);
741 defer source.deinit();
742
743 var combined = try valueAndGrad(allocator, &source, .{ .wrt = &.{ 0, 1, 2 } });
744 defer combined.deinit();
745
746 var inputs: [3][element_count]f32 = undefined;
747 for (&inputs, 0..) |*tensor_input, tensor_index| {
748 for (tensor_input, 0..) |*slot, element_index| {
749 slot.* = attentionNoise(tensor_index * 1000 + element_index);
750 }
751 }
752 var mask: [attention_tokens * attention_tokens]f32 = undefined;
753 for (0..attention_tokens) |row| {
754 for (0..attention_tokens) |col| {
755 mask[row * attention_tokens + col] = if (col <= row) 0.0 else attention_mask_penalty;
756 }
757 }
758
759 var loss_value: f32 = 0.0;
760 var gradients: [3][element_count]f32 = @splat(@splat(0.0));
761 var grad_outputs = [_][]u8{
762 std.mem.asBytes(&loss_value),
763 std.mem.sliceAsBytes(gradients[0][0..]),
764 std.mem.sliceAsBytes(gradients[1][0..]),
765 std.mem.sliceAsBytes(gradients[2][0..]),
766 };
767 try execute.runCpu(allocator, &combined, &.{
768 std.mem.sliceAsBytes(inputs[0][0..]),
769 std.mem.sliceAsBytes(inputs[1][0..]),
770 std.mem.sliceAsBytes(inputs[2][0..]),
771 std.mem.sliceAsBytes(mask[0..]),
772 }, grad_outputs[0..]);
773
774 var loss_executor = try execute.Cpu.init(allocator, &source);
775 defer loss_executor.deinit();
776
777 const step: f32 = 0.01;
778 for (0..3) |tensor_index| {
779 for (0..element_count) |element_index| {
780 var probes: [2]f32 = undefined;
781 for (&probes, [_]f32{ step, -step }) |*probe, offset| {
782 var perturbed = inputs;
783 perturbed[tensor_index][element_index] += offset;
784 var probe_loss: f32 = 0.0;
785 var probe_outputs = [_][]u8{std.mem.asBytes(&probe_loss)};
786 try loss_executor.launch(allocator, &.{
787 std.mem.sliceAsBytes(perturbed[0][0..]),
788 std.mem.sliceAsBytes(perturbed[1][0..]),
789 std.mem.sliceAsBytes(perturbed[2][0..]),
790 std.mem.sliceAsBytes(mask[0..]),
791 }, probe_outputs[0..]);
792 probe.* = probe_loss;
793 }
794 const finite_difference = (probes[0] - probes[1]) / (2.0 * step);
795 const analytic = gradients[tensor_index][element_index];
796 const tolerance = 0.002 + 0.02 * @abs(finite_difference);
797 try std.testing.expectApproxEqAbs(finite_difference, analytic, tolerance);
798 }
799 }
800 }
801
802 const custom_double_target = "accy.custom.double";
803
804 fn doubleSumLoss(builder: *trace.Builder, args: []const trace.Value) !trace.Value {
805 const doubled = try builder.customCall(custom_double_target, 1, &.{args[0]}, args[0].ty);
806 return try doubled.sum(.lane);
807 }
808
809 const DoubleJvpRule = struct {
810 applied: *usize,
811
812 pub fn bind(self: *@This(), ctx: anytype) !autodiff.Dual {
813 switch (ctx.op.kind) {
814 .custom_call => |custom| {
815 if (std.mem.eql(u8, custom.target, custom_double_target)) {
816 self.applied.* += 1;
817 const builder = ctx.builderHandle();
818 return .{
819 .primal = try builder.customCall(custom.target, custom.version, &.{ctx.args[0].primal}, ctx.op.result),
820 .tangent = try builder.customCall(custom.target, custom.version, &.{ctx.args[0].tangent}, ctx.op.result),
821 };
822 }
823 },
824 else => {},
825 }
826 return ctx.default();
827 }
828 };
829
830 const DoubleVjpRule = struct {
831 applied: *usize,
832
833 pub fn customCall(self: *@This(), ctx: anytype) !void {
834 self.applied.* += 1;
835 if (!ctx.operandIsActive(0)) return;
836 const builder = ctx.builderHandle();
837 const contribution = try builder.customCall(
838 ctx.custom.target,
839 ctx.custom.version,
840 &.{ctx.cotangent},
841 ctx.cotangent.ty,
842 );
843 try ctx.contribute(0, contribution);
844 }
845 };
846
847 test "tensor gradWithRules differentiates custom calls through both contracts" {
848 var source = try trace.define(std.testing.allocator, "grad_custom_double", &.{
849 types.spec(.f32, .{ .lane = 4 }),
850 }, doubleSumLoss);
851 defer source.deinit();
852
853 var jvp_applied: usize = 0;
854 var vjp_applied: usize = 0;
855 var differentiated = try gradWithRules(
856 std.testing.allocator,
857 &source,
858 .{ .wrt = &.{0} },
859 DoubleJvpRule{ .applied = &jvp_applied },
860 DoubleVjpRule{ .applied = &vjp_applied },
861 );
862 defer differentiated.deinit();
863
864 try std.testing.expectEqual(@as(usize, 1), jvp_applied);
865 try std.testing.expectEqual(@as(usize, 1), vjp_applied);
866 try std.testing.expectEqual(@as(usize, 1), differentiated.parameters.len);
867 try std.testing.expectEqual(@as(usize, 1), differentiated.outputs.len);
868 try types.expectExtents(&.{4}, differentiated.typeOf(differentiated.outputs[0]));
869
870 var custom_calls: usize = 0;
871 for (differentiated.operations) |op| {
872 switch (op.kind) {
873 .custom_call => |custom| {
874 try std.testing.expectEqualStrings(custom_double_target, custom.target);
875 custom_calls += 1;
876 },
877 else => {},
878 }
879 }
880 try std.testing.expect(custom_calls >= 1);
881 }
882
883 test "tensor gradWith accepts custom-call contracts through public hooks" {
884 var source = try trace.define(std.testing.allocator, "grad_public_custom_double", &.{
885 types.spec(.f32, .{ .lane = 4 }),
886 }, doubleSumLoss);
887 defer source.deinit();
888
889 var jvp_applied: usize = 0;
890 var vjp_applied: usize = 0;
891 var differentiated = try gradWith(
892 std.testing.allocator,
893 &source,
894 .{ .wrt = &.{0} },
895 .{
896 .jvp = interpret_mod.bind(DoubleJvpRule{ .applied = &jvp_applied }),
897 .vjp = DoubleVjpRule{ .applied = &vjp_applied },
898 },
899 );
900 defer differentiated.deinit();
901
902 try std.testing.expectEqual(@as(usize, 1), jvp_applied);
903 try std.testing.expectEqual(@as(usize, 1), vjp_applied);
904 try std.testing.expectEqual(@as(usize, 1), differentiated.parameters.len);
905 try std.testing.expectEqual(@as(usize, 1), differentiated.outputs.len);
906 try types.expectExtents(&.{4}, differentiated.typeOf(differentiated.outputs[0]));
907 }
908
909 test "tensor gradWithRules still rejects missing contracts by name" {
910 var source = try trace.define(std.testing.allocator, "grad_custom_opaque", &.{
911 types.spec(.f32, .{ .lane = 4 }),
912 }, doubleSumLoss);
913 defer source.deinit();
914
915 var vjp_applied: usize = 0;
916 try std.testing.expectError(error.CustomCallRequiresJvpContract, gradWithRules(
917 std.testing.allocator,
918 &source,
919 .{ .wrt = &.{0} },
920 .{},
921 DoubleVjpRule{ .applied = &vjp_applied },
922 ));
923
924 var jvp_applied: usize = 0;
925 try std.testing.expectError(error.CustomCallRequiresVjpContract, gradWithRules(
926 std.testing.allocator,
927 &source,
928 .{ .wrt = &.{0} },
929 DoubleJvpRule{ .applied = &jvp_applied },
930 .{},
931 ));
932 }
933
934 fn customDoubleKernel() type {
935 const Body = struct {
936 fn run(k: anytype, args: anytype) !void {
937 _ = try k.forEach1D("e", 4, args, struct {
938 fn each(inner: anytype, index: accy.kernel.Index1D, each_args: anytype) !void {
939 const x = try each_args.param(.x).load(inner, index);
940 const doubled = try x.add(inner, x);
941 try each_args.param(.dst).store(inner, doubled, index);
942 }
943 }.each);
944 }
945 };
946 return accy.kernel.logical.Program(.{
947 .name = "accy_custom_double_4_f32",
948 .parameters = .{
949 .dst = accy.kernel.dynamicBuffer(.f32),
950 .x = accy.kernel.dynamicBuffer(.f32),
951 },
952 .body = Body.run,
953 }).withSchedule(accy.kernel.logical.schedule.threadBlocks(.{ .x = 4 }));
954 }
955
956 test "tensor valueAndGrad executes loss and gradients in one launch on live CUDA" {
957 const allocator = std.testing.allocator;
958
959 try gating.skipIfBuildFlagDisabled(.cuda);
960 if (!gpu.cuda.platformSupported()) return gating.skip(.cuda, .unsupported_platform);
961 var state = gpu.cuda.State.initDevice(allocator, 0) catch |err| switch (err) {
962 error.RuntimeUnavailable => return gating.skip(.cuda, .cuda_device_missing),
963 else => return err,
964 };
965 defer state.deinit();
966 const handle = state.handle();
967
968 var source = try trace.define(allocator, "live_value_and_grad", &.{
969 types.spec(.f32, .{ .lane = 4 }),
970 types.spec(.f32, .{ .lane = 4 }),
971 }, elementwiseLoss);
972 defer source.deinit();
973
974 var combined = try valueAndGrad(allocator, &source, .{ .wrt = &.{ 0, 1 } });
975 defer combined.deinit();
976
977 const compiled = try lower.compileFragment(allocator, handle, &combined, .{});
978 var fragment = try accy.executable.loadFragment(allocator, handle, compiled, .{});
979 defer fragment.deinit();
980
981 var x = [_]f32{ 1.0, 2.0, 3.0, 4.0 };
982 var y = [_]f32{ 0.5, -1.0, 2.0, 0.25 };
983 var loss_value: f32 = -1.0;
984 var x_grad = @as([4]f32, @splat(-1.0));
985 var y_grad = @as([4]f32, @splat(-1.0));
986 var outputs = [_][]u8{
987 std.mem.asBytes(&loss_value),
988 std.mem.sliceAsBytes(x_grad[0..]),
989 std.mem.sliceAsBytes(y_grad[0..]),
990 };
991 try accy.executable.invoke(fragment, allocator, allocator, &.{
992 std.mem.sliceAsBytes(x[0..]),
993 std.mem.sliceAsBytes(y[0..]),
994 }, &outputs);
995
996 try std.testing.expectApproxEqAbs(@as(f32, 5.5), loss_value, 0.000001);
997 try std.testing.expectEqualSlices(f32, y[0..], x_grad[0..]);
998 try std.testing.expectEqualSlices(f32, x[0..], y_grad[0..]);
999 }
1000
1001 test "tensor gradWith custom-call gradient executes on live CUDA" {
1002 const allocator = std.testing.allocator;
1003
1004 try gating.skipIfBuildFlagDisabled(.cuda);
1005 if (!gpu.cuda.platformSupported()) return gating.skip(.cuda, .unsupported_platform);
1006 var state = gpu.cuda.State.initDevice(allocator, 0) catch |err| switch (err) {
1007 error.RuntimeUnavailable => return gating.skip(.cuda, .cuda_device_missing),
1008 else => return err,
1009 };
1010 defer state.deinit();
1011 const handle = state.handle();
1012
1013 var call_artifact = try customDoubleKernel().createKernelCallArtifact(
1014 allocator,
1015 accy.kernel.Limits.testing,
1016 handle,
1017 .{
1018 .target = custom_double_target,
1019 .version = 1,
1020 .format = .cuda_ptx,
1021 },
1022 );
1023 defer call_artifact.deinit();
1024 const registry = call_artifact.registry();
1025
1026 var source = try trace.define(allocator, "live_grad_custom_double", &.{
1027 types.spec(.f32, .{ .lane = 4 }),
1028 }, doubleSumLoss);
1029 defer source.deinit();
1030
1031 var jvp_applied: usize = 0;
1032 var vjp_applied: usize = 0;
1033 var differentiated = try gradWith(
1034 allocator,
1035 &source,
1036 .{ .wrt = &.{0} },
1037 .{
1038 .jvp = interpret_mod.bind(DoubleJvpRule{ .applied = &jvp_applied }),
1039 .vjp = DoubleVjpRule{ .applied = &vjp_applied },
1040 },
1041 );
1042 defer differentiated.deinit();
1043
1044 const options = lower.FragmentCompilerOptions{
1045 .kernel_call_registry = ®istry,
1046 };
1047 const compiled = try lower.compileFragment(allocator, handle, &differentiated, options);
1048 var fragment = try accy.executable.loadFragment(allocator, handle, compiled, options);
1049 defer fragment.deinit();
1050
1051 var x = [_]f32{ 1.0, 2.0, 3.0, 4.0 };
1052 const input_bytes = [_][]const u8{std.mem.sliceAsBytes(x[0..])};
1053 const bindings = try accy.executable.prepareInvocation(fragment, allocator, input_bytes[0..]);
1054 defer bindings.deinit();
1055 try bindings.launchWithOptions(allocator, .{});
1056
1057 var gradient = @as([4]f32, @splat(-1.0));
1058 try bindings.readOutput(0, std.mem.sliceAsBytes(gradient[0..]));
1059
1060 for (gradient) |value| {
1061 try std.testing.expectApproxEqAbs(@as(f32, 2.0), value, 0.000001);
1062 }
1063 }