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 }