lib/accy/src/kernel/library/segmented.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 const tuning = @import("tuning.zig");
 12 
 13 const DType = choir_abi.DType;
 14 const indexExtent = extent_mod.indexExtent;
 15 const runtimeExtentArgument = extent_mod.runtimeExtentArgument;
 16 
 17 pub const SegmentSumGranularity = enum {
 18     thread,
 19     warp,
 20 
 21     pub fn name(self: SegmentSumGranularity) []const u8 {
 22         return switch (self) {
 23             .thread => "thread",
 24             .warp => "warp",
 25         };
 26     }
 27 };
 28 
 29 pub const SegmentSum = struct {
 30     segments: u64,
 31     total: u64,
 32     dtype: DType = .f32,
 33     accumulation_dtype: DType = .f32,
 34     granularity: SegmentSumGranularity = .thread,
 35     threads: u32 = 256,
 36     segment_axis: []const u8 = "s",
 37     element_axis: []const u8 = "e",
 38 
 39     pub fn launchExtent(self: SegmentSum) u64 {
 40         return switch (self.granularity) {
 41             .thread => self.segments,
 42             .warp => self.segments * segment_sum_warp_size,
 43         };
 44     }
 45 };
 46 
 47 pub const segment_sum_family_version: u32 = 1;
 48 pub const segment_sum_warp_size: u32 = 32;
 49 const segment_sum_thread_caps = geometry_mod.ThreadCaps1D{};
 50 
 51 pub fn segmentSumGranularityValid(granularity: SegmentSumGranularity, threads: u32) bool {
 52     return switch (granularity) {
 53         .thread => threads != 0,
 54         .warp => threads != 0 and threads % segment_sum_warp_size == 0,
 55     };
 56 }
 57 
 58 pub fn segmentSumDTypeSupported(dtype: DType) bool {
 59     return switch (dtype) {
 60         .f32, .f16 => true,
 61         else => false,
 62     };
 63 }
 64 
 65 pub fn segmentSumAccumulationDType(dtype: DType) ?DType {
 66     return switch (dtype) {
 67         .f32, .f16 => .f32,
 68         else => null,
 69     };
 70 }
 71 
 72 fn segmentSumAccumulationZero(inner_builder: anytype, spec: SegmentSum) !kernel.Value {
 73     return switch (spec.accumulation_dtype) {
 74         .f32 => inner_builder.constantFloat(.f32, 0.0),
 75         .f16 => inner_builder.constantFloat(.f16, 0.0),
 76         else => error.UnsupportedDType,
 77     };
 78 }
 79 
 80 fn segmentSumAccumulationValue(inner_builder: anytype, spec: SegmentSum, value: anytype) !kernel.Value {
 81     return switch (spec.accumulation_dtype) {
 82         .f32 => if (comptime @TypeOf(value).scalar_dtype == .f32) value.raw() else (try value.cast(inner_builder, .f32)).raw(),
 83         .f16 => if (comptime @TypeOf(value).scalar_dtype == .f16) value.raw() else (try value.cast(inner_builder, .f16)).raw(),
 84         else => error.UnsupportedDType,
 85     };
 86 }
 87 
 88 fn segmentSumOutputValue(inner_builder: anytype, spec: SegmentSum, value: kernel.Value) !kernel.Value {
 89     if (spec.dtype == spec.accumulation_dtype) return value;
 90     return switch (spec.dtype) {
 91         .f32 => inner_builder.cast(value, .f32),
 92         .f16 => inner_builder.cast(value, .f16),
 93         else => error.UnsupportedDType,
 94     };
 95 }
 96 
 97 fn segmentOffsetIndex(inner_builder: anytype, loaded: anytype, total: kernel.Value) !kernel.Value {
 98     const zero_i32 = try inner_builder.constantInt(.i32, 0);
 99     const max_i32_index = try inner_builder.constantIndex(std.math.maxInt(i32));
100     const total_limit = try inner_builder.min(total, max_i32_index);
101     const total_i32 = try inner_builder.cast(total_limit, .i32);
102     const lower = try inner_builder.max(loaded.raw(), zero_i32);
103     const clamped = try inner_builder.min(lower, total_i32);
104     return inner_builder.castIndex(clamped);
105 }
106 
107 fn segment_sum_value_apply(fold_builder: anytype, element: kernel.Value, current: kernel.Value, ctx: anytype) !kernel.Value {
108     const value = try ctx.args.param(.data).load(fold_builder, element);
109     const value_acc = try segmentSumAccumulationValue(fold_builder, ctx.spec, value);
110     return fold_builder.add(current, value_acc);
111 }
112 
113 fn segmentSumValue(
114     inner_builder: anytype,
115     spec: SegmentSum,
116     args: anytype,
117     segment: kernel.Value,
118     total: kernel.Value,
119 ) !kernel.Value {
120     const one = try inner_builder.constantIndex(1);
121     const next = try inner_builder.add(segment, one);
122     const begin_loaded = try args.param(.offsets).load(inner_builder, segment);
123     const end_loaded = try args.param(.offsets).load(inner_builder, next);
124     const end_clamped = try segmentOffsetIndex(inner_builder, end_loaded, total);
125     const begin_offset = try segmentOffsetIndex(inner_builder, begin_loaded, total);
126     const begin_clamped = try inner_builder.min(begin_offset, end_clamped);
127 
128     const acc_zero = try segmentSumAccumulationZero(inner_builder, spec);
129     return inner_builder.fold(begin_clamped, end_clamped, one, acc_zero, .{
130         .spec = spec,
131         .args = args,
132     }, segment_sum_value_apply);
133 }
134 
135 fn segment_sum_warp_value_apply(fold_builder: anytype, element: kernel.Value, current: kernel.Value, ctx: anytype) !kernel.Value {
136     const value = try ctx.args.param(.data).load(fold_builder, element);
137     const value_acc = try segmentSumAccumulationValue(fold_builder, ctx.spec, value);
138     return fold_builder.add(current, value_acc);
139 }
140 
141 fn segmentSumWarpValue(
142     inner_builder: anytype,
143     spec: SegmentSum,
144     args: anytype,
145     segment: kernel.Value,
146     lane: kernel.Value,
147     total: kernel.Value,
148 ) !kernel.Value {
149     const one = try inner_builder.constantIndex(1);
150     const next = try inner_builder.add(segment, one);
151     const begin_loaded = try args.param(.offsets).load(inner_builder, segment);
152     const end_loaded = try args.param(.offsets).load(inner_builder, next);
153     const end_clamped = try segmentOffsetIndex(inner_builder, end_loaded, total);
154     const begin_offset = try segmentOffsetIndex(inner_builder, begin_loaded, total);
155     const begin_clamped = try inner_builder.min(begin_offset, end_clamped);
156     const lane_begin = try inner_builder.add(begin_clamped, lane);
157     const stride = try inner_builder.constantIndex(segment_sum_warp_size);
158 
159     const acc_zero = try segmentSumAccumulationZero(inner_builder, spec);
160     const partial = try inner_builder.fold(lane_begin, end_clamped, stride, acc_zero, .{
161         .spec = spec,
162         .args = args,
163     }, segment_sum_warp_value_apply);
164     return inner_builder.warpReduce(.add, partial);
165 }
166 
167 fn segment_sum_body_each(inner_builder: anytype, index: kernel.Index1D, ctx: anytype) !void {
168     const total = try inner_builder.constantIndex(try indexExtent(ctx.spec.total));
169     const sum = try segmentSumValue(inner_builder, ctx.spec, ctx.args, index.index, total);
170     const output = try segmentSumOutputValue(inner_builder, ctx.spec, sum);
171     try ctx.args.param(.dst).store(inner_builder, output, index);
172 }
173 
174 fn segmentSumBody(k: anytype, spec: SegmentSum, args: anytype) !void {
175     if (spec.granularity != .thread) return error.UnsupportedGranularity;
176     _ = try k.forEach1D(spec.segment_axis, spec.segments, .{ .spec = spec, .args = args }, segment_sum_body_each);
177 }
178 
179 fn segmentSumRuntimeBody(k: anytype, spec: SegmentSum, args: anytype) !void {
180     return switch (spec.granularity) {
181         .thread => segmentSumThreadRuntimeBody(k, spec, args),
182         .warp => segmentSumWarpRuntimeBody(k, spec, args),
183     };
184 }
185 
186 fn segment_sum_thread_runtime_body_active(inner_builder: anytype, ctx: anytype) !void {
187     const sum = try segmentSumValue(inner_builder, ctx.spec, ctx.args, ctx.segment, ctx.total);
188     const output = try segmentSumOutputValue(inner_builder, ctx.spec, sum);
189     try ctx.args.param(.dst).store(inner_builder, output, ctx.segment);
190 }
191 
192 fn segmentSumThreadRuntimeBody(k: anytype, spec: SegmentSum, args: anytype) !void {
193     const segment = try k.globalId(.x);
194     const segments_extent = try k.castIndex(args.param(.segments).raw());
195     const total = try k.castIndex(args.param(.total).raw());
196     const active = try k.compare(.lt, segment, segments_extent);
197     try k.guardDo(active, .{
198         .spec = spec,
199         .args = args,
200         .segment = segment,
201         .total = total,
202     }, segment_sum_thread_runtime_body_active);
203 }
204 
205 fn segment_sum_warp_runtime_body_active(inner_builder: anytype, ctx: anytype) !void {
206     const sum = try segmentSumWarpValue(inner_builder, ctx.spec, ctx.args, ctx.segment, ctx.lane, ctx.total);
207     const zero = try inner_builder.constantIndex(0);
208     const writer = try inner_builder.compare(.eq, ctx.lane, zero);
209     try inner_builder.guardDo(writer, .{
210         .spec = ctx.spec,
211         .args = ctx.args,
212         .segment = ctx.segment,
213         .sum = sum,
214     }, segment_sum_warp_runtime_body_writer);
215 }
216 
217 fn segment_sum_warp_runtime_body_writer(writer_builder: anytype, writer_ctx: anytype) !void {
218     const output = try segmentSumOutputValue(writer_builder, writer_ctx.spec, writer_ctx.sum);
219     try writer_ctx.args.param(.dst).store(writer_builder, output, writer_ctx.segment);
220 }
221 
222 fn segmentSumWarpRuntimeBody(k: anytype, spec: SegmentSum, args: anytype) !void {
223     if (!segmentSumGranularityValid(.warp, spec.threads)) return error.UnsupportedGranularity;
224     const element_thread = try k.globalId(.x);
225     const lane = try k.laneId();
226     const warp_size = try k.constantIndex(segment_sum_warp_size);
227     const segment = try k.div(element_thread, warp_size);
228     const segments_extent = try k.castIndex(args.param(.segments).raw());
229     const total = try k.castIndex(args.param(.total).raw());
230     const active = try k.compare(.lt, segment, segments_extent);
231     try k.guardDo(active, .{
232         .spec = spec,
233         .args = args,
234         .segment = segment,
235         .lane = lane,
236         .total = total,
237     }, segment_sum_warp_runtime_body_active);
238 }
239 
240 fn segmentSumFamilySchedule(instance: SegmentSum) kernel.logical.schedule.ThreadBlocks {
241     return kernel.logical.schedule.threadBlocks(.{ .x = instance.threads });
242 }
243 
244 fn segmentSumFamily(comptime dtype: DType) type {
245     return kernel.logical.Family(.{
246         .name = std.fmt.comptimePrint("accy_kernel_segmented_segment_sum_{s}", .{dtype.name()}),
247         .parameters = .{
248             .dst = kernel.dynamicBuffer(dtype),
249             .data = kernel.dynamicBuffer(dtype),
250             .offsets = kernel.dynamicBuffer(.i32),
251         },
252         .Instance = SegmentSum,
253         .schedule = segmentSumFamilySchedule,
254         .body = segmentSumBody,
255     });
256 }
257 
258 fn segmentSumRuntimeFamily(comptime dtype: DType) type {
259     return kernel.logical.Family(.{
260         .name = std.fmt.comptimePrint("accy_kernel_segmented_segment_sum_runtime_{s}", .{dtype.name()}),
261         .parameters = .{
262             .dst = kernel.dynamicBuffer(dtype),
263             .data = kernel.dynamicBuffer(dtype),
264             .offsets = kernel.dynamicBuffer(.i32),
265             .segments = kernel.scalar(.i32),
266             .total = kernel.scalar(.i32),
267         },
268         .Instance = SegmentSum,
269         .schedule = segmentSumFamilySchedule,
270         .body = segmentSumRuntimeBody,
271     });
272 }
273 
274 pub const SegmentSumFamilyF32 = segmentSumFamily(.f32);
275 pub const SegmentSumFamilyF16 = segmentSumFamily(.f16);
276 pub const SegmentSumRuntimeFamilyF32 = segmentSumRuntimeFamily(.f32);
277 pub const SegmentSumRuntimeFamilyF16 = segmentSumRuntimeFamily(.f16);
278 
279 pub fn segmentSumThreadsForSegments(segments: u64) u32 {
280     return geometry_mod.threadsForExtent(segments, segment_sum_thread_caps);
281 }
282 
283 pub fn segmentSumThreadCandidatesForSegments(segments: u64) geometry_mod.Thread1DCandidates {
284     return geometry_mod.threadCandidatesForExtent(segments, segment_sum_thread_caps);
285 }
286 
287 pub fn segmentSumInstanceTarget(allocator: std.mem.Allocator, instance: SegmentSum) ![]u8 {
288     return std.fmt.allocPrint(
289         allocator,
290         "accy.kernel.segmented.segment_sum{d}x{d}_{s}_{d}_{s}",
291         .{ instance.segments, instance.total, instance.granularity.name(), instance.threads, instance.dtype.name() },
292     );
293 }
294 
295 pub fn segmentSumInstanceEntryName(allocator: std.mem.Allocator, instance: SegmentSum) ![]u8 {
296     return std.fmt.allocPrint(
297         allocator,
298         "accy_kernel_segmented_segment_sum{d}x{d}_{s}_{d}_{s}",
299         .{ instance.segments, instance.total, instance.granularity.name(), instance.threads, instance.dtype.name() },
300     );
301 }
302 
303 pub fn segmentSumFamilyTarget(allocator: std.mem.Allocator, instance: SegmentSum) ![]u8 {
304     return std.fmt.allocPrint(
305         allocator,
306         "accy.kernel.segmented.segment_sum_family_{s}_{d}_{s}",
307         .{ instance.granularity.name(), instance.threads, instance.dtype.name() },
308     );
309 }
310 
311 pub fn segmentSumFamilyEntryName(allocator: std.mem.Allocator, instance: SegmentSum) ![]u8 {
312     return std.fmt.allocPrint(
313         allocator,
314         "accy_kernel_segmented_segment_sum_family_{s}_{d}_{s}",
315         .{ instance.granularity.name(), instance.threads, instance.dtype.name() },
316     );
317 }
318 
319 pub fn segmentSumTuningExtents(instance: SegmentSum) [2]u64 {
320     return .{ instance.segments, instance.total };
321 }
322 
323 pub fn segmentSumTuningOperation(instance: SegmentSum) entry.Operation {
324     _ = instance;
325     return .{ .segmented = .segment_sum };
326 }
327 
328 pub fn segmentSumFamilyTuningKey(
329     backing_allocator: std.mem.Allocator,
330     device_fingerprint: u64,
331     instance: SegmentSum,
332 ) !tuning.FamilyTuningKey {
333     const family_fingerprint = try segmentSumFamilyFingerprint(backing_allocator, instance);
334     const extents = segmentSumTuningExtents(instance);
335     return tuning.FamilyTuningKey.init(
336         device_fingerprint,
337         family_fingerprint,
338         entry.operationFingerprint(segmentSumTuningOperation(instance)),
339         instance.dtype,
340         segment_sum_family_version,
341         extents[0..],
342     ) orelse unreachable;
343 }
344 
345 pub fn segmentSumRuntimeArguments(instance: SegmentSum) ![2]choir_abi.ScalarArgument {
346     return .{
347         .{ .u32 = try runtimeExtentArgument(instance.segments) },
348         .{ .u32 = try runtimeExtentArgument(instance.total) },
349     };
350 }
351 
352 pub fn segmentSumShapeProfileDimensions(instance: SegmentSum) [2]artifact_product.KernelCallShapeProfileDimension {
353     const bounds = segmentSumRuntimeExtentBounds();
354     return .{
355         .{ .name = instance.segment_axis, .runtime_scalar_argument_index = 0, .bounds = bounds },
356         .{ .name = instance.element_axis, .runtime_scalar_argument_index = 1, .bounds = bounds },
357     };
358 }
359 
360 fn segmentSumRuntimeExtentBounds() shape.Bounds {
361     return .{ .min = 1, .max = extent_mod.runtime_extent_max };
362 }
363 
364 fn segmentSumDerivedLaunch(instance: SegmentSum) !artifact_product.KernelCallLaunch {
365     if (instance.threads == 0) return error.KernelLibraryLaunchThreadgroupMustBeNonzero;
366     if (!segmentSumGranularityValid(instance.granularity, instance.threads)) return error.UnsupportedGranularity;
367     const divisor = switch (instance.granularity) {
368         .thread => instance.threads,
369         .warp => instance.threads / segment_sum_warp_size,
370     };
371     return .{ .derived = .{
372         .grid = .{
373             .{ .runtime_u32_ceil_div = .{ .argument_index = 0, .divisor = divisor } },
374             .{ .fixed = 1 },
375             .{ .fixed = 1 },
376         },
377         .threadgroup = .{ instance.threads, 1, 1 },
378     } };
379 }
380 
381 pub fn createSegmentSumFamilyArtifact(
382     allocator: std.mem.Allocator,
383     handle: kernel.BackendHandle,
384     instance: SegmentSum,
385     options: entry.ArtifactOptions,
386 ) !kernel.OwnedKernelCallArtifact {
387     const target = try segmentSumFamilyTarget(allocator, instance);
388     defer allocator.free(target);
389     const entry_name = try segmentSumFamilyEntryName(allocator, instance);
390     defer allocator.free(entry_name);
391     const family_fingerprint = options.shape_family_fingerprint orelse try segmentSumFamilyFingerprint(allocator, instance);
392     const shape_profile_dimensions = segmentSumShapeProfileDimensions(instance);
393     const shape_profile = options.shape_profile orelse artifact_product.KernelCallShapeProfile{
394         .name = "segment_sum",
395         .fingerprint = family_fingerprint,
396         .dimensions = shape_profile_dimensions[0..],
397     };
398 
399     var graph = switch (instance.dtype) {
400         .f32 => try SegmentSumRuntimeFamilyF32.buildNamed(allocator, options.limits, entry_name, instance),
401         .f16 => try SegmentSumRuntimeFamilyF16.buildNamed(allocator, options.limits, entry_name, instance),
402         else => return error.UnsupportedDType,
403     };
404     defer graph.deinit();
405     return kernel.createKernelCallArtifact(allocator, handle, &graph, .{
406         .target = target,
407         .version = segment_sum_family_version,
408         .format = options.format,
409         .kernel_plan = options.kernel_plan,
410         .element_count_argument = options.element_count_argument,
411         .shape_family_fingerprint = family_fingerprint,
412         .shape_profile = shape_profile,
413         .launch = options.launch orelse try segmentSumDerivedLaunch(instance),
414         .runtime_scalar_argument_count = if (options.runtime_scalar_argument_count == 0) 2 else options.runtime_scalar_argument_count,
415         .static_arguments = options.static_arguments,
416     });
417 }
418 
419 pub fn segmentSumFamilyFingerprint(backing_allocator: std.mem.Allocator, instance: SegmentSum) !u64 {
420     var family = try segmentSumShapeFamily(backing_allocator, instance);
421     defer family.deinit();
422     return shape.fingerprint(family);
423 }
424 
425 pub fn segmentSumShapeFamily(backing_allocator: std.mem.Allocator, instance: SegmentSum) !shape.Family {
426     var builder = try shape.Builder.init(backing_allocator, "segment_sum");
427     errdefer builder.deinit();
428 
429     const segments = try builder.symbol(instance.segment_axis);
430     const elements = try builder.symbol(instance.element_axis);
431 
432     const segments_expr = try builder.symbolExpression(segments);
433     const elements_expr = try builder.symbolExpression(elements);
434     const one_expr = builder.constantExpression(1);
435     const offsets_expr = try builder.addExpression(segments_expr, one_expr);
436 
437     _ = try builder.tensor("data", &.{elements_expr});
438     _ = try builder.tensor("offsets", &.{offsets_expr});
439     _ = try builder.tensor("out", &.{segments_expr});
440     try builder.assumeBounds(segments_expr, segmentSumRuntimeExtentBounds());
441     try builder.assumeBounds(elements_expr, segmentSumRuntimeExtentBounds());
442 
443     return builder.finish();
444 }
445 
446 pub fn segmentSumFamilySpecialization(backing_allocator: std.mem.Allocator, instance: SegmentSum) !entry.OwnedSpecialization {
447     var owned = entry.OwnedSpecialization.init(backing_allocator);
448     errdefer owned.deinit();
449     const lifetime_allocator = owned.allocator();
450 
451     const inputs = try lifetime_allocator.alloc(entry.Shape, 2);
452     inputs[0] = try entry.runtimeShape1D(lifetime_allocator, instance.element_axis, instance.total);
453     inputs[1] = try entry.runtimeShape1D(lifetime_allocator, instance.segment_axis, instance.segments + 1);
454 
455     const outputs = try lifetime_allocator.alloc(entry.Shape, 1);
456     outputs[0] = try entry.runtimeShape1D(lifetime_allocator, instance.segment_axis, instance.segments);
457 
458     owned.value = .{
459         .dtype = instance.dtype,
460         .operation = .{ .segmented = .segment_sum },
461         .inputs = inputs,
462         .outputs = outputs,
463         .schedule = try entry.runtimeThreadBlocks1D(lifetime_allocator, instance.segment_axis, instance.launchExtent(), instance.threads),
464     };
465     owned.value.launch = owned.value.schedule.?.launch();
466     var family = try segmentSumShapeFamily(backing_allocator, instance);
467     errdefer family.deinit();
468     try owned.takeShapeFamily(&family);
469     return owned;
470 }
471 
472 fn ceilDivExtent(extent: u64, divisor: u32) u64 {
473     return (extent + divisor - 1) / divisor;
474 }
475 
476 fn launchMatches1D(launch: entry.Launch, extent: u64, threadgroup: u32) bool {
477     if (threadgroup == 0) return false;
478     if (@as(u64, threadgroup) > extent) return false;
479     const expected_grid = ceilDivExtent(extent, threadgroup);
480     return @as(u64, launch.grid[0]) == expected_grid and
481         launch.grid[1] == 1 and launch.grid[2] == 1 and
482         launch.threadgroup[1] == 1 and launch.threadgroup[2] == 1;
483 }
484 
485 pub fn segmentSumGranularityFromLaunch(launch: entry.Launch, segments: u64) ?SegmentSumGranularity {
486     const threadgroup = launch.threadgroup[0];
487     if (launchMatches1D(launch, segments, threadgroup)) return .thread;
488     if (threadgroup % segment_sum_warp_size == 0 and
489         launchMatches1D(launch, segments * segment_sum_warp_size, threadgroup))
490     {
491         return .warp;
492     }
493     return null;
494 }
495 
496 pub fn segmentSumInstanceFromSpecialization(specialization: entry.Specialization) ?SegmentSum {
497     if (!specialization.scheduleMatchesLaunch()) return null;
498     if (!specialization.operationIs(.{ .segmented = .segment_sum })) return null;
499     const dtype = specialization.dtype orelse return null;
500     if (!segmentSumDTypeSupported(dtype)) return null;
501     if (specialization.inputs.len != 2 or specialization.outputs.len != 1) return null;
502     if (specialization.reductions.len != 0) return null;
503     const data = specialization.inputs[0];
504     const offsets = specialization.inputs[1];
505     const output = specialization.outputs[0];
506     if (data.axes.len != 1 or offsets.axes.len != 1 or output.axes.len != 1) return null;
507     const total = data.axes[0].extent;
508     const segments = output.axes[0].extent;
509     if (offsets.axes[0].extent != segments + 1) return null;
510     if (!std.mem.eql(u8, offsets.axes[0].name, output.axes[0].name)) return null;
511     const launch = specialization.launch orelse return null;
512     if (launch.threadgroup[0] == 0) return null;
513     const granularity = segmentSumGranularityFromLaunch(launch, segments) orelse return null;
514     return .{
515         .segments = segments,
516         .total = total,
517         .dtype = dtype,
518         .granularity = granularity,
519         .threads = launch.threadgroup[0],
520         .segment_axis = output.axes[0].name,
521         .element_axis = data.axes[0].name,
522     };
523 }
524 
525 fn segmentSumSpecialization(comptime spec: SegmentSum) entry.Specialization {
526     return .{
527         .dtype = spec.dtype,
528         .operation = .{ .segmented = .segment_sum },
529         .inputs = &.{
530             entry.shape1D(spec.element_axis, spec.total),
531             entry.shape1D(spec.segment_axis, spec.segments + 1),
532         },
533         .outputs = &.{entry.shape1D(spec.segment_axis, spec.segments)},
534         .launch = entry.launch1D(ceilDivComptime(spec.segments, spec.threads), spec.threads),
535         .schedule = entry.threadBlocks1D(spec.segment_axis, spec.segments, spec.threads),
536     };
537 }
538 
539 fn ceilDivComptime(comptime numerator: u64, comptime denominator: u32) u32 {
540     return @intCast(numerator / denominator + @as(u64, @intFromBool(numerator % denominator != 0)));
541 }
542 
543 fn segmentSumProgram(comptime spec: SegmentSum) type {
544     const Body = struct {
545         fn run(k: anytype, args: anytype) !void {
546             try segmentSumBody(k, spec, args);
547         }
548     };
549 
550     return kernel.logical.Program(.{
551         .name = std.fmt.comptimePrint(
552             "accy_kernel_segmented_segment_sum{}x{}_{s}_{}_{s}",
553             .{ spec.segments, spec.total, spec.granularity.name(), spec.threads, spec.dtype.name() },
554         ),
555         .parameters = .{
556             .dst = kernel.dynamicBuffer(spec.dtype),
557             .data = kernel.dynamicBuffer(spec.dtype),
558             .offsets = kernel.dynamicBuffer(.i32),
559         },
560         .body = Body.run,
561     }).withSchedule(kernel.logical.schedule.threadBlocks(.{ .x = spec.threads }));
562 }
563 
564 pub fn segmentSumF32(comptime spec: SegmentSum) type {
565     return entry.Entry(segmentSumProgram(spec), .{
566         .target = std.fmt.comptimePrint(
567             "accy.kernel.segmented.segment_sum{}x{}_{s}_{}_{s}",
568             .{ spec.segments, spec.total, spec.granularity.name(), spec.threads, spec.dtype.name() },
569         ),
570         .layer = .logical,
571         .category = .segmented,
572         .specialization = segmentSumSpecialization(spec),
573     });
574 }
575 
576 pub const SegmentSum4F32 = segmentSumF32(.{ .segments = 4, .total = 16, .threads = 4 });
577 
578 test "segmented segment sum entry runs on CPU with ragged segments" {
579     var data: [16]f32 = undefined;
580     for (&data, 0..) |*value, index| value.* = @floatFromInt(index + 1);
581     var offsets = [_]i32{ 0, 3, 3, 10, 16 };
582     var dst = @as([4]f32, @splat(0));
583 
584     try SegmentSum4F32.runCpu(std.testing.allocator, SegmentSum4F32.Limits.testing, &.{
585         kernel.argumentBuffer(f32, dst[0..]),
586         kernel.argumentBuffer(f32, data[0..]),
587         kernel.argumentBuffer(i32, offsets[0..]),
588     });
589     try std.testing.expectEqualSlices(f32, &.{ 6, 0, 49, 81 }, dst[0..]);
590 }
591 
592 test "segmented segment sum clamps out-of-range offsets" {
593     var data = [_]f32{ 1, 2, 4, 8 };
594     var offsets = [_]i32{ -2, 2, 9, 4 };
595     var dst = @as([3]f32, @splat(0));
596 
597     const Entry3 = segmentSumF32(.{ .segments = 3, .total = 4, .threads = 3 });
598     try Entry3.runCpu(std.testing.allocator, Entry3.Limits.testing, &.{
599         kernel.argumentBuffer(f32, dst[0..]),
600         kernel.argumentBuffer(f32, data[0..]),
601         kernel.argumentBuffer(i32, offsets[0..]),
602     });
603     try std.testing.expectEqualSlices(f32, &.{ 3, 12, 0 }, dst[0..]);
604 }
605 
606 test "segmented segment sum runtime family executes explicit runtime extents" {
607     const allocator = std.testing.allocator;
608     const compiled = SegmentSum{ .segments = 1, .total = 1, .threads = 4 };
609     const runtime = SegmentSum{ .segments = 3, .total = 12, .threads = 4 };
610 
611     var graph = try SegmentSumRuntimeFamilyF32.build(allocator, SegmentSumRuntimeFamilyF32.Limits.testing, compiled);
612     defer graph.deinit();
613 
614     var data: [12]f32 = undefined;
615     for (&data, 0..) |*value, index| value.* = @floatFromInt(index);
616     var offsets = [_]i32{ 0, 5, 5, 12 };
617     var dst = @as([3]f32, @splat(0));
618 
619     var expected = [_]f32{ 0, 0, 0 };
620     for (0..3) |segment| {
621         const begin: usize = @intCast(offsets[segment]);
622         const end: usize = @intCast(offsets[segment + 1]);
623         for (begin..end) |element| expected[segment] += data[element];
624     }
625 
626     const launch_value = try entry.runtimeLaunch1D(runtime.segments, runtime.threads);
627     try graph.runCpuWithLaunch(allocator, &.{
628         kernel.argumentBuffer(f32, dst[0..]),
629         kernel.argumentBuffer(f32, data[0..]),
630         kernel.argumentBuffer(i32, offsets[0..]),
631         kernel.argumentI32(@intCast(runtime.segments)),
632         kernel.argumentI32(@intCast(runtime.total)),
633     }, .{
634         .grid = launch_value.grid,
635         .block = launch_value.threadgroup,
636     });
637     try std.testing.expectEqualSlices(f32, expected[0..], dst[0..]);
638 }
639 
640 test "segmented segment sum warp runtime family matches the thread oracle" {
641     const allocator = std.testing.allocator;
642     const compiled = SegmentSum{ .segments = 1, .total = 1, .granularity = .warp, .threads = 32 };
643     const runtime = SegmentSum{ .segments = 3, .total = 80, .granularity = .warp, .threads = 32 };
644 
645     var graph = try SegmentSumRuntimeFamilyF32.build(allocator, SegmentSumRuntimeFamilyF32.Limits.testing, compiled);
646     defer graph.deinit();
647 
648     var data: [80]f32 = undefined;
649     for (&data, 0..) |*value, index| value.* = @floatFromInt(index);
650     var offsets = [_]i32{ 0, 50, 50, 80 };
651     var dst = @as([3]f32, @splat(0));
652 
653     var expected = [_]f32{ 0, 0, 0 };
654     for (0..3) |segment| {
655         const begin: usize = @intCast(offsets[segment]);
656         const end: usize = @intCast(offsets[segment + 1]);
657         for (begin..end) |element| expected[segment] += data[element];
658     }
659 
660     const launch_value = try entry.runtimeLaunch1D(runtime.launchExtent(), runtime.threads);
661     try graph.runCpuWithLaunch(allocator, &.{
662         kernel.argumentBuffer(f32, dst[0..]),
663         kernel.argumentBuffer(f32, data[0..]),
664         kernel.argumentBuffer(i32, offsets[0..]),
665         kernel.argumentI32(@intCast(runtime.segments)),
666         kernel.argumentI32(@intCast(runtime.total)),
667     }, .{
668         .grid = launch_value.grid,
669         .block = launch_value.threadgroup,
670     });
671     try std.testing.expectEqualSlices(f32, expected[0..], dst[0..]);
672 }
673 
674 test "segmented segment sum warp identity carries the granularity" {
675     const instance = SegmentSum{ .segments = 1024, .total = 1_000_000, .granularity = .warp, .threads = 128 };
676 
677     const family_target = try segmentSumFamilyTarget(std.testing.allocator, instance);
678     defer std.testing.allocator.free(family_target);
679     try std.testing.expectEqualStrings("accy.kernel.segmented.segment_sum_family_warp_128_f32", family_target);
680 
681     const family_entry = try segmentSumFamilyEntryName(std.testing.allocator, instance);
682     defer std.testing.allocator.free(family_entry);
683     try std.testing.expectEqualStrings("accy_kernel_segmented_segment_sum_family_warp_128_f32", family_entry);
684 }
685 
686 test "segmented segment sum warp derived launch divides segments by warps per block" {
687     const allocator = std.testing.allocator;
688     var state = gpu.recording.BackendState{
689         .allocator = allocator,
690         .kind = .cuda,
691         .format = .cuda_ptx,
692     };
693     const instance = SegmentSum{ .segments = 1000, .total = 65536, .granularity = .warp, .threads = 128 };
694 
695     var family_artifact = try createSegmentSumFamilyArtifact(allocator, state.handle(), instance, .{ .limits = .testing });
696     defer family_artifact.deinit();
697 
698     const family_entry = family_artifact.entry();
699     try std.testing.expectEqualStrings("accy.kernel.segmented.segment_sum_family_warp_128_f32", family_entry.target);
700     switch (family_entry.launch) {
701         .derived => |launch| {
702             try std.testing.expectEqual(@as(u32, 128), launch.threadgroup[0]);
703             switch (launch.grid[0]) {
704                 .runtime_u32_ceil_div => |term| {
705                     try std.testing.expectEqual(@as(usize, 0), term.argument_index);
706                     try std.testing.expectEqual(@as(u32, 4), term.divisor);
707                 },
708                 else => return error.TestExpectedDerivedLaunch,
709             }
710         },
711         else => return error.TestExpectedDerivedLaunch,
712     }
713 }
714 
715 test "segmented segment sum warp instance round-trips through specialization" {
716     const instance = SegmentSum{ .segments = 100, .total = 4096, .granularity = .warp, .threads = 64 };
717     var owned = try segmentSumFamilySpecialization(std.testing.allocator, instance);
718     defer owned.deinit();
719 
720     const recovered = segmentSumInstanceFromSpecialization(owned.value) orelse return error.TestExpectedSegmentSumInstance;
721     try std.testing.expectEqual(SegmentSumGranularity.warp, recovered.granularity);
722     try std.testing.expectEqual(instance.segments, recovered.segments);
723     try std.testing.expectEqual(instance.total, recovered.total);
724     try std.testing.expectEqual(instance.threads, recovered.threads);
725 }
726 
727 test "segmented segment sum granularity recovery separates single-segment launches" {
728     const thread_instance = SegmentSum{ .segments = 1, .total = 8, .threads = 1 };
729     var thread_owned = try segmentSumFamilySpecialization(std.testing.allocator, thread_instance);
730     defer thread_owned.deinit();
731     const thread_recovered = segmentSumInstanceFromSpecialization(thread_owned.value) orelse return error.TestExpectedSegmentSumInstance;
732     try std.testing.expectEqual(SegmentSumGranularity.thread, thread_recovered.granularity);
733 
734     const warp_instance = SegmentSum{ .segments = 1, .total = 8, .granularity = .warp, .threads = 32 };
735     var warp_owned = try segmentSumFamilySpecialization(std.testing.allocator, warp_instance);
736     defer warp_owned.deinit();
737     const warp_recovered = segmentSumInstanceFromSpecialization(warp_owned.value) orelse return error.TestExpectedSegmentSumInstance;
738     try std.testing.expectEqual(SegmentSumGranularity.warp, warp_recovered.granularity);
739 }
740 
741 test "segmented segment sum family instance identity matches fixed entry strings" {
742     const instance = SegmentSum{ .segments = 4, .total = 16, .threads = 4 };
743 
744     const target = try segmentSumInstanceTarget(std.testing.allocator, instance);
745     defer std.testing.allocator.free(target);
746     try std.testing.expectEqualStrings(SegmentSum4F32.target, target);
747 
748     const entry_name = try segmentSumInstanceEntryName(std.testing.allocator, instance);
749     defer std.testing.allocator.free(entry_name);
750     try std.testing.expectEqualStrings(SegmentSum4F32.name, entry_name);
751 
752     try std.testing.expectEqual(SegmentSum4F32.version, segment_sum_family_version);
753 
754     const fresh = SegmentSum{ .segments = 1024, .total = 1_000_000, .threads = 128 };
755     const family_target = try segmentSumFamilyTarget(std.testing.allocator, fresh);
756     defer std.testing.allocator.free(family_target);
757     try std.testing.expectEqualStrings("accy.kernel.segmented.segment_sum_family_thread_128_f32", family_target);
758 
759     const family_entry = try segmentSumFamilyEntryName(std.testing.allocator, fresh);
760     defer std.testing.allocator.free(family_entry);
761     try std.testing.expectEqualStrings("accy_kernel_segmented_segment_sum_family_thread_128_f32", family_entry);
762 }
763 
764 test "segmented segment sum family tuning keys group candidate granularity" {
765     const allocator = std.testing.allocator;
766     const device = tuning.deviceFingerprint(.{ .identity = .{
767         .backend = .cuda,
768         .family = .nvidia_cuda,
769         .name = "segmented-family-tuning-test-device",
770         .vendor_id = 0x10de,
771         .device_id = 0x2684,
772     } });
773 
774     const thread = SegmentSum{ .segments = 64, .total = 4096, .granularity = .thread, .threads = 128 };
775     const warp = SegmentSum{ .segments = 64, .total = 4096, .granularity = .warp, .threads = 128 };
776     const thread_key = try segmentSumFamilyTuningKey(allocator, device, thread);
777     const warp_key = try segmentSumFamilyTuningKey(allocator, device, warp);
778     try std.testing.expect(thread_key.eql(warp_key));
779 
780     const other_extent = try segmentSumFamilyTuningKey(allocator, device, .{ .segments = 32, .total = 4096 });
781     try std.testing.expect(!thread_key.eql(other_extent));
782 
783     const other_device = try segmentSumFamilyTuningKey(
784         allocator,
785         tuning.deviceFingerprint(.{ .identity = .{
786             .backend = .cuda,
787             .family = .nvidia_cuda,
788             .name = "other-segmented-family-tuning-test-device",
789             .vendor_id = 0x10de,
790             .device_id = 0x1b80,
791         } }),
792         thread,
793     );
794     try std.testing.expect(!thread_key.eql(other_device));
795     try std.testing.expectEqual(thread_key.family_fingerprint, other_device.family_fingerprint);
796     try std.testing.expectEqual(thread_key.operation_fingerprint, other_device.operation_fingerprint);
797 }
798 
799 test "segmented segment sum family artifact carries runtime launch contract" {
800     const allocator = std.testing.allocator;
801     var state = gpu.recording.BackendState{
802         .allocator = allocator,
803         .kind = .cuda,
804         .format = .cuda_ptx,
805     };
806     const instance = SegmentSum{ .segments = 4, .total = 16, .threads = 4 };
807 
808     var family_artifact = try createSegmentSumFamilyArtifact(allocator, state.handle(), instance, .{ .limits = .testing });
809     defer family_artifact.deinit();
810 
811     const family_entry = family_artifact.entry();
812     try std.testing.expectEqualStrings("accy.kernel.segmented.segment_sum_family_thread_4_f32", family_entry.target);
813     try std.testing.expectEqualStrings("accy_kernel_segmented_segment_sum_family_thread_4_f32", family_entry.entry_name);
814     try std.testing.expectEqual(@as(u32, 5), family_entry.argument_count);
815     try std.testing.expectEqual(@as(u32, 2), family_entry.runtime_scalar_argument_count);
816     try std.testing.expect(family_entry.required_dtypes.contains(.f32));
817     try std.testing.expect(family_entry.required_dtypes.contains(.i32));
818     try std.testing.expect(family_entry.shape_family_fingerprint != null);
819     const profile = family_entry.shape_profile orelse return error.TestExpectedShapeProfile;
820     try std.testing.expectEqualStrings("segment_sum", profile.name);
821     try std.testing.expectEqual(@as(usize, 2), profile.dimensions.len);
822     switch (family_entry.launch) {
823         .derived => |launch| {
824             try std.testing.expectEqual(@as(u32, 4), launch.threadgroup[0]);
825             switch (launch.grid[0]) {
826                 .runtime_u32_ceil_div => |term| {
827                     try std.testing.expectEqual(@as(usize, 0), term.argument_index);
828                     try std.testing.expectEqual(@as(u32, 4), term.divisor);
829                 },
830                 else => return error.TestExpectedDerivedLaunch,
831             }
832         },
833         else => return error.TestExpectedDerivedLaunch,
834     }
835 }
836 
837 test "segmented segment sum instance round-trips through specialization" {
838     const instance = SegmentSum{ .segments = 100, .total = 4096, .threads = 32 };
839     var owned = try segmentSumFamilySpecialization(std.testing.allocator, instance);
840     defer owned.deinit();
841 
842     const recovered = segmentSumInstanceFromSpecialization(owned.value) orelse return error.TestExpectedSegmentSumInstance;
843     try std.testing.expectEqual(instance.segments, recovered.segments);
844     try std.testing.expectEqual(instance.total, recovered.total);
845     try std.testing.expectEqual(instance.dtype, recovered.dtype);
846     try std.testing.expectEqual(instance.threads, recovered.threads);
847 
848     try std.testing.expectEqual(@as(?SegmentSum, null), segmentSumInstanceFromSpecialization(.{}));
849 }
850 
851 test "segmented segment sum thread candidates stay bounded and lead with the default" {
852     const candidates = segmentSumThreadCandidatesForSegments(50_000);
853     try std.testing.expect(candidates.count > 2);
854     try std.testing.expectEqual(segmentSumThreadsForSegments(50_000), candidates.items[0]);
855     for (candidates.slice(), 0..) |candidate, index| {
856         try std.testing.expect(candidate != 0);
857         for (candidates.slice()[0..index]) |previous| try std.testing.expect(previous != candidate);
858     }
859 }