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 }