lib/accy/src/tensor/test.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const gpu = @import("gpu");
3 const tensor = @import("root.zig");
4
5 const fixture = @import("../fixture/root.zig");
6 const types = tensor.types;
7 const program = tensor.program;
8 const dsl = tensor.dsl;
9 const interpret = tensor.interpret;
10 const random = tensor.random;
11 const trace = tensor.trace;
12 const transform = tensor.transform;
13 const autodiff = tensor.autodiff;
14 const batch = tensor.batch;
15 const reverse = tensor.reverse;
16 const gradient = tensor.gradient;
17 const nn = tensor.nn;
18 const lower = tensor.lower;
19 const execute = tensor.execute;
20 const function = tensor.function;
21 const unroll = tensor.unroll;
22 const DType = tensor.DType;
23 const Type = tensor.Type;
24 const Spec = tensor.Spec;
25 const Dim = tensor.Dim;
26 const Graph = tensor.Graph;
27 const Program = tensor.Program;
28 const program_product_name = tensor.program_product_name;
29 const Id = tensor.Id;
30 const Operation = tensor.Operation;
31 const Unary = tensor.Unary;
32 const Binary = tensor.Binary;
33 const Reducer = tensor.Reducer;
34 const Builder = tensor.Builder;
35 const Value = tensor.Value;
36 const Dual = tensor.Dual;
37 const TransformContext = tensor.TransformContext;
38 const spec = tensor.spec;
39 const specDims = tensor.specDims;
40 const define = tensor.define;
41 const rewrite = tensor.rewrite;
42 const linearize = tensor.linearize;
43 const linearizeWith = tensor.linearizeWith;
44 const jvp = tensor.jvp;
45 const jvpWith = tensor.jvpWith;
46 const pullback = tensor.pullback;
47 const pullbackWith = tensor.pullbackWith;
48 const grad = tensor.grad;
49 const valueAndGrad = tensor.valueAndGrad;
50 const gradWith = tensor.gradWith;
51 const gradWithRules = tensor.gradWithRules;
52 const embedding = tensor.embedding;
53 const embeddingNamed = tensor.embeddingNamed;
54 const sparseCrossEntropy = tensor.sparseCrossEntropy;
55 const sparseCrossEntropyNamed = tensor.sparseCrossEntropyNamed;
56 const sparseCrossEntropyMean = tensor.sparseCrossEntropyMean;
57 const sparseCrossEntropyMeanNamed = tensor.sparseCrossEntropyMeanNamed;
58 const vmap = tensor.vmap;
59 const vmapWith = tensor.vmapWith;
60 const LinearizeOptions = tensor.LinearizeOptions;
61 const Linearization = tensor.Linearization;
62 const Pullback = tensor.Pullback;
63 const PullbackOptions = tensor.PullbackOptions;
64 const GradOptions = tensor.GradOptions;
65 const JvpOptions = tensor.JvpOptions;
66 const VmapOptions = tensor.VmapOptions;
67 const BatchAxis = tensor.BatchAxis;
68 const mappedAxis = tensor.mappedAxis;
69 const runCpu = tensor.runCpu;
70 const CpuExecutor = tensor.CpuExecutor;
71 const Function = tensor.Function;
72 const toSemanticModule = tensor.toSemanticModule;
73 const prepare = tensor.prepare;
74 const prepareWith = tensor.prepareWith;
75 const prepareFragment = tensor.prepareFragment;
76 const createArtifactJob = tensor.createArtifactJob;
77 const createArtifactJobFromPreparedJob = tensor.createArtifactJobFromPreparedJob;
78 const compileFragmentFromArtifactJob = tensor.compileFragmentFromArtifactJob;
79 const compileFragment = tensor.compileFragment;
80 const compileFragmentFromPreparedJob = tensor.compileFragmentFromPreparedJob;
81 const BackendPreparedJob = tensor.BackendPreparedJob;
82 const BackendPreparationRunOptions = tensor.BackendPreparationRunOptions;
83 const GeneratedScheduleKind = tensor.GeneratedScheduleKind;
84 const GeneratedSchedule = tensor.GeneratedSchedule;
85 const GeneratedKernelProgram = tensor.GeneratedKernelProgram;
86 const GeneratedKernelSummary = tensor.GeneratedKernelSummary;
87 const GeneratedKernelSummaries = tensor.GeneratedKernelSummaries;
88 const ArtifactJob = tensor.ArtifactJob;
89 const BackendHandle = tensor.BackendHandle;
90 const CompiledFragment = tensor.CompiledFragment;
91 const LoadedFragment = tensor.LoadedFragment;
92 const FragmentCompilerOptions = tensor.FragmentCompilerOptions;
93 const FragmentCompilerCache = tensor.FragmentCompilerCache;
94 const FragmentCompilerCacheUpdate = tensor.FragmentCompilerCacheUpdate;
95 const ArtifactKernelSource = tensor.ArtifactKernelSource;
96 const ArtifactKernelSummary = tensor.ArtifactKernelSummary;
97 const ArtifactKernelSummaries = tensor.ArtifactKernelSummaries;
98
99 test {
100 _ = @import("wire/test.zig");
101 _ = @import("session/test.zig");
102 _ = @import("dsl/test.zig");
103 _ = @import("interpret/test.zig");
104 _ = @import("random/test.zig");
105 _ = @import("trace/test.zig");
106 _ = @import("type/test.zig");
107 @import("test_discovery").discover(tensor);
108 }
109
110 fn denseBody(_: *Builder, args: []const Value) !Value {
111 const linear = try args[0].contract(args[1], .k);
112 return try (try linear.add(args[2])).tanh();
113 }
114
115 test "accy tensor namespace traces rewrites and lowers a dense program" {
116 var traced = try define(std.testing.allocator, "dense_root", &.{
117 spec(.f32, .{ .m = 2, .k = 4 }),
118 spec(.f32, .{ .k = 4, .n = 3 }),
119 spec(.f32, .{ .n = 3 }),
120 }, denseBody);
121 defer traced.deinit();
122
123 var rewritten = try rewrite(std.testing.allocator, &traced, struct {}{});
124 defer rewritten.deinit();
125
126 const module = try toSemanticModule(std.testing.allocator, &rewritten);
127 defer module.deinit();
128
129 try module.verify();
130 }
131
132 const DropAddZero = struct {
133 pub fn add(_: *@This(), ctx: *TransformContext) !?Value {
134 if (ctx.isZero(1)) return ctx.arg(0);
135 if (ctx.isZero(0)) return ctx.arg(1);
136 return null;
137 }
138 };
139
140 fn compositionBody(builder: *Builder, args: []const Value) !Value {
141 const zero = try builder.full(.f32, .{ .lane = 4 }, 0.0);
142 return try (try args[0].mul(args[1])).add(zero);
143 }
144
145 test "accy tensor transforms compose through rewrite linearize rewrite and lower" {
146 var traced = try define(std.testing.allocator, "composition", &.{
147 spec(.f32, .{ .lane = 4 }),
148 spec(.f32, .{ .lane = 4 }),
149 }, compositionBody);
150 defer traced.deinit();
151
152 var simplified = try rewrite(std.testing.allocator, &traced, DropAddZero{});
153 defer simplified.deinit();
154
155 var differentiated = try linearize(std.testing.allocator, &simplified, .{ .wrt = &.{ 0, 1 } });
156 defer differentiated.deinit();
157
158 var cleaned = try rewrite(std.testing.allocator, &differentiated.program, DropAddZero{});
159 defer cleaned.deinit();
160
161 const module = try toSemanticModule(std.testing.allocator, &cleaned);
162 defer module.deinit();
163
164 try std.testing.expectEqual(@as(usize, 2), differentiated.primal_parameter_count);
165 try std.testing.expectEqual(@as(usize, 2), differentiated.tangent_parameter_count);
166 try std.testing.expectEqual(@as(usize, 2), cleaned.outputs.len);
167 try module.verify();
168 }
169
170 test "accy tensor vmap composes with linearize and lowering" {
171 var traced = try define(std.testing.allocator, "vmap_linearize", &.{
172 spec(.f32, .{ .lane = 4 }),
173 spec(.f32, .{ .lane = 4 }),
174 }, compositionBody);
175 defer traced.deinit();
176
177 var differentiated = try linearize(std.testing.allocator, &traced, .{ .wrt = &.{ 0, 1 } });
178 defer differentiated.deinit();
179
180 var batched = try vmap(std.testing.allocator, &differentiated.program, .{
181 .axis_size = 8,
182 .in_axes = &.{ mappedAxis(0), mappedAxis(0), mappedAxis(0), mappedAxis(0) },
183 });
184 defer batched.deinit();
185
186 const module = try toSemanticModule(std.testing.allocator, &batched);
187 defer module.deinit();
188
189 try std.testing.expectEqual(@as(usize, 4), batched.parameters.len);
190 try std.testing.expectEqual(@as(usize, 2), batched.outputs.len);
191 try types.expectExtents(&.{ 8, 4 }, batched.typeOf(batched.outputs[0]));
192 try types.expectExtents(&.{ 8, 4 }, batched.typeOf(batched.outputs[1]));
193 try module.verify();
194 }
195
196 test "accy tensor pullback composes with vmap and lowering" {
197 var traced = try define(std.testing.allocator, "pullback_vmap", &.{
198 spec(.f32, .{ .lane = 4 }),
199 spec(.f32, .{ .lane = 4 }),
200 }, compositionBody);
201 defer traced.deinit();
202
203 var differentiated = try linearize(std.testing.allocator, &traced, .{ .wrt = &.{ 0, 1 } });
204 defer differentiated.deinit();
205
206 var transposed = try pullback(std.testing.allocator, &differentiated, .{});
207 defer transposed.deinit();
208
209 var batched = try vmap(std.testing.allocator, &transposed.program, .{
210 .axis_size = 8,
211 .in_axes = &.{ mappedAxis(0), mappedAxis(0), mappedAxis(0) },
212 });
213 defer batched.deinit();
214
215 const module = try toSemanticModule(std.testing.allocator, &batched);
216 defer module.deinit();
217
218 try std.testing.expectEqual(@as(usize, 3), batched.parameters.len);
219 try std.testing.expectEqual(@as(usize, 2), batched.outputs.len);
220 try types.expectExtents(&.{ 8, 4 }, batched.typeOf(batched.outputs[0]));
221 try types.expectExtents(&.{ 8, 4 }, batched.typeOf(batched.outputs[1]));
222 try module.verify();
223 }
224
225 fn gradientBody(_: *Builder, args: []const Value) !Value {
226 const product = try args[0].mul(args[1]);
227 return try product.sum(.lane);
228 }
229
230 test "accy tensor grad composes with vmap and lowering" {
231 var traced = try define(std.testing.allocator, "grad_vmap", &.{
232 spec(.f32, .{ .lane = 4 }),
233 spec(.f32, .{ .lane = 4 }),
234 }, gradientBody);
235 defer traced.deinit();
236
237 var differentiated = try grad(std.testing.allocator, &traced, .{ .wrt = &.{ 0, 1 } });
238 defer differentiated.deinit();
239
240 var batched = try vmap(std.testing.allocator, &differentiated, .{
241 .axis_size = 8,
242 .in_axes = &.{ mappedAxis(0), mappedAxis(0) },
243 });
244 defer batched.deinit();
245
246 const module = try toSemanticModule(std.testing.allocator, &batched);
247 defer module.deinit();
248
249 try std.testing.expectEqual(@as(usize, 2), batched.parameters.len);
250 try std.testing.expectEqual(@as(usize, 2), batched.outputs.len);
251 try types.expectExtents(&.{ 8, 4 }, batched.typeOf(batched.outputs[0]));
252 try types.expectExtents(&.{ 8, 4 }, batched.typeOf(batched.outputs[1]));
253 try module.verify();
254 }
255
256 const scan_differential_steps = 5;
257 const scan_differential_width = 4;
258
259 fn buildScanDifferentialScan(allocator: std.mem.Allocator) !Graph {
260 var builder = try Builder.init(allocator, "scan_differential");
261 errdefer builder.deinit();
262 const x0 = try builder.input(.f32, .{ .lane = scan_differential_width });
263 const c = try builder.input(.f32, .{ .lane = scan_differential_width });
264 const acc0 = try builder.full(.f32, .{ .lane = scan_differential_width }, 0.0);
265 const walked = try builder.scan(.{
266 .length = scan_differential_steps,
267 .init = .{ .x = x0, .acc = acc0, .c = c },
268 .body = scanDifferentialStep,
269 });
270 return try builder.finish(&.{ walked.x, walked.acc });
271 }
272
273 fn scanDifferentialStep(builder: *Builder, carry: anytype) !@TypeOf(carry) {
274 const half = try builder.scalar(.f32, 0.5);
275 const next_x = try (try (try carry.x.mul(carry.x)).mul(half)).add(carry.c);
276 return .{
277 .x = next_x,
278 .acc = try carry.acc.add(next_x),
279 .c = carry.c,
280 };
281 }
282
283 fn buildScanDifferentialUnrolled(allocator: std.mem.Allocator) !Graph {
284 var builder = try Builder.init(allocator, "scan_differential");
285 errdefer builder.deinit();
286 const x0 = try builder.input(.f32, .{ .lane = scan_differential_width });
287 const c = try builder.input(.f32, .{ .lane = scan_differential_width });
288 const half = try builder.scalar(.f32, 0.5);
289 var x = x0;
290 var acc = try builder.full(.f32, .{ .lane = scan_differential_width }, 0.0);
291 for (0..scan_differential_steps) |_| {
292 x = try (try (try x.mul(x)).mul(half)).add(c);
293 acc = try acc.add(x);
294 }
295 return try builder.finish(&.{ x, acc });
296 }
297
298 fn buildScatterAddDuplicateProgram(allocator: std.mem.Allocator) !Graph {
299 var builder = try Builder.init(allocator, "scatter_add_duplicate");
300 errdefer builder.deinit();
301 const seed = try builder.input(.f32, .{ .vocab = 4, .channel = 2 });
302 const ids = try builder.input(.i32, .{ .token = 4 });
303 const updates = try builder.input(.f32, .{ .token = 4, .channel = 2 });
304 const out = try seed.scatterAdd(ids, updates, .vocab);
305 return try builder.finish(&.{out});
306 }
307
308 test "accy tensor scatter add accumulates duplicate indices on cpu" {
309 try fixture.requireNativeCpuArtifacts();
310 const allocator = std.testing.allocator;
311 var graph = try buildScatterAddDuplicateProgram(allocator);
312 defer graph.deinit();
313
314 const seed = [4][2]f32{
315 .{ 0.5, 0.5 },
316 .{ 1.0, 1.0 },
317 .{ 2.0, 2.0 },
318 .{ 3.0, 3.0 },
319 };
320 const ids = [_]i32{ 1, 2, 1, 3 };
321 const updates = [4][2]f32{
322 .{ 1.0, 10.0 },
323 .{ 2.0, 20.0 },
324 .{ 3.0, 30.0 },
325 .{ 4.0, 40.0 },
326 };
327 const expected = [4][2]f32{
328 .{ 0.5, 0.5 },
329 .{ 5.0, 41.0 },
330 .{ 4.0, 22.0 },
331 .{ 7.0, 43.0 },
332 };
333 var out: [4][2]f32 = @splat(@splat(0));
334 var outputs = [_][]u8{std.mem.sliceAsBytes(out[0..])};
335 try execute.runCpu(allocator, &graph, &.{
336 std.mem.sliceAsBytes(seed[0..]),
337 std.mem.sliceAsBytes(ids[0..]),
338 std.mem.sliceAsBytes(updates[0..]),
339 }, outputs[0..]);
340
341 for (0..4) |row| {
342 try std.testing.expectEqualSlices(f32, expected[row][0..], out[row][0..]);
343 }
344 }
345
346 test "accy tensor scheduled scatter add matches the expanded lowering on live CUDA" {
347 const accy = @import("../root.zig");
348 const gating = @import("accy_validation_gating");
349 const allocator = std.testing.allocator;
350
351 try gating.skipIfBuildFlagDisabled(.cuda);
352 if (!gpu.cuda.platformSupported()) return gating.skip(.cuda, .unsupported_platform);
353 try fixture.requireNativeCpuArtifacts();
354 var state = gpu.cuda.State.initDevice(allocator, 0) catch |err| switch (err) {
355 error.RuntimeUnavailable => return gating.skip(.cuda, .cuda_device_missing),
356 else => return err,
357 };
358 defer state.deinit();
359 const handle = state.handle();
360
361 var graph = try buildScatterAddDuplicateProgram(allocator);
362 defer graph.deinit();
363
364 const seed = [4][2]f32{
365 .{ 0.5, 0.5 },
366 .{ 1.0, 1.0 },
367 .{ 2.0, 2.0 },
368 .{ 3.0, 3.0 },
369 };
370 const ids = [_]i32{ 1, 2, 1, 3 };
371 const updates = [4][2]f32{
372 .{ 1.0, 10.0 },
373 .{ 2.0, 20.0 },
374 .{ 3.0, 30.0 },
375 .{ 4.0, 40.0 },
376 };
377 const inputs = [_][]const u8{
378 std.mem.sliceAsBytes(seed[0..]),
379 std.mem.sliceAsBytes(ids[0..]),
380 std.mem.sliceAsBytes(updates[0..]),
381 };
382
383 var expanded: [4][2]f32 = @splat(@splat(0));
384 var expanded_outputs = [_][]u8{std.mem.sliceAsBytes(expanded[0..])};
385 try execute.runCpu(allocator, &graph, inputs[0..], expanded_outputs[0..]);
386
387 var family_artifact = try accy.kernel.library.indexing.createScatterAddFamilyArtifact(allocator, handle, .{
388 .axis_size = 4,
389 .updates = 4,
390 .inner = 2,
391 .dtype = .f32,
392 .threads = 4,
393 }, .{ .limits = .testing, .format = .cuda_ptx });
394 defer family_artifact.deinit();
395 const registry = family_artifact.registry();
396
397 const options = lower.FragmentCompilerOptions{
398 .artifact_format = .cuda_ptx,
399 .kernel_call_registry = ®istry,
400 .scatter_add_schedule = .{ .thread_blocks = 4 },
401 };
402 const compiled = try lower.compileFragment(allocator, handle, &graph, options);
403 var scheduled_fragment = try accy.executable.loadFragment(allocator, handle, compiled, options);
404 defer scheduled_fragment.deinit();
405
406 var scheduled: [4][2]f32 = @splat(@splat(0));
407 var scheduled_outputs = [_][]u8{std.mem.sliceAsBytes(scheduled[0..])};
408 try accy.executable.invoke(scheduled_fragment, allocator, allocator, inputs[0..], scheduled_outputs[0..]);
409
410 for (0..4) |row| {
411 try std.testing.expectEqualSlices(f32, expanded[row][0..], scheduled[row][0..]);
412 }
413 }
414
415 fn vmapGatherBody(_: *Builder, args: []const Value) !Value {
416 return args[0].gather(args[1], .vocab);
417 }
418
419 fn vmapScatterAddBody(_: *Builder, args: []const Value) !Value {
420 return args[0].scatterAdd(args[1], args[2], .vocab);
421 }
422
423 test "accy tensor vmap shared batch gather matches the host reference on cpu" {
424 try fixture.requireNativeCpuArtifacts();
425 const allocator = std.testing.allocator;
426 var source = try define(allocator, "vmap_gather_shared_cpu", &.{
427 spec(.f32, .{ .vocab = 4, .channel = 2 }),
428 spec(.i32, .{ .token = 3 }),
429 }, vmapGatherBody);
430 defer source.deinit();
431
432 var batched = try vmap(allocator, &source, .{
433 .axis_size = 2,
434 .in_axes = &.{ mappedAxis(0), mappedAxis(0) },
435 });
436 defer batched.deinit();
437
438 var input: [2][4][2]f32 = undefined;
439 for (0..2) |b| {
440 for (0..4) |v| {
441 for (0..2) |c| {
442 input[b][v][c] = @floatFromInt(100 * b + 10 * v + c);
443 }
444 }
445 }
446 const ids = [2][3]i32{
447 .{ 0, 2, 3 },
448 .{ 3, 0, 3 },
449 };
450
451 var expected: [2][3][2]f32 = undefined;
452 for (0..2) |b| {
453 for (0..3) |t| {
454 for (0..2) |c| {
455 expected[b][t][c] = input[b][@intCast(ids[b][t])][c];
456 }
457 }
458 }
459
460 var out: [2][3][2]f32 = @splat(@splat(@splat(0)));
461 var outputs = [_][]u8{std.mem.sliceAsBytes(out[0..])};
462 try execute.runCpu(allocator, &batched, &.{
463 std.mem.sliceAsBytes(input[0..]),
464 std.mem.sliceAsBytes(ids[0..]),
465 }, outputs[0..]);
466
467 for (0..2) |b| {
468 for (0..3) |t| {
469 try std.testing.expectEqualSlices(f32, expected[b][t][0..], out[b][t][0..]);
470 }
471 }
472 }
473
474 test "accy tensor vmap all batched scatter add accumulates duplicate indices on cpu" {
475 try fixture.requireNativeCpuArtifacts();
476 const allocator = std.testing.allocator;
477 var source = try define(allocator, "vmap_scatter_add_all_batched_cpu", &.{
478 spec(.f32, .{ .vocab = 4, .channel = 2 }),
479 spec(.i32, .{ .token = 3 }),
480 spec(.f32, .{ .token = 3, .channel = 2 }),
481 }, vmapScatterAddBody);
482 defer source.deinit();
483
484 var batched = try vmap(allocator, &source, .{
485 .axis_size = 2,
486 .in_axes = &.{ mappedAxis(0), mappedAxis(0), mappedAxis(0) },
487 });
488 defer batched.deinit();
489
490 var seed: [2][4][2]f32 = undefined;
491 var updates: [2][3][2]f32 = undefined;
492 for (0..2) |b| {
493 for (0..4) |v| {
494 for (0..2) |c| {
495 seed[b][v][c] = @floatFromInt(100 * b + 10 * v + c);
496 }
497 }
498 for (0..3) |t| {
499 for (0..2) |c| {
500 updates[b][t][c] = @floatFromInt(1000 + 100 * b + 10 * t + c);
501 }
502 }
503 }
504 const ids = [2][3]i32{
505 .{ 1, 1, 3 },
506 .{ 0, 2, 2 },
507 };
508
509 var expected = seed;
510 for (0..2) |b| {
511 for (0..3) |t| {
512 for (0..2) |c| {
513 expected[b][@intCast(ids[b][t])][c] += updates[b][t][c];
514 }
515 }
516 }
517
518 var out: [2][4][2]f32 = @splat(@splat(@splat(0)));
519 var outputs = [_][]u8{std.mem.sliceAsBytes(out[0..])};
520 try execute.runCpu(allocator, &batched, &.{
521 std.mem.sliceAsBytes(seed[0..]),
522 std.mem.sliceAsBytes(ids[0..]),
523 std.mem.sliceAsBytes(updates[0..]),
524 }, outputs[0..]);
525
526 for (0..2) |b| {
527 for (0..4) |v| {
528 try std.testing.expectEqualSlices(f32, expected[b][v][0..], out[b][v][0..]);
529 }
530 }
531 }
532
533 test "accy tensor vmap batched index scatter add broadcasts the shared input on cpu" {
534 try fixture.requireNativeCpuArtifacts();
535 const allocator = std.testing.allocator;
536 var source = try define(allocator, "vmap_scatter_add_batched_indices_cpu", &.{
537 spec(.f32, .{ .vocab = 4, .channel = 2 }),
538 spec(.i32, .{ .token = 3 }),
539 spec(.f32, .{ .token = 3, .channel = 2 }),
540 }, vmapScatterAddBody);
541 defer source.deinit();
542
543 var batched = try vmap(allocator, &source, .{
544 .axis_size = 2,
545 .in_axes = &.{ .none, mappedAxis(0), mappedAxis(0) },
546 });
547 defer batched.deinit();
548
549 var input: [4][2]f32 = undefined;
550 for (0..4) |v| {
551 for (0..2) |c| {
552 input[v][c] = @floatFromInt(10 * v + c);
553 }
554 }
555 var updates: [2][3][2]f32 = undefined;
556 for (0..2) |b| {
557 for (0..3) |t| {
558 for (0..2) |c| {
559 updates[b][t][c] = @floatFromInt(1000 + 100 * b + 10 * t + c);
560 }
561 }
562 }
563 const ids = [2][3]i32{
564 .{ 0, 0, 2 },
565 .{ 1, 3, 1 },
566 };
567
568 var expected: [2][4][2]f32 = undefined;
569 for (0..2) |b| {
570 for (0..4) |v| {
571 expected[b][v] = input[v];
572 }
573 for (0..3) |t| {
574 for (0..2) |c| {
575 expected[b][@intCast(ids[b][t])][c] += updates[b][t][c];
576 }
577 }
578 }
579
580 var out: [2][4][2]f32 = @splat(@splat(@splat(0)));
581 var outputs = [_][]u8{std.mem.sliceAsBytes(out[0..])};
582 try execute.runCpu(allocator, &batched, &.{
583 std.mem.sliceAsBytes(input[0..]),
584 std.mem.sliceAsBytes(ids[0..]),
585 std.mem.sliceAsBytes(updates[0..]),
586 }, outputs[0..]);
587
588 for (0..2) |b| {
589 for (0..4) |v| {
590 try std.testing.expectEqualSlices(f32, expected[b][v][0..], out[b][v][0..]);
591 }
592 }
593 }
594
595 fn sparseCrossEntropyLossBody(_: *Builder, args: []const Value) !Value {
596 return args[0].sparseCrossEntropyLoss(args[1], .vocab);
597 }
598
599 test "accy tensor sparse cross entropy losses match the host reference on cpu" {
600 try fixture.requireNativeCpuArtifacts();
601 const accy = @import("../root.zig");
602 const allocator = std.testing.allocator;
603 var graph = try define(allocator, "sparse_cross_entropy_cpu", &.{
604 spec(.f32, .{ .sample = 4, .vocab = 5 }),
605 spec(.i32, .{ .sample = 4 }),
606 }, sparseCrossEntropyLossBody);
607 defer graph.deinit();
608
609 var logits = [_]f32{
610 0.5, -1.0, 2.0, 0.0, 1.5,
611 -0.25, 0.75, -2.0, 3.0, 0.125,
612 1.0, 1.0, 1.0, 1.0, 1.0,
613 -3.0, 4.0, 0.5, -0.5, 2.5,
614 };
615 var targets = [4]i32{ 2, 0, 4, 1 };
616 var expected: [4]f32 = undefined;
617 accy.kernel.library.loss.hostRowSparseCrossEntropy(4, 5, logits[0..], targets[0..], expected[0..]);
618
619 var losses = @as([4]f32, @splat(0));
620 var outputs = [_][]u8{std.mem.sliceAsBytes(losses[0..])};
621 try execute.runCpu(allocator, &graph, &.{
622 std.mem.sliceAsBytes(logits[0..]),
623 std.mem.sliceAsBytes(targets[0..]),
624 }, outputs[0..]);
625
626 for (expected, losses) |want, got| {
627 try std.testing.expectApproxEqAbs(want, got, 0.0001);
628 }
629 }
630
631 fn sparseCrossEntropyMeanBody(_: *Builder, args: []const Value) !Value {
632 return sparseCrossEntropyMean(args[0], args[1], .vocab);
633 }
634
635 test "accy tensor sparse cross entropy gradient matches softmax minus one hot on cpu" {
636 try fixture.requireNativeCpuArtifacts();
637 const allocator = std.testing.allocator;
638 var loss_graph = try define(allocator, "sparse_cross_entropy_grad_cpu", &.{
639 spec(.f32, .{ .sample = 4, .vocab = 5 }),
640 spec(.i32, .{ .sample = 4 }),
641 }, sparseCrossEntropyMeanBody);
642 defer loss_graph.deinit();
643
644 var gradient_graph = try grad(allocator, &loss_graph, .{ .wrt = &.{0} });
645 defer gradient_graph.deinit();
646
647 var logits = [4][5]f32{
648 .{ 0.5, -1.0, 2.0, 0.0, 1.5 },
649 .{ -0.25, 0.75, -2.0, 3.0, 0.125 },
650 .{ 1.0, 1.0, 1.0, 1.0, 1.0 },
651 .{ -3.0, 4.0, 0.5, -0.5, 2.5 },
652 };
653 var targets = [4]i32{ 2, 0, 4, 1 };
654
655 var expected: [4][5]f32 = undefined;
656 for (0..4) |row| {
657 var row_max = logits[row][0];
658 for (logits[row]) |value| row_max = @max(row_max, value);
659 var denom: f32 = 0;
660 for (logits[row]) |value| denom += @exp(value - row_max);
661 for (0..5) |class| {
662 const softmax = @exp(logits[row][class] - row_max) / denom;
663 const one_hot: f32 = if (targets[row] == class) 1.0 else 0.0;
664 expected[row][class] = (softmax - one_hot) / 4.0;
665 }
666 }
667
668 var logits_grad: [4][5]f32 = @splat(@splat(0));
669 var outputs = [_][]u8{std.mem.sliceAsBytes(logits_grad[0..])};
670 try execute.runCpu(allocator, &gradient_graph, &.{
671 std.mem.sliceAsBytes(logits[0..]),
672 std.mem.sliceAsBytes(targets[0..]),
673 }, outputs[0..]);
674
675 for (0..4) |row| {
676 for (0..5) |class| {
677 try std.testing.expectApproxEqAbs(expected[row][class], logits_grad[row][class], 0.0001);
678 }
679 }
680 }
681
682 test "accy tensor vmap sparse cross entropy matches the host reference on cpu" {
683 try fixture.requireNativeCpuArtifacts();
684 const accy = @import("../root.zig");
685 const allocator = std.testing.allocator;
686 var source = try define(allocator, "vmap_sparse_cross_entropy_cpu", &.{
687 spec(.f32, .{ .sample = 2, .vocab = 5 }),
688 spec(.i32, .{ .sample = 2 }),
689 }, sparseCrossEntropyLossBody);
690 defer source.deinit();
691
692 var batched = try vmap(allocator, &source, .{
693 .axis_size = 2,
694 .in_axes = &.{ mappedAxis(0), mappedAxis(0) },
695 });
696 defer batched.deinit();
697
698 var logits = [_]f32{
699 0.5, -1.0, 2.0, 0.0, 1.5,
700 -0.25, 0.75, -2.0, 3.0, 0.125,
701 1.0, 1.0, 1.0, 1.0, 1.0,
702 -3.0, 4.0, 0.5, -0.5, 2.5,
703 };
704 var targets = [4]i32{ 2, 0, 4, 1 };
705 var expected: [4]f32 = undefined;
706 accy.kernel.library.loss.hostRowSparseCrossEntropy(4, 5, logits[0..], targets[0..], expected[0..]);
707
708 var losses = @as([4]f32, @splat(0));
709 var outputs = [_][]u8{std.mem.sliceAsBytes(losses[0..])};
710 try execute.runCpu(allocator, &batched, &.{
711 std.mem.sliceAsBytes(logits[0..]),
712 std.mem.sliceAsBytes(targets[0..]),
713 }, outputs[0..]);
714
715 for (expected, losses) |want, got| {
716 try std.testing.expectApproxEqAbs(want, got, 0.0001);
717 }
718 }
719
720 test "accy tensor scheduled sparse cross entropy matches the expanded lowering on live CUDA" {
721 const accy = @import("../root.zig");
722 const gating = @import("accy_validation_gating");
723 const allocator = std.testing.allocator;
724
725 try gating.skipIfBuildFlagDisabled(.cuda);
726 if (!gpu.cuda.platformSupported()) return gating.skip(.cuda, .unsupported_platform);
727 try fixture.requireNativeCpuArtifacts();
728 var state = gpu.cuda.State.initDevice(allocator, 0) catch |err| switch (err) {
729 error.RuntimeUnavailable => return gating.skip(.cuda, .cuda_device_missing),
730 else => return err,
731 };
732 defer state.deinit();
733 const handle = state.handle();
734
735 var graph = try define(allocator, "sparse_cross_entropy_cuda", &.{
736 spec(.f32, .{ .sample = 4, .vocab = 5 }),
737 spec(.i32, .{ .sample = 4 }),
738 }, sparseCrossEntropyLossBody);
739 defer graph.deinit();
740
741 var logits = [_]f32{
742 0.5, -1.0, 2.0, 0.0, 1.5,
743 -0.25, 0.75, -2.0, 3.0, 0.125,
744 1.0, 1.0, 1.0, 1.0, 1.0,
745 -3.0, 4.0, 0.5, -0.5, 2.5,
746 };
747 var targets = [4]i32{ 2, 0, 4, 1 };
748 const inputs = [_][]const u8{
749 std.mem.sliceAsBytes(logits[0..]),
750 std.mem.sliceAsBytes(targets[0..]),
751 };
752
753 var expanded = @as([4]f32, @splat(0));
754 var expanded_outputs = [_][]u8{std.mem.sliceAsBytes(expanded[0..])};
755 try execute.runCpu(allocator, &graph, inputs[0..], expanded_outputs[0..]);
756
757 var family_artifact = try accy.kernel.library.loss.createRowSparseCrossEntropyFamilyArtifact(allocator, handle, .{
758 .rows = 4,
759 .classes = 5,
760 .threads = 4,
761 }, .{ .limits = .testing, .format = .cuda_ptx });
762 defer family_artifact.deinit();
763 const registry = family_artifact.registry();
764
765 const options = lower.FragmentCompilerOptions{
766 .artifact_format = .cuda_ptx,
767 .kernel_call_registry = ®istry,
768 .row_sparse_cross_entropy_schedule = .{ .thread_blocks = 4 },
769 };
770 const compiled = try lower.compileFragment(allocator, handle, &graph, options);
771 var scheduled_fragment = try accy.executable.loadFragment(allocator, handle, compiled, options);
772 defer scheduled_fragment.deinit();
773
774 var scheduled = @as([4]f32, @splat(0));
775 var scheduled_outputs = [_][]u8{std.mem.sliceAsBytes(scheduled[0..])};
776 try accy.executable.invoke(scheduled_fragment, allocator, allocator, inputs[0..], scheduled_outputs[0..]);
777
778 for (expanded, scheduled) |want, got| {
779 try std.testing.expectApproxEqAbs(want, got, 0.0001);
780 }
781 }
782
783 fn runScanDifferentialProgram(
784 allocator: std.mem.Allocator,
785 graph: *Graph,
786 xs: []const f32,
787 cs: []const f32,
788 final_x: []f32,
789 final_acc: []f32,
790 ) !void {
791 var outputs = [_][]u8{ std.mem.sliceAsBytes(final_x), std.mem.sliceAsBytes(final_acc) };
792 try execute.runCpu(
793 allocator,
794 graph,
795 &.{ std.mem.sliceAsBytes(xs), std.mem.sliceAsBytes(cs) },
796 outputs[0..],
797 );
798 }
799
800 test "accy tensor scan lowers to the same numbers as its hand unrolled twin" {
801 try fixture.requireNativeCpuArtifacts();
802 const allocator = std.testing.allocator;
803 const xs = [_]f32{ 0.25, -0.5, 1.0, 0.125 };
804 const cs = [_]f32{ 0.1, 0.2, -0.3, 0.05 };
805
806 var scan_graph = try buildScanDifferentialScan(allocator);
807 defer scan_graph.deinit();
808 var scan_x = @as([scan_differential_width]f32, @splat(0));
809 var scan_acc = @as([scan_differential_width]f32, @splat(0));
810 try runScanDifferentialProgram(allocator, &scan_graph, xs[0..], cs[0..], scan_x[0..], scan_acc[0..]);
811
812 var unrolled_graph = try buildScanDifferentialUnrolled(allocator);
813 defer unrolled_graph.deinit();
814 var unrolled_x = @as([scan_differential_width]f32, @splat(0));
815 var unrolled_acc = @as([scan_differential_width]f32, @splat(0));
816 try runScanDifferentialProgram(allocator, &unrolled_graph, xs[0..], cs[0..], unrolled_x[0..], unrolled_acc[0..]);
817
818 var host_x: [scan_differential_width]f32 = xs;
819 var host_acc = @as([scan_differential_width]f32, @splat(0));
820 for (0..scan_differential_steps) |_| {
821 for (0..scan_differential_width) |lane| {
822 host_x[lane] = host_x[lane] * host_x[lane] * 0.5 + cs[lane];
823 host_acc[lane] += host_x[lane];
824 }
825 }
826
827 try std.testing.expectEqualSlices(f32, unrolled_x[0..], scan_x[0..]);
828 try std.testing.expectEqualSlices(f32, unrolled_acc[0..], scan_acc[0..]);
829 for (0..scan_differential_width) |lane| {
830 try std.testing.expectApproxEqAbs(host_x[lane], scan_x[lane], 1e-6);
831 try std.testing.expectApproxEqAbs(host_acc[lane], scan_acc[lane], 1e-6);
832 }
833 }
834
835 fn zeroLengthScanStep(_: *Builder, carry: Value) !Value {
836 return carry.add(carry);
837 }
838
839 test "accy tensor scan with zero length yields its carry inits" {
840 try fixture.requireNativeCpuArtifacts();
841 const allocator = std.testing.allocator;
842 var builder = try Builder.init(allocator, "scan_zero_length");
843 defer builder.deinit();
844 const x0 = try builder.input(.f32, .{ .lane = scan_differential_width });
845 const walked = try builder.scan(.{
846 .length = 0,
847 .init = x0,
848 .body = zeroLengthScanStep,
849 });
850 const one = try builder.scalar(.f32, 1.0);
851 const bumped = try walked.add(one);
852 var graph = try builder.finish(&.{bumped});
853 defer graph.deinit();
854
855 const xs = [_]f32{ 0.25, -0.5, 1.0, 0.125 };
856 var out = @as([scan_differential_width]f32, @splat(0));
857 var outputs = [_][]u8{std.mem.sliceAsBytes(out[0..])};
858 try execute.runCpu(allocator, &graph, &.{std.mem.sliceAsBytes(xs[0..])}, outputs[0..]);
859
860 for (0..scan_differential_width) |lane| {
861 try std.testing.expectApproxEqAbs(xs[lane] + 1.0, out[lane], 0.0);
862 }
863 }
864
865 fn buildScanLossScan(allocator: std.mem.Allocator) !Graph {
866 var builder = try Builder.init(allocator, "scan_loss");
867 errdefer builder.deinit();
868 const x0 = try builder.input(.f32, .{ .lane = scan_differential_width });
869 const c = try builder.input(.f32, .{ .lane = scan_differential_width });
870 const acc0 = try builder.full(.f32, .{ .lane = scan_differential_width }, 0.0);
871 const walked = try builder.scan(.{
872 .length = scan_differential_steps,
873 .init = .{ .x = x0, .acc = acc0, .c = c },
874 .body = scanDifferentialStep,
875 });
876 const loss = try walked.acc.sum(.lane);
877 return try builder.finish(&.{loss});
878 }
879
880 fn buildScanLossUnrolled(allocator: std.mem.Allocator) !Graph {
881 var builder = try Builder.init(allocator, "scan_loss");
882 errdefer builder.deinit();
883 const x0 = try builder.input(.f32, .{ .lane = scan_differential_width });
884 const c = try builder.input(.f32, .{ .lane = scan_differential_width });
885 const half = try builder.scalar(.f32, 0.5);
886 var x = x0;
887 var acc = try builder.full(.f32, .{ .lane = scan_differential_width }, 0.0);
888 for (0..scan_differential_steps) |_| {
889 x = try (try (try x.mul(x)).mul(half)).add(c);
890 acc = try acc.add(x);
891 }
892 const loss = try acc.sum(.lane);
893 return try builder.finish(&.{loss});
894 }
895
896 test "accy tensor grad differentiates scan like its hand unrolled twin" {
897 try fixture.requireNativeCpuArtifacts();
898 const allocator = std.testing.allocator;
899 const xs = [_]f32{ 0.25, -0.5, 1.0, 0.125 };
900 const cs = [_]f32{ 0.1, 0.2, -0.3, 0.05 };
901
902 var scan_loss = try buildScanLossScan(allocator);
903 defer scan_loss.deinit();
904 var scan_grad = try grad(allocator, &scan_loss, .{ .wrt = &.{ 0, 1 } });
905 defer scan_grad.deinit();
906 try std.testing.expect(!scan_grad.containsScan());
907 var scan_dx = @as([scan_differential_width]f32, @splat(0));
908 var scan_dc = @as([scan_differential_width]f32, @splat(0));
909 var scan_outputs = [_][]u8{ std.mem.sliceAsBytes(scan_dx[0..]), std.mem.sliceAsBytes(scan_dc[0..]) };
910 try execute.runCpu(allocator, &scan_grad, &.{ std.mem.sliceAsBytes(xs[0..]), std.mem.sliceAsBytes(cs[0..]) }, scan_outputs[0..]);
911
912 var unrolled_loss = try buildScanLossUnrolled(allocator);
913 defer unrolled_loss.deinit();
914 var unrolled_grad = try grad(allocator, &unrolled_loss, .{ .wrt = &.{ 0, 1 } });
915 defer unrolled_grad.deinit();
916 var unrolled_dx = @as([scan_differential_width]f32, @splat(0));
917 var unrolled_dc = @as([scan_differential_width]f32, @splat(0));
918 var unrolled_outputs = [_][]u8{ std.mem.sliceAsBytes(unrolled_dx[0..]), std.mem.sliceAsBytes(unrolled_dc[0..]) };
919 try execute.runCpu(allocator, &unrolled_grad, &.{ std.mem.sliceAsBytes(xs[0..]), std.mem.sliceAsBytes(cs[0..]) }, unrolled_outputs[0..]);
920
921 try std.testing.expectEqualSlices(f32, unrolled_dx[0..], scan_dx[0..]);
922 try std.testing.expectEqualSlices(f32, unrolled_dc[0..], scan_dc[0..]);
923 }
924
925 test "accy tensor jvp linearizes scan like its hand unrolled twin" {
926 try fixture.requireNativeCpuArtifacts();
927 const allocator = std.testing.allocator;
928 const xs = [_]f32{ 0.25, -0.5, 1.0, 0.125 };
929 const cs = [_]f32{ 0.1, 0.2, -0.3, 0.05 };
930 const dxs = [_]f32{ 1.0, 0.5, -0.25, 2.0 };
931 const dcs = [_]f32{ 0.0, 1.0, 0.5, -1.0 };
932 const inputs = [_][]const u8{
933 std.mem.sliceAsBytes(xs[0..]),
934 std.mem.sliceAsBytes(cs[0..]),
935 std.mem.sliceAsBytes(dxs[0..]),
936 std.mem.sliceAsBytes(dcs[0..]),
937 };
938
939 var scan_graph = try buildScanDifferentialScan(allocator);
940 defer scan_graph.deinit();
941 var scan_linear = try linearize(allocator, &scan_graph, .{ .wrt = &.{ 0, 1 } });
942 defer scan_linear.deinit();
943 try std.testing.expect(!scan_linear.program.containsScan());
944 var scan_out: [4][scan_differential_width]f32 = @splat(@splat(0));
945 var scan_outputs = [_][]u8{
946 std.mem.sliceAsBytes(scan_out[0][0..]),
947 std.mem.sliceAsBytes(scan_out[1][0..]),
948 std.mem.sliceAsBytes(scan_out[2][0..]),
949 std.mem.sliceAsBytes(scan_out[3][0..]),
950 };
951 try execute.runCpu(allocator, &scan_linear.program, inputs[0..], scan_outputs[0..]);
952
953 var unrolled_graph = try buildScanDifferentialUnrolled(allocator);
954 defer unrolled_graph.deinit();
955 var unrolled_linear = try linearize(allocator, &unrolled_graph, .{ .wrt = &.{ 0, 1 } });
956 defer unrolled_linear.deinit();
957 var unrolled_out: [4][scan_differential_width]f32 = @splat(@splat(0));
958 var unrolled_outputs = [_][]u8{
959 std.mem.sliceAsBytes(unrolled_out[0][0..]),
960 std.mem.sliceAsBytes(unrolled_out[1][0..]),
961 std.mem.sliceAsBytes(unrolled_out[2][0..]),
962 std.mem.sliceAsBytes(unrolled_out[3][0..]),
963 };
964 try execute.runCpu(allocator, &unrolled_linear.program, inputs[0..], unrolled_outputs[0..]);
965
966 for (0..4) |output_index| {
967 try std.testing.expectEqualSlices(f32, unrolled_out[output_index][0..], scan_out[output_index][0..]);
968 }
969 }
970
971 test "accy tensor vmap batches scan structurally" {
972 try fixture.requireNativeCpuArtifacts();
973 const allocator = std.testing.allocator;
974 const batch_size = 3;
975 const lane_count = batch_size * scan_differential_width;
976 var xs: [lane_count]f32 = undefined;
977 var cs: [lane_count]f32 = undefined;
978 for (0..lane_count) |index| {
979 xs[index] = 0.05 * @as(f32, @floatFromInt(index + 1));
980 cs[index] = 0.02 * @as(f32, @floatFromInt(index + 1)) - 0.1;
981 }
982 const inputs = [_][]const u8{ std.mem.sliceAsBytes(xs[0..]), std.mem.sliceAsBytes(cs[0..]) };
983
984 var scan_graph = try buildScanDifferentialScan(allocator);
985 defer scan_graph.deinit();
986 var scan_batched = try vmap(allocator, &scan_graph, .{
987 .axis_size = batch_size,
988 .in_axes = &.{ mappedAxis(0), mappedAxis(0) },
989 });
990 defer scan_batched.deinit();
991 try std.testing.expect(scan_batched.containsScan());
992 var scan_x = @as([lane_count]f32, @splat(0));
993 var scan_acc = @as([lane_count]f32, @splat(0));
994 var scan_outputs = [_][]u8{ std.mem.sliceAsBytes(scan_x[0..]), std.mem.sliceAsBytes(scan_acc[0..]) };
995 try execute.runCpu(allocator, &scan_batched, inputs[0..], scan_outputs[0..]);
996
997 var unrolled_graph = try buildScanDifferentialUnrolled(allocator);
998 defer unrolled_graph.deinit();
999 var unrolled_batched = try vmap(allocator, &unrolled_graph, .{
1000 .axis_size = batch_size,
1001 .in_axes = &.{ mappedAxis(0), mappedAxis(0) },
1002 });
1003 defer unrolled_batched.deinit();
1004 var unrolled_x = @as([lane_count]f32, @splat(0));
1005 var unrolled_acc = @as([lane_count]f32, @splat(0));
1006 var unrolled_outputs = [_][]u8{ std.mem.sliceAsBytes(unrolled_x[0..]), std.mem.sliceAsBytes(unrolled_acc[0..]) };
1007 try execute.runCpu(allocator, &unrolled_batched, inputs[0..], unrolled_outputs[0..]);
1008
1009 try std.testing.expectEqualSlices(f32, unrolled_x[0..], scan_x[0..]);
1010 try std.testing.expectEqualSlices(f32, unrolled_acc[0..], scan_acc[0..]);
1011 }
1012
1013 fn denseGradientBody(_: *Builder, args: []const Value) !Value {
1014 const product = try args[0].contract(args[1], .k);
1015 return try product.sum(.{ .m, .n });
1016 }
1017
1018 test "accy tensor grad lowers dense scalar loss" {
1019 var traced = try define(std.testing.allocator, "grad_dense", &.{
1020 spec(.f32, .{ .m = 2, .k = 4 }),
1021 spec(.f32, .{ .k = 4, .n = 3 }),
1022 }, denseGradientBody);
1023 defer traced.deinit();
1024
1025 var differentiated = try grad(std.testing.allocator, &traced, .{ .wrt = &.{ 0, 1 } });
1026 defer differentiated.deinit();
1027
1028 const module = try toSemanticModule(std.testing.allocator, &differentiated);
1029 defer module.deinit();
1030
1031 try std.testing.expectEqual(@as(usize, 2), differentiated.parameters.len);
1032 try std.testing.expectEqual(@as(usize, 2), differentiated.outputs.len);
1033 try types.expectExtents(&.{ 2, 4 }, differentiated.typeOf(differentiated.outputs[0]));
1034 try types.expectExtents(&.{ 4, 3 }, differentiated.typeOf(differentiated.outputs[1]));
1035 try module.verify();
1036 }
1037
1038 fn batchedDenseGradientBody(_: *Builder, args: []const Value) !Value {
1039 const product = try args[0].contract(args[1], .k);
1040 return try product.sum(.{ .b, .m, .n });
1041 }
1042
1043 test "accy tensor grad lowers batched dense scalar loss" {
1044 var traced = try define(std.testing.allocator, "grad_batched_dense", &.{
1045 spec(.f32, .{ .b = 5, .m = 2, .k = 4 }),
1046 spec(.f32, .{ .b = 5, .k = 4, .n = 3 }),
1047 }, batchedDenseGradientBody);
1048 defer traced.deinit();
1049
1050 var differentiated = try grad(std.testing.allocator, &traced, .{ .wrt = &.{ 0, 1 } });
1051 defer differentiated.deinit();
1052
1053 const module = try toSemanticModule(std.testing.allocator, &differentiated);
1054 defer module.deinit();
1055
1056 try std.testing.expectEqual(@as(usize, 2), differentiated.parameters.len);
1057 try std.testing.expectEqual(@as(usize, 2), differentiated.outputs.len);
1058 try types.expectExtents(&.{ 5, 2, 4 }, differentiated.typeOf(differentiated.outputs[0]));
1059 try types.expectExtents(&.{ 5, 4, 3 }, differentiated.typeOf(differentiated.outputs[1]));
1060 try module.verify();
1061 }
1062
1063 test "accy tensor vmap composes with dense grad and lowering" {
1064 var traced = try define(std.testing.allocator, "vmap_grad_dense", &.{
1065 spec(.f32, .{ .m = 2, .k = 4 }),
1066 spec(.f32, .{ .k = 4, .n = 3 }),
1067 }, denseGradientBody);
1068 defer traced.deinit();
1069
1070 var differentiated = try grad(std.testing.allocator, &traced, .{ .wrt = &.{ 0, 1 } });
1071 defer differentiated.deinit();
1072
1073 var batched = try vmap(std.testing.allocator, &differentiated, .{
1074 .axis_size = 8,
1075 .in_axes = &.{ mappedAxis(0), .none },
1076 });
1077 defer batched.deinit();
1078
1079 const module = try toSemanticModule(std.testing.allocator, &batched);
1080 defer module.deinit();
1081
1082 try std.testing.expectEqual(@as(usize, 2), batched.parameters.len);
1083 try std.testing.expectEqual(@as(usize, 2), batched.outputs.len);
1084 try types.expectExtents(&.{ 8, 2, 4 }, batched.typeOf(batched.outputs[0]));
1085 try types.expectExtents(&.{ 8, 4, 3 }, batched.typeOf(batched.outputs[1]));
1086 try module.verify();
1087 }
1088
1089 const lm_batch = 2;
1090 const lm_token = 3;
1091 const lm_vocab = 11;
1092 const lm_channel = 4;
1093
1094 fn tinyLanguageModelLoss(_: *Builder, args: []const Value) !Value {
1095 const token_ids = args[0];
1096 const target_ids = args[1];
1097 const embedding_table = args[2];
1098 const projection = args[3];
1099 const hidden = try nn.embedding(embedding_table, token_ids, .vocab);
1100 const flat_hidden = try hidden.merge(.{ .batch, .token }, .sample);
1101 const flat_targets = try target_ids.merge(.{ .batch, .token }, .sample);
1102 const logits = try flat_hidden.contract(projection, .channel);
1103 return nn.sparseCrossEntropyMean(logits, flat_targets, .vocab);
1104 }
1105
1106 test "accy tensor builds and differentiates a tiny language model loss" {
1107 var traced = try define(std.testing.allocator, "tiny_lm_loss", &.{
1108 spec(.i32, .{ .batch = lm_batch, .token = lm_token }),
1109 spec(.i32, .{ .batch = lm_batch, .token = lm_token }),
1110 spec(.f32, .{ .vocab = lm_vocab, .channel = lm_channel }),
1111 spec(.f32, .{ .channel = lm_channel, .vocab = lm_vocab }),
1112 }, tinyLanguageModelLoss);
1113 defer traced.deinit();
1114
1115 try types.expectExtents(&.{}, traced.typeOf(traced.outputs[0]));
1116
1117 const lowered = try toSemanticModule(std.testing.allocator, &traced);
1118 defer lowered.deinit();
1119 try lowered.verify();
1120
1121 var differentiated = try grad(std.testing.allocator, &traced, .{ .wrt = &.{ 2, 3 } });
1122 defer differentiated.deinit();
1123
1124 const differentiated_lowered = try toSemanticModule(std.testing.allocator, &differentiated);
1125 defer differentiated_lowered.deinit();
1126
1127 try std.testing.expectEqual(@as(usize, 4), differentiated.parameters.len);
1128 try std.testing.expectEqual(@as(usize, 2), differentiated.outputs.len);
1129 try types.expectExtents(&.{ lm_vocab, lm_channel }, differentiated.typeOf(differentiated.outputs[0]));
1130 try types.expectExtents(&.{ lm_channel, lm_vocab }, differentiated.typeOf(differentiated.outputs[1]));
1131 try differentiated_lowered.verify();
1132 }
1133
1134 const attention_pos = 3;
1135 const attention_ctx = 4;
1136 const attention_head = 2;
1137 const attention_val = 2;
1138
1139 fn buildNamedAttention(allocator: std.mem.Allocator) !Graph {
1140 var builder = try Builder.init(allocator, "named_attention");
1141 errdefer builder.deinit();
1142 const q = try builder.input(.f32, .{ .pos = attention_pos, .head = attention_head });
1143 const k = try builder.input(.f32, .{ .ctx = attention_ctx, .head = attention_head });
1144 const v = try builder.input(.f32, .{ .ctx = attention_ctx, .val = attention_val });
1145
1146 const scores = try q.contract(k, .head);
1147 const stable = try scores.sub(try scores.max(.ctx));
1148 const weights = try stable.exp();
1149 const probs = try weights.div(try weights.sum(.ctx));
1150 const out = try probs.contract(v, .ctx);
1151 return try builder.finish(&.{out});
1152 }
1153
1154 fn hostNamedAttention(
1155 q: *const [attention_pos][attention_head]f32,
1156 k: *const [attention_ctx][attention_head]f32,
1157 v: *const [attention_ctx][attention_val]f32,
1158 out: *[attention_pos][attention_val]f32,
1159 ) void {
1160 for (0..attention_pos) |i| {
1161 var scores: [attention_ctx]f32 = undefined;
1162 var row_max: f32 = -std.math.inf(f32);
1163 for (0..attention_ctx) |j| {
1164 var dot: f32 = 0;
1165 for (0..attention_head) |h| dot += q[i][h] * k[j][h];
1166 scores[j] = dot;
1167 row_max = @max(row_max, dot);
1168 }
1169 var denom: f32 = 0;
1170 for (0..attention_ctx) |j| {
1171 scores[j] = @exp(scores[j] - row_max);
1172 denom += scores[j];
1173 }
1174 for (0..attention_val) |x| {
1175 var acc: f32 = 0;
1176 for (0..attention_ctx) |j| acc += (scores[j] / denom) * v[j][x];
1177 out[i][x] = acc;
1178 }
1179 }
1180 }
1181
1182 fn runProgramCuda(
1183 allocator: std.mem.Allocator,
1184 graph: *Graph,
1185 inputs: []const []const u8,
1186 outputs: [][]u8,
1187 ) !void {
1188 const accy = @import("../root.zig");
1189 const gating = @import("accy_validation_gating");
1190 try gating.skipIfBuildFlagDisabled(.cuda);
1191 if (!gpu.cuda.platformSupported()) return gating.skip(.cuda, .unsupported_platform);
1192 var state = gpu.cuda.State.initDevice(allocator, 0) catch |err| switch (err) {
1193 error.RuntimeUnavailable => return gating.skip(.cuda, .cuda_device_missing),
1194 else => return err,
1195 };
1196 defer state.deinit();
1197 const module = try toSemanticModule(allocator, graph);
1198 const options = accy.executable.FragmentCompilerOptions{
1199 .artifact_format = .cuda_ptx,
1200 };
1201 const compiled = try accy.executable.compileFragmentFromSemanticModule(allocator, state.handle(), module, options);
1202 var fragment = try accy.executable.loadFragment(allocator, state.handle(), compiled, options);
1203 defer fragment.deinit();
1204 try accy.executable.invoke(fragment, allocator, allocator, inputs, outputs);
1205 }
1206
1207 test "accy tensor named attention verifies and matches the host softmax reference" {
1208 const allocator = std.testing.allocator;
1209
1210 var graph = try buildNamedAttention(allocator);
1211 defer graph.deinit();
1212 try types.expectExtents(&.{ attention_pos, attention_val }, graph.typeOf(graph.outputs[0]));
1213 try std.testing.expectEqualStrings("pos", graph.typeOf(graph.outputs[0]).dims[0].name);
1214 try std.testing.expectEqualStrings("val", graph.typeOf(graph.outputs[0]).dims[1].name);
1215
1216 const module = try toSemanticModule(allocator, &graph);
1217 defer module.deinit();
1218 try module.verify();
1219
1220 const q = [attention_pos][attention_head]f32{
1221 .{ 0.2, -0.4 },
1222 .{ 1.1, 0.3 },
1223 .{ -0.7, 0.9 },
1224 };
1225 const k = [attention_ctx][attention_head]f32{
1226 .{ 0.5, 0.1 },
1227 .{ -0.3, 0.8 },
1228 .{ 0.9, -0.6 },
1229 .{ 0.0, 0.4 },
1230 };
1231 const v = [attention_ctx][attention_val]f32{
1232 .{ 1.0, 0.0 },
1233 .{ 0.0, 1.0 },
1234 .{ 0.5, 0.5 },
1235 .{ -1.0, 2.0 },
1236 };
1237
1238 var expected: [attention_pos][attention_val]f32 = undefined;
1239 hostNamedAttention(&q, &k, &v, &expected);
1240
1241 var out: [attention_pos][attention_val]f32 = @splat(@splat(0));
1242 var outputs = [_][]u8{std.mem.sliceAsBytes(out[0..])};
1243 try runProgramCuda(allocator, &graph, &.{
1244 std.mem.sliceAsBytes(q[0..]),
1245 std.mem.sliceAsBytes(k[0..]),
1246 std.mem.sliceAsBytes(v[0..]),
1247 }, outputs[0..]);
1248
1249 for (0..attention_pos) |i| {
1250 for (0..attention_val) |x| {
1251 try std.testing.expectApproxEqAbs(expected[i][x], out[i][x], 1e-5);
1252 }
1253 }
1254 }
1255
1256 fn buildBatchedContract(allocator: std.mem.Allocator) !Graph {
1257 var builder = try Builder.init(allocator, "batched_contract");
1258 errdefer builder.deinit();
1259 const lhs = try builder.input(.f32, .{ .b = 2, .m = 2, .k = 3 });
1260 const rhs = try builder.input(.f32, .{ .b = 2, .k = 3, .n = 2 });
1261 const out = try lhs.contract(rhs, .k);
1262 return try builder.finish(&.{out});
1263 }
1264
1265 test "accy tensor contract batches shared axes like a per-slice matmul" {
1266 const allocator = std.testing.allocator;
1267
1268 var graph = try buildBatchedContract(allocator);
1269 defer graph.deinit();
1270 const out_ty = graph.typeOf(graph.outputs[0]);
1271 try types.expectExtents(&.{ 2, 2, 2 }, out_ty);
1272 try std.testing.expectEqualStrings("b", out_ty.dims[0].name);
1273 try std.testing.expectEqualStrings("m", out_ty.dims[1].name);
1274 try std.testing.expectEqualStrings("n", out_ty.dims[2].name);
1275
1276 const module = try toSemanticModule(allocator, &graph);
1277 defer module.deinit();
1278 try module.verify();
1279
1280 var lhs: [2][2][3]f32 = undefined;
1281 var rhs: [2][3][2]f32 = undefined;
1282 var seed: f32 = 0.1;
1283 for (0..2) |b| {
1284 for (0..2) |m| {
1285 for (0..3) |k| {
1286 lhs[b][m][k] = seed;
1287 seed += 0.07;
1288 }
1289 }
1290 for (0..3) |k| {
1291 for (0..2) |n| {
1292 rhs[b][k][n] = seed - 0.5;
1293 seed += 0.05;
1294 }
1295 }
1296 }
1297
1298 var expected: [2][2][2]f32 = @splat(@splat(@splat(0)));
1299 for (0..2) |b| {
1300 for (0..2) |m| {
1301 for (0..2) |n| {
1302 for (0..3) |k| {
1303 expected[b][m][n] += lhs[b][m][k] * rhs[b][k][n];
1304 }
1305 }
1306 }
1307 }
1308
1309 var out: [2][2][2]f32 = @splat(@splat(@splat(0)));
1310 var outputs = [_][]u8{std.mem.sliceAsBytes(out[0..])};
1311 try runProgramCuda(allocator, &graph, &.{
1312 std.mem.sliceAsBytes(lhs[0..]),
1313 std.mem.sliceAsBytes(rhs[0..]),
1314 }, outputs[0..]);
1315
1316 for (0..2) |b| {
1317 for (0..2) |m| {
1318 for (0..2) |n| {
1319 try std.testing.expectApproxEqAbs(expected[b][m][n], out[b][m][n], 1e-5);
1320 }
1321 }
1322 }
1323 }
1324
1325 test "accy tensor declaration coverage" {
1326 std.testing.refAllDecls(tensor);
1327 }