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 }