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 }