lib/accy/src/kernel/library/loss.zig

daab053ee43316e1809a84551d573ddd1e5bf3d2

  1 const std = @import("std");
  2 const gpu = @import("gpu");
  3 const choir_abi = @import("choir_abi");
  4 
  5 const artifact_product = @import("../../artifact/model/root.zig");
  6 const shape = @import("../../choir/shape/root.zig");
  7 const entry = @import("entry.zig");
  8 const extent_mod = @import("extent.zig");
  9 const geometry_mod = @import("geometry.zig");
 10 const kernel = @import("../root.zig");
 11 
 12 const DType = choir_abi.DType;
 13 const runtimeExtentArgument = extent_mod.runtimeExtentArgument;
 14 
 15 const loss_thread_caps = geometry_mod.ThreadCaps1D{};
 16 
 17 pub const RowSparseCrossEntropy = struct {
 18     rows: u64,
 19     classes: u64,
 20     dtype: DType = .f32,
 21     threads: u32 = 256,
 22     row_axis: []const u8 = "r",
 23     class_axis: []const u8 = "c",
 24 
 25     pub fn total(self: RowSparseCrossEntropy) u64 {
 26         return self.rows;
 27     }
 28 };
 29 
 30 pub const row_sparse_cross_entropy_family_version: u32 = 1;
 31 
 32 pub fn rowSparseCrossEntropyDTypeSupported(dtype: DType) bool {
 33     return dtype == .f32;
 34 }
 35 
 36 pub fn rowSparseCrossEntropyInstanceValid(instance: RowSparseCrossEntropy) bool {
 37     if (!rowSparseCrossEntropyDTypeSupported(instance.dtype)) return false;
 38     if (instance.rows == 0 or instance.classes == 0) return false;
 39     return instance.threads != 0;
 40 }
 41 
 42 pub fn rowSparseCrossEntropyThreadsForRows(rows: u64) u32 {
 43     return geometry_mod.threadsForExtent(rows, loss_thread_caps);
 44 }
 45 
 46 pub fn rowSparseCrossEntropyThreadCandidatesForRows(rows: u64) geometry_mod.Thread1DCandidates {
 47     return geometry_mod.threadCandidatesForExtent(rows, loss_thread_caps);
 48 }
 49 
 50 pub fn rowSparseCrossEntropyFamilyTarget(allocator: std.mem.Allocator, instance: RowSparseCrossEntropy) ![]u8 {
 51     return std.fmt.allocPrint(
 52         allocator,
 53         "accy.kernel.loss.row_sparse_cross_entropy_family_{d}_{s}",
 54         .{ instance.threads, instance.dtype.name() },
 55     );
 56 }
 57 
 58 pub fn rowSparseCrossEntropyFamilyEntryName(allocator: std.mem.Allocator, instance: RowSparseCrossEntropy) ![]u8 {
 59     return std.fmt.allocPrint(
 60         allocator,
 61         "accy_kernel_loss_row_sparse_cross_entropy_family_{d}_{s}",
 62         .{ instance.threads, instance.dtype.name() },
 63     );
 64 }
 65 
 66 pub fn rowSparseCrossEntropyRuntimeArguments(instance: RowSparseCrossEntropy) ![2]choir_abi.ScalarArgument {
 67     return .{
 68         .{ .u32 = try runtimeExtentArgument(instance.rows) },
 69         .{ .u32 = try runtimeExtentArgument(instance.classes) },
 70     };
 71 }
 72 
 73 pub fn rowSparseCrossEntropyShapeProfileDimensions(
 74     instance: RowSparseCrossEntropy,
 75 ) [2]artifact_product.KernelCallShapeProfileDimension {
 76     const bounds = lossRuntimeExtentBounds();
 77     return .{
 78         .{ .name = instance.row_axis, .runtime_scalar_argument_index = 0, .bounds = bounds },
 79         .{ .name = instance.class_axis, .runtime_scalar_argument_index = 1, .bounds = bounds },
 80     };
 81 }
 82 
 83 fn lossRuntimeExtentBounds() shape.Bounds {
 84     return .{ .min = 1, .max = extent_mod.runtime_extent_max };
 85 }
 86 
 87 fn rowSparseCrossEntropyDerivedLaunch(instance: RowSparseCrossEntropy) !artifact_product.KernelCallLaunch {
 88     if (instance.threads == 0) return error.KernelLibraryLaunchThreadgroupMustBeNonzero;
 89     return .{ .derived = .{
 90         .grid = .{
 91             .{ .runtime_u32_ceil_div = .{ .argument_index = 0, .divisor = instance.threads } },
 92             .{ .fixed = 1 },
 93             .{ .fixed = 1 },
 94         },
 95         .threadgroup = .{ instance.threads, 1, 1 },
 96     } };
 97 }
 98 
 99 pub fn rowSparseCrossEntropyShapeFamily(
100     backing_allocator: std.mem.Allocator,
101     instance: RowSparseCrossEntropy,
102 ) !shape.Family {
103     var builder = try shape.Builder.init(backing_allocator, "row_sparse_cross_entropy");
104     errdefer builder.deinit();
105 
106     const row = try builder.symbol(instance.row_axis);
107     const class = try builder.symbol(instance.class_axis);
108     const row_expr = try builder.symbolExpression(row);
109     const class_expr = try builder.symbolExpression(class);
110 
111     _ = try builder.tensor("logits", &.{ row_expr, class_expr });
112     _ = try builder.tensor("targets", &.{row_expr});
113     _ = try builder.tensor("losses", &.{row_expr});
114     try builder.assumeBounds(row_expr, lossRuntimeExtentBounds());
115     try builder.assumeBounds(class_expr, lossRuntimeExtentBounds());
116     return builder.finish();
117 }
118 
119 pub fn rowSparseCrossEntropyFamilyFingerprint(
120     backing_allocator: std.mem.Allocator,
121     instance: RowSparseCrossEntropy,
122 ) !u64 {
123     var family = try rowSparseCrossEntropyShapeFamily(backing_allocator, instance);
124     defer family.deinit();
125     return shape.fingerprint(family);
126 }
127 
128 pub fn rowSparseCrossEntropyFamilySpecialization(
129     backing_allocator: std.mem.Allocator,
130     instance: RowSparseCrossEntropy,
131 ) !entry.OwnedSpecialization {
132     var owned = entry.OwnedSpecialization.init(backing_allocator);
133     errdefer owned.deinit();
134     const lifetime_allocator = owned.allocator();
135 
136     const inputs = try lifetime_allocator.alloc(entry.Shape, 2);
137     inputs[0] = try entry.runtimeShape2D(
138         lifetime_allocator,
139         instance.row_axis,
140         instance.rows,
141         instance.class_axis,
142         instance.classes,
143     );
144     inputs[1] = try entry.runtimeShape1D(lifetime_allocator, instance.row_axis, instance.rows);
145 
146     const outputs = try lifetime_allocator.alloc(entry.Shape, 1);
147     outputs[0] = try entry.runtimeShape1D(lifetime_allocator, instance.row_axis, instance.rows);
148 
149     owned.value = .{
150         .dtype = instance.dtype,
151         .operation = .{ .loss = .row_sparse_cross_entropy },
152         .inputs = inputs,
153         .outputs = outputs,
154         .schedule = try entry.runtimeThreadBlocks1D(lifetime_allocator, instance.row_axis, instance.rows, instance.threads),
155     };
156     owned.value.launch = owned.value.schedule.?.launch();
157     var family = try rowSparseCrossEntropyShapeFamily(backing_allocator, instance);
158     errdefer family.deinit();
159     try owned.takeShapeFamily(&family);
160     return owned;
161 }
162 
163 pub fn rowSparseCrossEntropyInstanceFromSpecialization(
164     specialization: entry.Specialization,
165 ) ?RowSparseCrossEntropy {
166     if (!specialization.scheduleMatchesLaunch()) return null;
167     if (!specialization.operationIs(.{ .loss = .row_sparse_cross_entropy })) return null;
168     const dtype = specialization.dtype orelse return null;
169     if (!rowSparseCrossEntropyDTypeSupported(dtype)) return null;
170     if (specialization.inputs.len != 2 or specialization.outputs.len != 1) return null;
171     if (specialization.reductions.len != 0) return null;
172     const logits = specialization.inputs[0];
173     const targets = specialization.inputs[1];
174     const losses = specialization.outputs[0];
175     if (logits.axes.len != 2 or targets.axes.len != 1 or losses.axes.len != 1) return null;
176     const rows = logits.axes[0].extent;
177     const classes = logits.axes[1].extent;
178     if (targets.axes[0].extent != rows or losses.axes[0].extent != rows) return null;
179     if (!std.mem.eql(u8, logits.axes[0].name, targets.axes[0].name)) return null;
180     if (!std.mem.eql(u8, logits.axes[0].name, losses.axes[0].name)) return null;
181     const launch = specialization.launch orelse return null;
182     if (launch.threadgroup[0] == 0) return null;
183     return .{
184         .rows = rows,
185         .classes = classes,
186         .dtype = dtype,
187         .threads = launch.threadgroup[0],
188         .row_axis = logits.axes[0].name,
189         .class_axis = logits.axes[1].name,
190     };
191 }
192 
193 fn rowSparseCrossEntropyFamilySchedule(instance: RowSparseCrossEntropy) kernel.logical.schedule.ThreadBlocks {
194     return kernel.logical.schedule.threadBlocks(.{ .x = instance.threads });
195 }
196 
197 fn row_sparse_cross_entropy_runtime_body_active(b: anytype, ctx: anytype) !void {
198     const classes_extent = try b.castIndex(ctx.args.param(.classes).raw());
199     const row_base = try b.mul(ctx.row, classes_extent);
200     const one_index = try b.constantIndex(1);
201 
202     const first_logit = (try ctx.args.param(.logits).load(b, row_base)).raw();
203     const row_max = try b.fold(one_index, classes_extent, one_index, first_logit, .{
204         .args = ctx.args,
205         .row_base = row_base,
206     }, row_sparse_cross_entropy_runtime_body_row_max);
207 
208     const zero_index = try b.constantIndex(0);
209     const zero_value = try b.constantFloat(ctx.dtype, 0.0);
210     const exp_sum = try b.fold(zero_index, classes_extent, one_index, zero_value, .{
211         .args = ctx.args,
212         .row_base = row_base,
213         .row_max = row_max,
214     }, row_sparse_cross_entropy_runtime_body_exp_sum);
215 
216     const target = try b.castIndex((try ctx.args.param(.targets).load(b, ctx.row)).raw());
217     const non_negative = try b.compare(.ge, target, zero_index);
218     const in_range = try b.compare(.lt, target, classes_extent);
219     const valid = try b.and_(non_negative, in_range);
220     const clamped = try b.select(valid, target, zero_index);
221     const target_logit = (try ctx.args.param(.logits).load(b, try b.add(row_base, clamped))).raw();
222     const target_shifted = try b.select(valid, try b.sub(target_logit, row_max), zero_value);
223 
224     const loss = try b.sub(try b.log(exp_sum), target_shifted);
225     try ctx.args.param(.losses).store(b, loss, ctx.row);
226 }
227 
228 fn row_sparse_cross_entropy_runtime_body_row_max(fb: anytype, class_index: kernel.Value, acc: kernel.Value, fold_ctx: anytype) !kernel.Value {
229     const position = try fb.add(fold_ctx.row_base, class_index);
230     const value = (try fold_ctx.args.param(.logits).load(fb, position)).raw();
231     return fb.max(acc, value);
232 }
233 
234 fn row_sparse_cross_entropy_runtime_body_exp_sum(fb: anytype, class_index: kernel.Value, acc: kernel.Value, fold_ctx: anytype) !kernel.Value {
235     const position = try fb.add(fold_ctx.row_base, class_index);
236     const value = (try fold_ctx.args.param(.logits).load(fb, position)).raw();
237     const shifted = try fb.sub(value, fold_ctx.row_max);
238     return fb.add(acc, try fb.exp(shifted));
239 }
240 
241 fn rowSparseCrossEntropyRuntimeBody(k: anytype, spec: RowSparseCrossEntropy, args: anytype) !void {
242     const row = try k.globalId(.x);
243     const rows_extent = try k.castIndex(args.param(.rows).raw());
244     const active = try k.compare(.lt, row, rows_extent);
245     try k.guardDo(active, .{
246         .args = args,
247         .row = row,
248         .dtype = spec.dtype,
249     }, row_sparse_cross_entropy_runtime_body_active);
250 }
251 
252 fn rowSparseCrossEntropyRuntimeFamily(comptime dtype: DType) type {
253     return kernel.logical.Family(.{
254         .name = std.fmt.comptimePrint("accy_kernel_loss_row_sparse_cross_entropy_runtime_{s}", .{dtype.name()}),
255         .parameters = .{
256             .losses = kernel.dynamicBuffer(dtype),
257             .logits = kernel.dynamicBuffer(dtype),
258             .targets = kernel.dynamicBuffer(.i32),
259             .rows = kernel.scalar(.i32),
260             .classes = kernel.scalar(.i32),
261         },
262         .Instance = RowSparseCrossEntropy,
263         .schedule = rowSparseCrossEntropyFamilySchedule,
264         .body = rowSparseCrossEntropyRuntimeBody,
265     });
266 }
267 
268 pub const RowSparseCrossEntropyRuntimeFamilyF32 = rowSparseCrossEntropyRuntimeFamily(.f32);
269 
270 pub fn createRowSparseCrossEntropyFamilyArtifact(
271     allocator: std.mem.Allocator,
272     handle: kernel.BackendHandle,
273     instance: RowSparseCrossEntropy,
274     options: entry.ArtifactOptions,
275 ) !kernel.OwnedKernelCallArtifact {
276     if (!rowSparseCrossEntropyInstanceValid(instance)) return error.InvalidKernelLibraryEntry;
277     const target = try rowSparseCrossEntropyFamilyTarget(allocator, instance);
278     defer allocator.free(target);
279     const entry_name = try rowSparseCrossEntropyFamilyEntryName(allocator, instance);
280     defer allocator.free(entry_name);
281     const family_fingerprint = options.shape_family_fingerprint orelse try rowSparseCrossEntropyFamilyFingerprint(allocator, instance);
282     const shape_profile_dimensions = rowSparseCrossEntropyShapeProfileDimensions(instance);
283     const shape_profile = options.shape_profile orelse artifact_product.KernelCallShapeProfile{
284         .name = "row_sparse_cross_entropy",
285         .fingerprint = family_fingerprint,
286         .dimensions = shape_profile_dimensions[0..],
287     };
288 
289     var graph = switch (instance.dtype) {
290         .f32 => try RowSparseCrossEntropyRuntimeFamilyF32.buildNamed(allocator, options.limits, entry_name, instance),
291         else => return error.UnsupportedDType,
292     };
293     defer graph.deinit();
294     return kernel.createKernelCallArtifact(allocator, handle, &graph, .{
295         .target = target,
296         .version = row_sparse_cross_entropy_family_version,
297         .format = options.format,
298         .kernel_plan = options.kernel_plan,
299         .element_count_argument = options.element_count_argument,
300         .shape_family_fingerprint = family_fingerprint,
301         .shape_profile = shape_profile,
302         .launch = options.launch orelse try rowSparseCrossEntropyDerivedLaunch(instance),
303         .runtime_scalar_argument_count = if (options.runtime_scalar_argument_count == 0) 2 else options.runtime_scalar_argument_count,
304         .static_arguments = options.static_arguments,
305     });
306 }
307 
308 pub fn hostRowSparseCrossEntropy(
309     rows: usize,
310     classes: usize,
311     logits: []const f32,
312     targets: []const i32,
313     losses: []f32,
314 ) void {
315     for (0..rows) |row| {
316         const row_logits = logits[row * classes ..][0..classes];
317         var row_max = row_logits[0];
318         for (row_logits) |value| row_max = @max(row_max, value);
319         var exp_sum: f32 = 0;
320         for (row_logits) |value| exp_sum += @exp(value - row_max);
321         const target = targets[row];
322         const target_shifted = if (target >= 0 and target < classes)
323             row_logits[@intCast(target)] - row_max
324         else
325             0;
326         losses[row] = @log(exp_sum) - target_shifted;
327     }
328 }
329 
330 test "loss row sparse cross entropy runtime family matches the host oracle" {
331     const allocator = std.testing.allocator;
332     const compiled = RowSparseCrossEntropy{ .rows = 1, .classes = 1, .threads = 32 };
333     const runtime = RowSparseCrossEntropy{ .rows = 4, .classes = 5, .threads = 32 };
334 
335     var graph = try RowSparseCrossEntropyRuntimeFamilyF32.build(allocator, RowSparseCrossEntropyRuntimeFamilyF32.Limits.testing, compiled);
336     defer graph.deinit();
337 
338     var logits = [_]f32{
339         0.5,   -1.0, 2.0,  0.0,  1.5,
340         -0.25, 0.75, -2.0, 3.0,  0.125,
341         1.0,   1.0,  1.0,  1.0,  1.0,
342         -3.0,  4.0,  0.5,  -0.5, 2.5,
343     };
344     var targets = [4]i32{ 2, 0, 4, -1 };
345     var expected: [4]f32 = undefined;
346     hostRowSparseCrossEntropy(4, 5, logits[0..], targets[0..], expected[0..]);
347 
348     var losses = @as([4]f32, @splat(0));
349     const launch_value = try entry.runtimeLaunch1D(runtime.total(), runtime.threads);
350     try graph.runCpuWithLaunch(allocator, &.{
351         kernel.argumentBuffer(f32, losses[0..]),
352         kernel.argumentBuffer(f32, logits[0..]),
353         kernel.argumentBuffer(i32, targets[0..]),
354         kernel.argumentI32(@intCast(runtime.rows)),
355         kernel.argumentI32(@intCast(runtime.classes)),
356     }, .{
357         .grid = launch_value.grid,
358         .block = launch_value.threadgroup,
359     });
360 
361     for (expected, losses) |want, got| {
362         try std.testing.expectApproxEqAbs(want, got, 0.0001);
363     }
364 }
365 
366 test "loss row sparse cross entropy instance round-trips through specialization" {
367     const instance = RowSparseCrossEntropy{ .rows = 96, .classes = 11, .threads = 64 };
368     var owned = try rowSparseCrossEntropyFamilySpecialization(std.testing.allocator, instance);
369     defer owned.deinit();
370 
371     const recovered = rowSparseCrossEntropyInstanceFromSpecialization(owned.value) orelse {
372         return error.TestExpectedRowSparseCrossEntropyInstance;
373     };
374     try std.testing.expectEqual(instance.rows, recovered.rows);
375     try std.testing.expectEqual(instance.classes, recovered.classes);
376     try std.testing.expectEqual(instance.dtype, recovered.dtype);
377     try std.testing.expectEqual(instance.threads, recovered.threads);
378 }
379 
380 test "loss row sparse cross entropy family artifact carries runtime launch contract" {
381     const allocator = std.testing.allocator;
382     var state = gpu.recording.BackendState{
383         .allocator = allocator,
384         .kind = .cuda,
385         .format = .cuda_ptx,
386     };
387     const instance = RowSparseCrossEntropy{ .rows = 8, .classes = 16, .threads = 8 };
388 
389     var family_artifact = try createRowSparseCrossEntropyFamilyArtifact(allocator, state.handle(), instance, .{ .limits = .testing });
390     defer family_artifact.deinit();
391 
392     const family_entry = family_artifact.entry();
393     try std.testing.expectEqualStrings("accy.kernel.loss.row_sparse_cross_entropy_family_8_f32", family_entry.target);
394     try std.testing.expectEqual(@as(u32, 2), family_entry.runtime_scalar_argument_count);
395     try std.testing.expect(family_entry.required_dtypes.contains(.f32));
396     const profile = family_entry.shape_profile orelse return error.TestExpectedShapeProfile;
397     try std.testing.expectEqualStrings("row_sparse_cross_entropy", profile.name);
398     try std.testing.expectEqual(@as(usize, 2), profile.dimensions.len);
399     switch (family_entry.launch) {
400         .derived => |launch| {
401             try std.testing.expectEqual(@as(u32, 8), launch.threadgroup[0]);
402             switch (launch.grid[0]) {
403                 .runtime_u32_ceil_div => |term| {
404                     try std.testing.expectEqual(@as(usize, 0), term.argument_index);
405                     try std.testing.expectEqual(@as(u32, 8), term.divisor);
406                 },
407                 else => return error.TestExpectedDerivedLaunch,
408             }
409         },
410         else => return error.TestExpectedDerivedLaunch,
411     }
412 }