lib/accy/src/kernel/library/catalog/family/loss.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2
3 const descriptor_mod = @import("../root.zig");
4 const geometry = @import("../../root.zig").geometry;
5 const loss = @import("../../root.zig").loss;
6 const match_mod = @import("../match/root.zig");
7 const query_mod = descriptor_mod;
8
9 const Descriptor = descriptor_mod.Descriptor;
10 const OwnedDescriptor = descriptor_mod.OwnedDescriptor;
11 const RowSparseCrossEntropyQuery = query_mod.RowSparseCrossEntropyQuery;
12
13 pub const RowSparseCrossEntropyCandidateDescriptors = struct {
14 count: usize = 0,
15 items: [geometry.max_thread_candidates]OwnedDescriptor = undefined,
16
17 pub fn slice(self: *const RowSparseCrossEntropyCandidateDescriptors) []const OwnedDescriptor {
18 return self.items[0..self.count];
19 }
20
21 pub fn deinit(self: *RowSparseCrossEntropyCandidateDescriptors) void {
22 for (self.items[0..self.count]) |*descriptor| descriptor.deinit();
23 self.* = undefined;
24 }
25 };
26
27 pub fn selectRowSparseCrossEntropy(
28 backing_allocator: std.mem.Allocator,
29 query: RowSparseCrossEntropyQuery,
30 ) !?OwnedDescriptor {
31 if (!canonicalRowSparseCrossEntropy(query)) return null;
32 const instance = rowSparseCrossEntropyFamilyInstance(query) orelse return null;
33 return try rowSparseCrossEntropyDescriptorForInstance(backing_allocator, instance, query);
34 }
35
36 pub fn selectRowSparseCrossEntropyCandidates(
37 backing_allocator: std.mem.Allocator,
38 query: RowSparseCrossEntropyQuery,
39 ) !RowSparseCrossEntropyCandidateDescriptors {
40 var result = RowSparseCrossEntropyCandidateDescriptors{};
41 errdefer result.deinit();
42
43 if (!canonicalRowSparseCrossEntropy(query)) return result;
44 if (query.schedule != null) {
45 if (try selectRowSparseCrossEntropy(backing_allocator, query)) |descriptor| {
46 result.items[result.count] = descriptor;
47 result.count += 1;
48 }
49 return result;
50 }
51
52 const thread_candidates = loss.rowSparseCrossEntropyThreadCandidatesForRows(query.rows);
53 for (thread_candidates.slice()) |threads| {
54 const instance = loss.RowSparseCrossEntropy{
55 .rows = query.rows,
56 .classes = query.classes,
57 .dtype = query.dtype,
58 .threads = threads,
59 };
60 if (!loss.rowSparseCrossEntropyInstanceValid(instance)) continue;
61 if (result.count >= result.items.len) break;
62 if (try rowSparseCrossEntropyDescriptorForInstance(backing_allocator, instance, query)) |descriptor| {
63 result.items[result.count] = descriptor;
64 result.count += 1;
65 }
66 }
67 return result;
68 }
69
70 fn canonicalRowSparseCrossEntropy(query: RowSparseCrossEntropyQuery) bool {
71 if (!loss.rowSparseCrossEntropyDTypeSupported(query.dtype)) return false;
72 return query.rows != 0 and query.classes != 0;
73 }
74
75 fn rowSparseCrossEntropyDescriptorForInstance(
76 backing_allocator: std.mem.Allocator,
77 instance: loss.RowSparseCrossEntropy,
78 query: RowSparseCrossEntropyQuery,
79 ) !?OwnedDescriptor {
80 var specialization = try loss.rowSparseCrossEntropyFamilySpecialization(backing_allocator, instance);
81 errdefer specialization.deinit();
82 const lifetime_allocator = specialization.allocator();
83 const descriptor = Descriptor{
84 .name = try loss.rowSparseCrossEntropyFamilyEntryName(lifetime_allocator, instance),
85 .metadata = .{
86 .target = try loss.rowSparseCrossEntropyFamilyTarget(lifetime_allocator, instance),
87 .version = loss.row_sparse_cross_entropy_family_version,
88 .layer = .logical,
89 .category = .loss,
90 .specialization = specialization.value,
91 },
92 };
93 if (!match_mod.rowSparseCrossEntropyDescriptorMatches(descriptor, query)) {
94 specialization.deinit();
95 return null;
96 }
97 return .{ .descriptor = descriptor, .specialization = specialization };
98 }
99
100 fn rowSparseCrossEntropyFamilyInstance(query: RowSparseCrossEntropyQuery) ?loss.RowSparseCrossEntropy {
101 var instance = loss.RowSparseCrossEntropy{
102 .rows = query.rows,
103 .classes = query.classes,
104 .dtype = query.dtype,
105 .threads = loss.rowSparseCrossEntropyThreadsForRows(query.rows),
106 };
107 if (query.schedule) |requested| {
108 switch (requested) {
109 .thread_blocks => |threads| {
110 if (threads == 0) return null;
111 if (@as(u64, threads) > instance.rows) return null;
112 instance.threads = threads;
113 },
114 }
115 }
116 if (!loss.rowSparseCrossEntropyInstanceValid(instance)) return null;
117 return instance;
118 }
119
120 test "catalog row sparse cross entropy family selection answers canonical queries" {
121 const allocator = std.testing.allocator;
122
123 var selected = (try selectRowSparseCrossEntropy(allocator, .{
124 .dtype = .f32,
125 .rows = 512,
126 .classes = 1000,
127 })) orelse return error.TestExpectedRowSparseCrossEntropyDescriptor;
128 defer selected.deinit();
129 try std.testing.expect(std.mem.startsWith(u8, selected.descriptor.metadata.target, "accy.kernel.loss.row_sparse_cross_entropy_family_"));
130 try std.testing.expect(selected.descriptor.metadata.specialization.operationIs(.{ .loss = .row_sparse_cross_entropy }));
131
132 var scheduled = (try selectRowSparseCrossEntropy(allocator, .{
133 .dtype = .f32,
134 .rows = 512,
135 .classes = 1000,
136 .schedule = .{ .thread_blocks = 64 },
137 })) orelse return error.TestExpectedRowSparseCrossEntropyDescriptor;
138 defer scheduled.deinit();
139 try std.testing.expectEqualStrings("accy.kernel.loss.row_sparse_cross_entropy_family_64_f32", scheduled.descriptor.metadata.target);
140
141 try std.testing.expectEqual(@as(?OwnedDescriptor, null), try selectRowSparseCrossEntropy(allocator, .{
142 .dtype = .f16,
143 .rows = 512,
144 .classes = 1000,
145 }));
146 try std.testing.expectEqual(@as(?OwnedDescriptor, null), try selectRowSparseCrossEntropy(allocator, .{
147 .dtype = .f32,
148 .rows = 0,
149 .classes = 1000,
150 }));
151 }
152
153 test "catalog row sparse cross entropy candidates enumerate thread geometries" {
154 const allocator = std.testing.allocator;
155
156 var candidates = try selectRowSparseCrossEntropyCandidates(allocator, .{
157 .dtype = .f32,
158 .rows = 4096,
159 .classes = 128,
160 });
161 defer candidates.deinit();
162 try std.testing.expect(candidates.count >= 2);
163 for (candidates.slice()) |candidate| {
164 try std.testing.expect(candidate.descriptor.metadata.specialization.operationIs(.{ .loss = .row_sparse_cross_entropy }));
165 }
166
167 var pinned = try selectRowSparseCrossEntropyCandidates(allocator, .{
168 .dtype = .f32,
169 .rows = 4096,
170 .classes = 128,
171 .schedule = .{ .thread_blocks = 128 },
172 });
173 defer pinned.deinit();
174 try std.testing.expectEqual(@as(usize, 1), pinned.count);
175 try std.testing.expectEqualStrings(
176 "accy.kernel.loss.row_sparse_cross_entropy_family_128_f32",
177 pinned.slice()[0].descriptor.metadata.target,
178 );
179 }