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 }