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 }