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 }