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 = &registry,
 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 = &registry,
 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 }