lib/accy/src/tensor/random/test.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const choir_abi = @import("choir_abi");
3 const accy = @import("../../root.zig");
4 const tensor = @import("../root.zig");
5 const namespace = @import("root.zig");
6 const DType = tensor.DType;
7 const Builder = tensor.Builder;
8 const Algorithm = namespace.Algorithm;
9 const KeySpec = namespace.KeySpec;
10 const Key = namespace.Key;
11 const input = namespace.input;
12 const seed = namespace.seed;
13 const bind = namespace.bind;
14 const uniform = namespace.uniform;
15
16 test "tensor random seed creates scalar key constants" {
17 var builder = try Builder.init(std.testing.allocator, "random_seed");
18 defer builder.deinit();
19
20 const seeded = try seed(&builder, 0x01234567_89abcdef, .{});
21 try std.testing.expectEqual(DType.key, seeded.value.ty.dtype);
22 try std.testing.expectEqual(@as(usize, 0), seeded.value.ty.rank());
23 const payload = builder.operations.items[seeded.value.id.index].kind.constant.payload;
24 const keys = std.mem.bytesAsSlice(choir_abi.Key, @constCast(payload));
25 try std.testing.expectEqual(@as(u64, 0x01234567_89abcdef), keys[0].bits);
26 }
27
28 test "tensor random split names the sample axis and lowers to custom calls" {
29 var builder = try Builder.init(std.testing.allocator, "random_split_uniform");
30 defer builder.deinit();
31
32 const seeded = try seed(&builder, 7, .{ .threads = 64 });
33 const keys = try seeded.split(.{ .sample = 8 });
34 const out = try keys.uniform(.f32);
35
36 try std.testing.expectEqual(DType.key, keys.value.ty.dtype);
37 try tensor.types.expectExtents(&.{8}, keys.value.ty);
38 try std.testing.expectEqualStrings("sample", keys.value.ty.dims[0].name);
39 try std.testing.expectEqual(DType.f32, out.ty.dtype);
40 try std.testing.expectEqualStrings("sample", out.ty.dims[0].name);
41
42 const split_call = builder.operations.items[keys.value.id.index].kind.custom_call;
43 try std.testing.expectEqualStrings("accy.kernel.random.philox_key_split_family_10r_64_key", split_call.target);
44 try std.testing.expectEqual(accy.kernel.library.random.philox_key_split_family_version, split_call.version);
45
46 const uniform_call = builder.operations.items[out.id.index].kind.custom_call;
47 try std.testing.expectEqualStrings("accy.kernel.random.philox_key_uniform_family_10r_64_f32", uniform_call.target);
48 try std.testing.expectEqual(accy.kernel.library.random.philox_key_uniform_family_version, uniform_call.version);
49
50 var program = try builder.finish(&.{out});
51 defer program.deinit();
52
53 const lowered = try tensor.toSemanticModule(std.testing.allocator, &program);
54 defer lowered.deinit();
55
56 try lowered.verify();
57 }
58
59 test "tensor random counter uniform lowers to key and counter custom call" {
60 var builder = try Builder.init(std.testing.allocator, "random_counter_uniform");
61 defer builder.deinit();
62
63 const seeded = try seed(&builder, 7, .{ .threads = 64 });
64 const counter = try builder.scalar(.i32, 3);
65 const out = try seeded.counterUniform(counter, .{ .draw = 8 }, .f32);
66
67 try std.testing.expectEqual(DType.f32, out.ty.dtype);
68 try tensor.types.expectExtents(&.{8}, out.ty);
69 try std.testing.expectEqualStrings("draw", out.ty.dims[0].name);
70
71 const uniform_call = builder.operations.items[out.id.index].kind.custom_call;
72 try std.testing.expectEqualStrings("accy.kernel.random.philox_key_counter_uniform_family_10r_64_f32", uniform_call.target);
73 try std.testing.expectEqual(accy.kernel.library.random.philox_key_counter_uniform_family_version, uniform_call.version);
74 try std.testing.expectEqual(@as(usize, 2), uniform_call.operands.len);
75
76 var program = try builder.finish(&.{out});
77 defer program.deinit();
78
79 const lowered = try tensor.toSemanticModule(std.testing.allocator, &program);
80 defer lowered.deinit();
81
82 try lowered.verify();
83 }