lib/accy/src/kernel/library/catalog/family/random/squares.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2
3 const descriptor_mod = @import("../../root.zig");
4 const random = @import("../../../root.zig").random;
5 const match_mod = @import("../../match/root.zig");
6 const query_mod = descriptor_mod;
7
8 const candidate_mod = @import("candidate.zig");
9
10 const Descriptor = descriptor_mod.Descriptor;
11 const OwnedDescriptor = descriptor_mod.OwnedDescriptor;
12 const RandomCandidateDescriptors = candidate_mod.RandomCandidateDescriptors;
13 const RandomQuery = query_mod.RandomQuery;
14
15 pub fn select(backing_allocator: std.mem.Allocator, query: RandomQuery) !?OwnedDescriptor {
16 return descriptorForInstance(backing_allocator, familyInstance(query) orelse return null, query);
17 }
18
19 pub fn appendCandidates(
20 backing_allocator: std.mem.Allocator,
21 query: RandomQuery,
22 result: *RandomCandidateDescriptors,
23 ) !void {
24 const thread_candidates = random.squaresThreadCandidatesForCount(query.count);
25 for (thread_candidates.slice()) |threads| {
26 const instance = random.Squares{
27 .count = query.count,
28 .dtype = query.dtype,
29 .threads = threads,
30 };
31 if (try descriptorForInstance(backing_allocator, instance, query)) |descriptor| {
32 result.items[result.count] = descriptor;
33 result.count += 1;
34 }
35 }
36 }
37
38 fn descriptorForInstance(
39 backing_allocator: std.mem.Allocator,
40 instance: random.Squares,
41 query: RandomQuery,
42 ) !?OwnedDescriptor {
43 var specialization = try random.squaresFamilySpecialization(backing_allocator, instance);
44 errdefer specialization.deinit();
45 const lifetime_allocator = specialization.allocator();
46 const descriptor = Descriptor{
47 .name = try random.squaresFamilyEntryName(lifetime_allocator, instance),
48 .metadata = .{
49 .target = try random.squaresFamilyTarget(lifetime_allocator, instance),
50 .version = random.squares_family_version,
51 .layer = .logical,
52 .category = .random,
53 .specialization = specialization.value,
54 },
55 };
56 if (!match_mod.randomDescriptorMatches(descriptor, query)) {
57 specialization.deinit();
58 return null;
59 }
60 return .{ .descriptor = descriptor, .specialization = specialization };
61 }
62
63 fn familyInstance(query: RandomQuery) ?random.Squares {
64 var instance = random.Squares{
65 .count = query.count,
66 .dtype = query.dtype,
67 .threads = random.squaresThreadsForCount(query.count),
68 };
69 if (query.schedule) |requested| {
70 switch (requested) {
71 .thread_blocks => |threads| {
72 if (threads == 0) return null;
73 if (@as(u64, threads) > instance.generators()) return null;
74 instance.threads = threads;
75 },
76 }
77 }
78 return instance;
79 }