lib/accy/src/kernel/library/histogram/family/runtime/body.zig

daab053ee43316e1809a84551d573ddd1e5bf3d2

  1 const std = @import("std");
  2 const choir_abi = @import("choir_abi");
  3 
  4 const accy = @import("../../../../../root.zig");
  5 const library = @import("../../../root.zig");
  6 const family = @import("../root.zig");
  7 
  8 const entry = library.entry;
  9 const kernel = accy.kernel;
 10 const Histogram = family.Histogram;
 11 const HistogramBinningPolicy = family.HistogramBinningPolicy;
 12 const DType = choir_abi.DType;
 13 
 14 fn histogramFamilySchedule(instance: Histogram) kernel.logical.schedule.ThreadBlocks {
 15     return kernel.logical.schedule.threadBlocks(.{ .x = instance.threads });
 16 }
 17 
 18 fn histogramBinIndex(b: anytype, value: kernel.Value, lo: kernel.Value, width: kernel.Value) !kernel.Value {
 19     const offset = try b.sub(value, lo);
 20     const ratio = try b.div(offset, width);
 21     const bin_i32 = try b.cast(ratio, .i32);
 22     return b.castIndex(bin_i32);
 23 }
 24 
 25 fn histogramUpperBound(b: anytype, bins: kernel.Value, lo: kernel.Value, width: kernel.Value, policy: HistogramBinningPolicy) !kernel.Value {
 26     return switch (policy) {
 27         .lower_inclusive_upper_exclusive => {
 28             const bins_f32 = try b.cast(bins, .f32);
 29             const span = try b.mul(bins_f32, width);
 30             return b.add(lo, span);
 31         },
 32     };
 33 }
 34 
 35 fn histogramAccumulate(b: anytype, ctx: anytype, target: kernel.Value) !void {
 36     const one = try b.constantInt(.i32, 1);
 37     _ = try b.atomicRmwIndex(.add, one, ctx.dst, target);
 38 }
 39 
 40 fn histogram_element_body_in_low(inner: anytype, low_ctx: anytype) !void {
 41     const width = low_ctx.args.param(.width).raw();
 42     const upper = try histogramUpperBound(inner, low_ctx.bins, low_ctx.lo, width, low_ctx.binning);
 43     const in_high = try inner.compare(.lt, low_ctx.value, upper);
 44     try inner.guardDo(in_high, .{
 45         .args = low_ctx.args,
 46         .value = low_ctx.value,
 47         .lo = low_ctx.lo,
 48         .width = width,
 49         .bins = low_ctx.bins,
 50         .dst = low_ctx.dst,
 51     }, histogram_element_body_in_high);
 52 }
 53 
 54 fn histogram_element_body_in_high(range_builder: anytype, range_ctx: anytype) !void {
 55     const target = try histogramBinIndex(range_builder, range_ctx.value, range_ctx.lo, range_ctx.width);
 56     const in_range = try range_builder.compare(.lt, target, range_ctx.bins);
 57     try range_builder.guardDo(in_range, .{ .dst = range_ctx.dst, .target = target }, histogram_element_body_in_range);
 58 }
 59 
 60 fn histogram_element_body_in_range(atomic_builder: anytype, atomic_ctx: anytype) !void {
 61     try histogramAccumulate(atomic_builder, atomic_ctx, atomic_ctx.target);
 62 }
 63 
 64 fn histogramElementBody(b: anytype, ctx: anytype) !void {
 65     const loaded = try ctx.args.param(.data).load(b, ctx.element);
 66     const value = loaded.raw();
 67     const lo = ctx.args.param(.lo).raw();
 68     const in_low = try b.compare(.ge, value, lo);
 69     try b.guardDo(in_low, .{
 70         .args = ctx.args,
 71         .value = value,
 72         .lo = lo,
 73         .bins = ctx.bins,
 74         .dst = ctx.dst,
 75         .binning = ctx.binning,
 76     }, histogram_element_body_in_low);
 77 }
 78 
 79 fn histogram_runtime_body_accumulate(guard_builder: anytype, ctx: anytype) !void {
 80     try histogramElementBody(guard_builder, ctx);
 81 }
 82 
 83 fn histogramDirectRuntimeBody(k: anytype, spec: Histogram, args: anytype) !void {
 84     const element = try k.globalId(.x);
 85     const count = try k.castIndex(args.param(.count).raw());
 86     const bins = try k.castIndex(args.param(.bins).raw());
 87     const active = try k.compare(.lt, element, count);
 88     try k.guardDo(active, .{
 89         .args = args,
 90         .element = element,
 91         .bins = bins,
 92         .dst = args.param(.dst).raw(),
 93         .binning = spec.binning,
 94     }, histogram_runtime_body_accumulate);
 95 }
 96 
 97 fn histogram_shared_runtime_body_zero_bin(loop_builder: anytype, bin: kernel.Value, acc: kernel.Value, ctx: anytype) !kernel.Value {
 98     try loop_builder.storeIndex(ctx.zero_count, ctx.shared_bins, bin);
 99     return acc;
100 }
101 
102 fn histogram_shared_runtime_body_merge(loop_builder: anytype, bin: kernel.Value, acc: kernel.Value, ctx: anytype) !kernel.Value {
103     const partial = try loop_builder.loadIndex(ctx.shared_bins, bin);
104     _ = try loop_builder.atomicRmwIndex(.add, partial, ctx.args.param(.dst).raw(), bin);
105     return acc;
106 }
107 
108 fn histogramSharedRuntimeBody(k: anytype, spec: Histogram, args: anytype) !void {
109     const shared_bins = try k.sharedBuffer(.i32, spec.bins);
110     const zero_count = try k.constantInt(.i32, 0);
111     const thread = try k.castIndex(try k.threadId(.x));
112     const stride = try k.castIndex(try k.blockDim(.x));
113     const bins = try k.castIndex(args.param(.bins).raw());
114 
115     _ = try k.fold(thread, bins, stride, zero_count, .{
116         .shared_bins = shared_bins,
117         .zero_count = zero_count,
118     }, histogram_shared_runtime_body_zero_bin);
119     try k.barrier(.block);
120 
121     const element = try k.globalId(.x);
122     const count = try k.castIndex(args.param(.count).raw());
123     const active = try k.compare(.lt, element, count);
124     try k.guardDo(active, .{
125         .args = args,
126         .element = element,
127         .bins = bins,
128         .dst = shared_bins,
129         .binning = spec.binning,
130     }, histogram_runtime_body_accumulate);
131     try k.barrier(.block);
132 
133     _ = try k.fold(thread, bins, stride, zero_count, .{
134         .args = args,
135         .shared_bins = shared_bins,
136     }, histogram_shared_runtime_body_merge);
137 }
138 
139 fn histogramRuntimeBody(k: anytype, spec: Histogram, args: anytype) !void {
140     switch (spec.variant) {
141         .direct => try histogramDirectRuntimeBody(k, spec, args),
142         .shared_bins => try histogramSharedRuntimeBody(k, spec, args),
143     }
144 }
145 
146 fn histogramRuntimeFamily(comptime data_dtype: DType) type {
147     return kernel.logical.Family(.{
148         .name = std.fmt.comptimePrint("accy_kernel_histogram_runtime_{s}", .{data_dtype.name()}),
149         .parameters = .{
150             .dst = kernel.dynamicBuffer(.i32),
151             .data = kernel.dynamicBuffer(data_dtype),
152             .bins = kernel.scalar(.i32),
153             .count = kernel.scalar(.i32),
154             .lo = kernel.scalar(.f32),
155             .width = kernel.scalar(.f32),
156         },
157         .Instance = Histogram,
158         .schedule = histogramFamilySchedule,
159         .body = histogramRuntimeBody,
160     });
161 }
162 
163 pub const HistogramRuntimeFamilyF32 = histogramRuntimeFamily(.f32);
164 
165 test "histogram runtime family bins values exactly on the oracle" {
166     const allocator = std.testing.allocator;
167     const compiled = Histogram{ .bins = 8, .count = 1, .variant = .shared_bins, .threads = 8 };
168     const runtime = Histogram{ .bins = 8, .count = 20, .lo = 0, .width = 0.5, .variant = .shared_bins, .threads = 8 };
169 
170     var graph = try HistogramRuntimeFamilyF32.build(allocator, HistogramRuntimeFamilyF32.Limits.testing, compiled);
171     defer graph.deinit();
172 
173     var dst = @as([8]i32, @splat(0));
174     var data: [20]f32 = undefined;
175     for (&data, 0..) |*value, element| {
176         value.* = -0.4 + @as(f32, @floatFromInt(element)) * 0.25;
177     }
178 
179     var expected = @as([8]i32, @splat(0));
180     for (data) |value| {
181         if (family.histogramBinForValue(runtime, value)) |bin| expected[bin] += 1;
182     }
183 
184     const launch_value = try entry.runtimeLaunch1D(runtime.count, runtime.threads);
185     try graph.runCpuWithLaunch(allocator, &.{
186         kernel.argumentBuffer(i32, dst[0..]),
187         kernel.argumentBuffer(f32, data[0..]),
188         kernel.argumentI32(@intCast(runtime.bins)),
189         kernel.argumentI32(@intCast(runtime.count)),
190         .{ .scalar = .{ .f32 = runtime.lo } },
191         .{ .scalar = .{ .f32 = runtime.width } },
192     }, .{
193         .grid = launch_value.grid,
194         .block = launch_value.threadgroup,
195     });
196     try std.testing.expectEqualSlices(i32, expected[0..], dst[0..]);
197 }
198 
199 test "histogram binning policy drops closed upper edge" {
200     const allocator = std.testing.allocator;
201     const instance = Histogram{ .bins = 4, .count = 6, .lo = 1.0, .width = 0.5, .variant = .direct, .threads = 4 };
202     var graph = try HistogramRuntimeFamilyF32.build(allocator, HistogramRuntimeFamilyF32.Limits.testing, instance);
203     defer graph.deinit();
204 
205     var dst = @as([4]i32, @splat(0));
206     var data = [_]f32{ 0.99, 1.0, 1.49, 1.5, 2.99, 3.0 };
207     var expected = @as([4]i32, @splat(0));
208     for (data) |value| {
209         if (family.histogramBinForValue(instance, value)) |bin| expected[bin] += 1;
210     }
211 
212     const launch_value = try entry.runtimeLaunch1D(instance.count, instance.threads);
213     try graph.runCpuWithLaunch(allocator, &.{
214         kernel.argumentBuffer(i32, dst[0..]),
215         kernel.argumentBuffer(f32, data[0..]),
216         kernel.argumentI32(@intCast(instance.bins)),
217         kernel.argumentI32(@intCast(instance.count)),
218         .{ .scalar = .{ .f32 = instance.lo } },
219         .{ .scalar = .{ .f32 = instance.width } },
220     }, .{
221         .grid = launch_value.grid,
222         .block = launch_value.threadgroup,
223     });
224     try std.testing.expectEqualSlices(i32, expected[0..], dst[0..]);
225     try std.testing.expectEqual(@as(i32, 1), dst[3]);
226 }