lib/accy/src/kernel/library/random/key/artifact.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const random = @import("../root.zig");
2 const types = @import("types.zig");
3 const runtime = @import("runtime.zig");
4
5 const base = random.base;
6 const std = base.std;
7 const artifact_product = base.artifact_product;
8 const entry = base.entry;
9 const kernel = base.kernel;
10
11 const PhiloxKeySplit = types.PhiloxKeySplit;
12 const PhiloxKeyUniform = types.PhiloxKeyUniform;
13 const PhiloxKeyCounterUniform = types.PhiloxKeyCounterUniform;
14 const PhiloxKeySplitRuntimeFamily = runtime.PhiloxKeySplitRuntimeFamily;
15 const PhiloxKeyUniformRuntimeFamilyF32 = runtime.PhiloxKeyUniformRuntimeFamilyF32;
16 const PhiloxKeyUniformRuntimeFamilyI32 = runtime.PhiloxKeyUniformRuntimeFamilyI32;
17 const PhiloxKeyCounterUniformRuntimeFamilyF32 = runtime.PhiloxKeyCounterUniformRuntimeFamilyF32;
18 const PhiloxKeyCounterUniformRuntimeFamilyI32 = runtime.PhiloxKeyCounterUniformRuntimeFamilyI32;
19 const philox_key_split_family_version = types.philox_key_split_family_version;
20 const philox_key_uniform_family_version = types.philox_key_uniform_family_version;
21 const philox_key_counter_uniform_family_version = types.philox_key_counter_uniform_family_version;
22 const philoxKeySplitFamilyTarget = @import("names.zig").philoxKeySplitFamilyTarget;
23 const philoxKeySplitFamilyEntryName = @import("names.zig").philoxKeySplitFamilyEntryName;
24 const philoxKeyUniformFamilyTarget = @import("names.zig").philoxKeyUniformFamilyTarget;
25 const philoxKeyUniformFamilyEntryName = @import("names.zig").philoxKeyUniformFamilyEntryName;
26 const philoxKeyCounterUniformFamilyTarget = @import("names.zig").philoxKeyCounterUniformFamilyTarget;
27 const philoxKeyCounterUniformFamilyEntryName = @import("names.zig").philoxKeyCounterUniformFamilyEntryName;
28 const profile = @import("profile.zig");
29 const randomFoldDerivedLaunch = base.randomFoldDerivedLaunch;
30
31 pub fn createPhiloxKeySplitFamilyArtifact(
32 allocator: std.mem.Allocator,
33 handle: kernel.BackendHandle,
34 instance: PhiloxKeySplit,
35 options: entry.ArtifactOptions,
36 ) !kernel.OwnedKernelCallArtifact {
37 const target = try philoxKeySplitFamilyTarget(allocator, instance);
38 defer allocator.free(target);
39 const entry_name = try philoxKeySplitFamilyEntryName(allocator, instance);
40 defer allocator.free(entry_name);
41 const family_fingerprint = options.shape_family_fingerprint orelse try profile.philoxKeySplitFamilyFingerprint(allocator, instance);
42 const shape_profile_dimensions = profile.philoxKeySplitShapeProfileDimensions(instance);
43 const shape_profile = options.shape_profile orelse artifact_product.KernelCallShapeProfile{
44 .name = "philox_key_split",
45 .fingerprint = family_fingerprint,
46 .dimensions = shape_profile_dimensions[0..],
47 };
48
49 var graph = try PhiloxKeySplitRuntimeFamily.buildNamed(allocator, options.limits, entry_name, instance);
50 defer graph.deinit();
51 return kernel.createKernelCallArtifact(allocator, handle, &graph, .{
52 .target = target,
53 .version = philox_key_split_family_version,
54 .format = options.format,
55 .kernel_plan = options.kernel_plan,
56 .element_count_argument = options.element_count_argument,
57 .shape_family_fingerprint = family_fingerprint,
58 .shape_profile = shape_profile,
59 .launch = options.launch orelse try randomFoldDerivedLaunch(instance.threads),
60 .runtime_scalar_argument_count = if (options.runtime_scalar_argument_count == 0) 1 else options.runtime_scalar_argument_count,
61 .static_arguments = options.static_arguments,
62 });
63 }
64
65 pub fn createPhiloxKeyUniformFamilyArtifact(
66 allocator: std.mem.Allocator,
67 handle: kernel.BackendHandle,
68 instance: PhiloxKeyUniform,
69 options: entry.ArtifactOptions,
70 ) !kernel.OwnedKernelCallArtifact {
71 const target = try philoxKeyUniformFamilyTarget(allocator, instance);
72 defer allocator.free(target);
73 const entry_name = try philoxKeyUniformFamilyEntryName(allocator, instance);
74 defer allocator.free(entry_name);
75 const family_fingerprint = options.shape_family_fingerprint orelse try profile.philoxKeyUniformFamilyFingerprint(allocator, instance);
76 const shape_profile_dimensions = profile.philoxKeyUniformShapeProfileDimensions(instance);
77 const shape_profile = options.shape_profile orelse artifact_product.KernelCallShapeProfile{
78 .name = "philox_key_uniform",
79 .fingerprint = family_fingerprint,
80 .dimensions = shape_profile_dimensions[0..],
81 };
82
83 var graph = switch (instance.dtype) {
84 .f32 => try PhiloxKeyUniformRuntimeFamilyF32.buildNamed(allocator, options.limits, entry_name, instance),
85 .i32 => try PhiloxKeyUniformRuntimeFamilyI32.buildNamed(allocator, options.limits, entry_name, instance),
86 else => return error.UnsupportedDType,
87 };
88 defer graph.deinit();
89 return kernel.createKernelCallArtifact(allocator, handle, &graph, .{
90 .target = target,
91 .version = philox_key_uniform_family_version,
92 .format = options.format,
93 .kernel_plan = options.kernel_plan,
94 .element_count_argument = options.element_count_argument,
95 .shape_family_fingerprint = family_fingerprint,
96 .shape_profile = shape_profile,
97 .launch = options.launch orelse try randomFoldDerivedLaunch(instance.threads),
98 .runtime_scalar_argument_count = if (options.runtime_scalar_argument_count == 0) 1 else options.runtime_scalar_argument_count,
99 .static_arguments = options.static_arguments,
100 });
101 }
102
103 pub fn createPhiloxKeyCounterUniformFamilyArtifact(
104 allocator: std.mem.Allocator,
105 handle: kernel.BackendHandle,
106 instance: PhiloxKeyCounterUniform,
107 options: entry.ArtifactOptions,
108 ) !kernel.OwnedKernelCallArtifact {
109 const target = try philoxKeyCounterUniformFamilyTarget(allocator, instance);
110 defer allocator.free(target);
111 const entry_name = try philoxKeyCounterUniformFamilyEntryName(allocator, instance);
112 defer allocator.free(entry_name);
113 const family_fingerprint = options.shape_family_fingerprint orelse try profile.philoxKeyCounterUniformFamilyFingerprint(allocator, instance);
114 const shape_profile_dimensions = profile.philoxKeyCounterUniformShapeProfileDimensions(instance);
115 const shape_profile = options.shape_profile orelse artifact_product.KernelCallShapeProfile{
116 .name = "philox_key_counter_uniform",
117 .fingerprint = family_fingerprint,
118 .dimensions = shape_profile_dimensions[0..],
119 };
120
121 var graph = switch (instance.dtype) {
122 .f32 => try PhiloxKeyCounterUniformRuntimeFamilyF32.buildNamed(allocator, options.limits, entry_name, instance),
123 .i32 => try PhiloxKeyCounterUniformRuntimeFamilyI32.buildNamed(allocator, options.limits, entry_name, instance),
124 else => return error.UnsupportedDType,
125 };
126 defer graph.deinit();
127 return kernel.createKernelCallArtifact(allocator, handle, &graph, .{
128 .target = target,
129 .version = philox_key_counter_uniform_family_version,
130 .format = options.format,
131 .kernel_plan = options.kernel_plan,
132 .element_count_argument = options.element_count_argument,
133 .shape_family_fingerprint = family_fingerprint,
134 .shape_profile = shape_profile,
135 .launch = options.launch orelse try randomFoldDerivedLaunch(instance.threads),
136 .runtime_scalar_argument_count = if (options.runtime_scalar_argument_count == 0) 1 else options.runtime_scalar_argument_count,
137 .static_arguments = options.static_arguments,
138 });
139 }