lib/accy/src/kernel/library/random/feistel/artifact.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const random = @import("../root.zig");
2 const model = @import("model.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 Feistel = model.Feistel;
12 const FeistelRuntimeFamilyI32 = runtime.FeistelRuntimeFamilyI32;
13 const feistel_family_version = model.feistel_family_version;
14 const feistelFamilyTarget = model.feistelFamilyTarget;
15 const feistelFamilyEntryName = model.feistelFamilyEntryName;
16 const feistelFamilyFingerprint = model.feistelFamilyFingerprint;
17 const feistelShapeProfileDimensions = model.feistelShapeProfileDimensions;
18 const feistelInstanceValid = model.feistelInstanceValid;
19 const randomFoldDerivedLaunch = base.randomFoldDerivedLaunch;
20
21 pub fn createFeistelFamilyArtifact(
22 allocator: std.mem.Allocator,
23 handle: kernel.BackendHandle,
24 instance: Feistel,
25 options: entry.ArtifactOptions,
26 ) !kernel.OwnedKernelCallArtifact {
27 if (!feistelInstanceValid(instance)) return error.UnsupportedExtent;
28 const target = try feistelFamilyTarget(allocator, instance);
29 defer allocator.free(target);
30 const entry_name = try feistelFamilyEntryName(allocator, instance);
31 defer allocator.free(entry_name);
32 const family_fingerprint = options.shape_family_fingerprint orelse try feistelFamilyFingerprint(allocator, instance);
33 const shape_profile_dimensions = feistelShapeProfileDimensions(instance);
34 const shape_profile = options.shape_profile orelse artifact_product.KernelCallShapeProfile{
35 .name = "feistel",
36 .fingerprint = family_fingerprint,
37 .dimensions = shape_profile_dimensions[0..],
38 };
39
40 var graph = switch (instance.dtype) {
41 .i32 => try FeistelRuntimeFamilyI32.buildNamed(allocator, options.limits, entry_name, instance),
42 else => return error.UnsupportedDType,
43 };
44 defer graph.deinit();
45 return kernel.createKernelCallArtifact(allocator, handle, &graph, .{
46 .target = target,
47 .version = feistel_family_version,
48 .format = options.format,
49 .kernel_plan = options.kernel_plan,
50 .element_count_argument = options.element_count_argument,
51 .shape_family_fingerprint = family_fingerprint,
52 .shape_profile = shape_profile,
53 .launch = options.launch orelse try randomFoldDerivedLaunch(instance.threads),
54 .runtime_scalar_argument_count = if (options.runtime_scalar_argument_count == 0) 3 else options.runtime_scalar_argument_count,
55 .static_arguments = options.static_arguments,
56 });
57 }