lib/accy/src/kernel/library/histogram/family/specialization.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2
3 const library = @import("../../root.zig");
4 const model = @import("model.zig");
5 const shape = @import("shape.zig");
6
7 const entry = library.entry;
8 const Histogram = model.Histogram;
9 const HistogramBinningPolicy = model.HistogramBinningPolicy;
10 const HistogramVariant = model.HistogramVariant;
11
12 pub fn histogramFamilySpecialization(backing_allocator: std.mem.Allocator, instance: Histogram) !entry.OwnedSpecialization {
13 var owned = entry.OwnedSpecialization.init(backing_allocator);
14 errdefer owned.deinit();
15 const lifetime_allocator = owned.allocator();
16
17 const inputs = try lifetime_allocator.alloc(entry.Shape, 1);
18 inputs[0] = try entry.runtimeShape1D(lifetime_allocator, instance.element_axis, instance.count);
19
20 const outputs = try lifetime_allocator.alloc(entry.Shape, 1);
21 outputs[0] = try entry.runtimeShape1D(lifetime_allocator, instance.bin_axis, instance.bins);
22
23 const static_parameters = try lifetime_allocator.alloc(entry.StaticParameter, 1);
24 static_parameters[0] = try entry.runtimeStaticParameter(
25 lifetime_allocator,
26 model.histogram_binning_policy_parameter,
27 model.histogramBinningPolicyCode(instance.binning),
28 );
29
30 owned.value = .{
31 .dtype = instance.dtype,
32 .operation = .{ .indexing = .histogram },
33 .inputs = inputs,
34 .outputs = outputs,
35 .static_parameters = static_parameters,
36 .schedule = try entry.runtimeThreadBlocks1D(lifetime_allocator, instance.element_axis, instance.count, instance.threads),
37 .structure = @tagName(instance.variant),
38 };
39 owned.value.launch = owned.value.schedule.?.launch();
40 var family = try shape.histogramShapeFamily(backing_allocator, instance);
41 errdefer family.deinit();
42 try owned.takeShapeFamily(&family);
43 return owned;
44 }
45
46 pub fn histogramInstanceFromSpecialization(specialization: entry.Specialization) ?Histogram {
47 if (!specialization.scheduleMatchesLaunch()) return null;
48 if (!specialization.operationIs(.{ .indexing = .histogram })) return null;
49 const dtype = specialization.dtype orelse return null;
50 if (!model.histogramDTypeSupported(dtype)) return null;
51 if (specialization.inputs.len != 1 or specialization.outputs.len != 1) return null;
52 if (specialization.reductions.len != 0) return null;
53 if (specialization.static_parameters.len != 1) return null;
54 const binning = model.histogramBinningPolicyFromCode(
55 specialization.staticParameterValue(model.histogram_binning_policy_parameter) orelse return null,
56 ) orelse return null;
57 const data = specialization.inputs[0];
58 const bins_shape = specialization.outputs[0];
59 if (data.axes.len != 1 or bins_shape.axes.len != 1) return null;
60 const launch = specialization.launch orelse return null;
61 if (launch.threadgroup[0] == 0) return null;
62 const structure = specialization.structure orelse return null;
63 const variant = std.meta.stringToEnum(HistogramVariant, structure) orelse return null;
64 const instance = Histogram{
65 .bins = bins_shape.axes[0].extent,
66 .count = data.axes[0].extent,
67 .dtype = dtype,
68 .binning = binning,
69 .variant = variant,
70 .threads = launch.threadgroup[0],
71 .bin_axis = bins_shape.axes[0].name,
72 .element_axis = data.axes[0].name,
73 };
74 if (!model.histogramInstanceValid(instance)) return null;
75 return instance;
76 }
77
78 test "histogram instance round-trips through its specialization" {
79 var owned = try histogramFamilySpecialization(std.testing.allocator, .{
80 .bins = 64,
81 .count = 4096,
82 .variant = .shared_bins,
83 .threads = 64,
84 });
85 defer owned.deinit();
86 const recovered = histogramInstanceFromSpecialization(owned.value) orelse {
87 return error.TestExpectedHistogramInstance;
88 };
89 try std.testing.expectEqual(@as(u64, 64), recovered.bins);
90 try std.testing.expectEqual(@as(u64, 4096), recovered.count);
91 try std.testing.expectEqual(HistogramBinningPolicy.lower_inclusive_upper_exclusive, recovered.binning);
92 try std.testing.expectEqual(HistogramVariant.shared_bins, recovered.variant);
93 try std.testing.expectEqual(@as(u32, 64), recovered.threads);
94 try std.testing.expectEqual(@as(?Histogram, null), histogramInstanceFromSpecialization(.{}));
95 }