lib/accy/src/kernel/library/random/key/profile.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const choir_abi = @import("choir_abi");
2 const random = @import("../root.zig");
3 const types = @import("types.zig");
4
5 const base = random.base;
6 const block = random.block;
7
8 const std = base.std;
9 const artifact_product = base.artifact_product;
10 const shape = base.shape;
11 const entry = base.entry;
12 const runtimeExtentArgument = base.runtimeExtentArgument;
13 const randomRuntimeExtentBounds = base.randomRuntimeExtentBounds;
14 const randomShapeFamily = base.randomShapeFamily;
15 const randomDTypeSupported = block.randomDTypeSupported;
16
17 const PhiloxKeySplit = types.PhiloxKeySplit;
18 const PhiloxKeyUniform = types.PhiloxKeyUniform;
19 const PhiloxKeyCounterUniform = types.PhiloxKeyCounterUniform;
20
21 pub fn philoxKeySplitRuntimeArguments(instance: PhiloxKeySplit) ![1]choir_abi.ScalarArgument {
22 return .{.{ .u32 = try runtimeExtentArgument(instance.count) }};
23 }
24
25 pub fn philoxKeyUniformRuntimeArguments(instance: PhiloxKeyUniform) ![1]choir_abi.ScalarArgument {
26 return .{.{ .u32 = try runtimeExtentArgument(instance.count) }};
27 }
28
29 pub fn philoxKeyCounterUniformRuntimeArguments(instance: PhiloxKeyCounterUniform) ![1]choir_abi.ScalarArgument {
30 return .{.{ .u32 = try runtimeExtentArgument(instance.count) }};
31 }
32
33 pub fn philoxKeySplitShapeProfileDimensions(instance: PhiloxKeySplit) [1]artifact_product.KernelCallShapeProfileDimension {
34 return .{.{ .name = instance.count_axis, .runtime_scalar_argument_index = 0, .bounds = randomRuntimeExtentBounds() }};
35 }
36
37 pub fn philoxKeyUniformShapeProfileDimensions(instance: PhiloxKeyUniform) [1]artifact_product.KernelCallShapeProfileDimension {
38 return .{.{ .name = instance.count_axis, .runtime_scalar_argument_index = 0, .bounds = randomRuntimeExtentBounds() }};
39 }
40
41 pub fn philoxKeyCounterUniformShapeProfileDimensions(instance: PhiloxKeyCounterUniform) [1]artifact_product.KernelCallShapeProfileDimension {
42 return .{.{ .name = instance.count_axis, .runtime_scalar_argument_index = 0, .bounds = randomRuntimeExtentBounds() }};
43 }
44
45 pub fn philoxKeySplitFamilyFingerprint(backing_allocator: std.mem.Allocator, instance: PhiloxKeySplit) !u64 {
46 var family = try philoxKeySplitShapeFamily(backing_allocator, instance);
47 defer family.deinit();
48 return shape.fingerprint(family);
49 }
50
51 pub fn philoxKeyUniformFamilyFingerprint(backing_allocator: std.mem.Allocator, instance: PhiloxKeyUniform) !u64 {
52 var family = try philoxKeyUniformShapeFamily(backing_allocator, instance);
53 defer family.deinit();
54 return shape.fingerprint(family);
55 }
56
57 pub fn philoxKeyCounterUniformFamilyFingerprint(backing_allocator: std.mem.Allocator, instance: PhiloxKeyCounterUniform) !u64 {
58 var family = try philoxKeyCounterUniformShapeFamily(backing_allocator, instance);
59 defer family.deinit();
60 return shape.fingerprint(family);
61 }
62
63 pub fn philoxKeySplitShapeFamily(backing_allocator: std.mem.Allocator, instance: PhiloxKeySplit) !shape.Family {
64 return randomShapeFamily(backing_allocator, "philox_key_split", instance.count_axis);
65 }
66
67 pub fn philoxKeyUniformShapeFamily(backing_allocator: std.mem.Allocator, instance: PhiloxKeyUniform) !shape.Family {
68 return randomShapeFamily(backing_allocator, "philox_key_uniform", instance.count_axis);
69 }
70
71 pub fn philoxKeyCounterUniformShapeFamily(backing_allocator: std.mem.Allocator, instance: PhiloxKeyCounterUniform) !shape.Family {
72 return randomShapeFamily(backing_allocator, "philox_key_counter_uniform", instance.count_axis);
73 }
74
75 pub fn philoxKeySplitFamilySpecialization(backing_allocator: std.mem.Allocator, instance: PhiloxKeySplit) !entry.OwnedSpecialization {
76 var owned = entry.OwnedSpecialization.init(backing_allocator);
77 errdefer owned.deinit();
78 const lifetime_allocator = owned.allocator();
79
80 const inputs = try lifetime_allocator.alloc(entry.Shape, 1);
81 inputs[0] = entry.runtimeShapeScalar();
82 const outputs = try lifetime_allocator.alloc(entry.Shape, 1);
83 outputs[0] = try entry.runtimeShape1D(lifetime_allocator, instance.count_axis, instance.count);
84
85 owned.value = .{
86 .dtype = .key,
87 .operation = .{ .random = .{ .philox_key_split = instance.rounds } },
88 .inputs = inputs,
89 .outputs = outputs,
90 .schedule = try entry.runtimeThreadBlocks1D(lifetime_allocator, "e", instance.generators(), instance.threads),
91 };
92 owned.value.launch = owned.value.schedule.?.launch();
93 var family = try philoxKeySplitShapeFamily(backing_allocator, instance);
94 errdefer family.deinit();
95 try owned.takeShapeFamily(&family);
96 return owned;
97 }
98
99 pub fn philoxKeyUniformFamilySpecialization(backing_allocator: std.mem.Allocator, instance: PhiloxKeyUniform) !entry.OwnedSpecialization {
100 var owned = entry.OwnedSpecialization.init(backing_allocator);
101 errdefer owned.deinit();
102 const lifetime_allocator = owned.allocator();
103
104 const inputs = try lifetime_allocator.alloc(entry.Shape, 1);
105 inputs[0] = try entry.runtimeShape1D(lifetime_allocator, instance.count_axis, instance.count);
106 const outputs = try lifetime_allocator.alloc(entry.Shape, 1);
107 outputs[0] = try entry.runtimeShape1D(lifetime_allocator, instance.count_axis, instance.count);
108
109 owned.value = .{
110 .dtype = instance.dtype,
111 .operation = .{ .random = .{ .philox_key_uniform = instance.rounds } },
112 .inputs = inputs,
113 .outputs = outputs,
114 .schedule = try entry.runtimeThreadBlocks1D(lifetime_allocator, "e", instance.generators(), instance.threads),
115 };
116 owned.value.launch = owned.value.schedule.?.launch();
117 var family = try philoxKeyUniformShapeFamily(backing_allocator, instance);
118 errdefer family.deinit();
119 try owned.takeShapeFamily(&family);
120 return owned;
121 }
122
123 pub fn philoxKeyCounterUniformFamilySpecialization(backing_allocator: std.mem.Allocator, instance: PhiloxKeyCounterUniform) !entry.OwnedSpecialization {
124 var owned = entry.OwnedSpecialization.init(backing_allocator);
125 errdefer owned.deinit();
126 const lifetime_allocator = owned.allocator();
127
128 const inputs = try lifetime_allocator.alloc(entry.Shape, 2);
129 inputs[0] = entry.runtimeShapeScalar();
130 inputs[1] = entry.runtimeShapeScalar();
131 const outputs = try lifetime_allocator.alloc(entry.Shape, 1);
132 outputs[0] = try entry.runtimeShape1D(lifetime_allocator, instance.count_axis, instance.count);
133
134 owned.value = .{
135 .dtype = instance.dtype,
136 .operation = .{ .random = .{ .philox_key_counter_uniform = instance.rounds } },
137 .inputs = inputs,
138 .outputs = outputs,
139 .schedule = try entry.runtimeThreadBlocks1D(lifetime_allocator, "e", instance.generators(), instance.threads),
140 };
141 owned.value.launch = owned.value.schedule.?.launch();
142 var family = try philoxKeyCounterUniformShapeFamily(backing_allocator, instance);
143 errdefer family.deinit();
144 try owned.takeShapeFamily(&family);
145 return owned;
146 }
147
148 pub fn philoxKeySplitInstanceFromSpecialization(specialization: entry.Specialization) ?PhiloxKeySplit {
149 if (!specialization.scheduleMatchesLaunch()) return null;
150 const operation = specialization.operation orelse return null;
151 const rounds = switch (operation) {
152 .random => |random_operation| switch (random_operation) {
153 .philox_key_split => |rounds| rounds,
154 else => return null,
155 },
156 else => return null,
157 };
158 const dtype = specialization.dtype orelse return null;
159 if (dtype != .key) return null;
160 if (specialization.inputs.len != 1 or specialization.outputs.len != 1) return null;
161 if (specialization.inputs[0].axes.len != 0) return null;
162 if (specialization.reductions.len != 0) return null;
163 const output = specialization.outputs[0];
164 if (output.axes.len != 1) return null;
165 const launch = specialization.launch orelse return null;
166 if (launch.threadgroup[0] == 0) return null;
167 return .{
168 .count = output.axes[0].extent,
169 .rounds = rounds,
170 .threads = launch.threadgroup[0],
171 .count_axis = output.axes[0].name,
172 };
173 }
174
175 pub fn philoxKeyCounterUniformInstanceFromSpecialization(specialization: entry.Specialization) ?PhiloxKeyCounterUniform {
176 if (!specialization.scheduleMatchesLaunch()) return null;
177 const operation = specialization.operation orelse return null;
178 const rounds = switch (operation) {
179 .random => |random_operation| switch (random_operation) {
180 .philox_key_counter_uniform => |rounds| rounds,
181 else => return null,
182 },
183 else => return null,
184 };
185 const dtype = specialization.dtype orelse return null;
186 if (!randomDTypeSupported(dtype)) return null;
187 if (specialization.inputs.len != 2 or specialization.outputs.len != 1) return null;
188 if (specialization.reductions.len != 0) return null;
189 if (specialization.inputs[0].axes.len != 0 or specialization.inputs[1].axes.len != 0) return null;
190 const output = specialization.outputs[0];
191 if (output.axes.len != 1) return null;
192 const launch = specialization.launch orelse return null;
193 if (launch.threadgroup[0] == 0) return null;
194 return .{
195 .count = output.axes[0].extent,
196 .rounds = rounds,
197 .dtype = dtype,
198 .threads = launch.threadgroup[0],
199 .count_axis = output.axes[0].name,
200 };
201 }
202
203 pub fn philoxKeyUniformInstanceFromSpecialization(specialization: entry.Specialization) ?PhiloxKeyUniform {
204 if (!specialization.scheduleMatchesLaunch()) return null;
205 const operation = specialization.operation orelse return null;
206 const rounds = switch (operation) {
207 .random => |random_operation| switch (random_operation) {
208 .philox_key_uniform => |rounds| rounds,
209 else => return null,
210 },
211 else => return null,
212 };
213 const dtype = specialization.dtype orelse return null;
214 if (!randomDTypeSupported(dtype)) return null;
215 if (specialization.inputs.len != 1 or specialization.outputs.len != 1) return null;
216 if (specialization.reductions.len != 0) return null;
217 const input = specialization.inputs[0];
218 const output = specialization.outputs[0];
219 if (input.axes.len != 1 or output.axes.len != 1) return null;
220 if (input.axes[0].extent != output.axes[0].extent) return null;
221 const launch = specialization.launch orelse return null;
222 if (launch.threadgroup[0] == 0) return null;
223 return .{
224 .count = output.axes[0].extent,
225 .rounds = rounds,
226 .dtype = dtype,
227 .threads = launch.threadgroup[0],
228 .count_axis = output.axes[0].name,
229 };
230 }