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 }