lib/accy/src/kernel/library/histogram/family/tuning.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2
3 const library = @import("../../root.zig");
4 const model = @import("model.zig");
5 const naming = @import("naming.zig");
6 const shape = @import("shape.zig");
7
8 const entry = library.entry;
9 const geometry = library.geometry;
10 const tuning = library.tuning;
11 const Histogram = model.Histogram;
12 const HistogramVariant = model.HistogramVariant;
13
14 const histogram_thread_caps = geometry.ThreadCaps1D{};
15
16 pub const HistogramResolvedSchedule = struct {
17 variant: HistogramVariant,
18 threads: u32,
19 };
20
21 pub fn histogramTuningExtents(instance: Histogram) [2]u64 {
22 return .{ instance.bins, instance.count };
23 }
24
25 pub fn histogramTuningOperation(instance: Histogram) entry.Operation {
26 _ = instance;
27 return .{ .indexing = .histogram };
28 }
29
30 pub fn histogramFamilyTuningKey(
31 backing_allocator: std.mem.Allocator,
32 device_fingerprint: u64,
33 instance: Histogram,
34 ) !tuning.FamilyTuningKey {
35 const family_fingerprint = try shape.histogramFamilyFingerprint(backing_allocator, instance);
36 const extents = histogramTuningExtents(instance);
37 return tuning.FamilyTuningKey.init(
38 device_fingerprint,
39 family_fingerprint,
40 entry.operationFingerprint(histogramTuningOperation(instance)),
41 instance.dtype,
42 model.histogram_family_version,
43 extents[0..],
44 ) orelse unreachable;
45 }
46
47 pub fn resolveHistogramSchedule(
48 backing_allocator: std.mem.Allocator,
49 reader: tuning.FamilyTuningReader,
50 instance: Histogram,
51 ) !?HistogramResolvedSchedule {
52 const key = try histogramFamilyTuningKey(backing_allocator, reader.device_fingerprint, instance);
53 const record = reader.table.find(key) orelse return null;
54 const thread_candidates = histogramThreadCandidatesForCount(instance.count);
55 const variants = [_]HistogramVariant{ .direct, .shared_bins };
56 for (variants) |variant| {
57 for (thread_candidates.slice()) |threads| {
58 var candidate = instance;
59 candidate.variant = variant;
60 candidate.threads = threads;
61 if (variant == .shared_bins and candidate.bins > model.histogram_shared_bins_cap) continue;
62 const target = try naming.histogramFamilyTarget(backing_allocator, candidate);
63 defer backing_allocator.free(target);
64 if (std.mem.eql(u8, target, record.target)) {
65 return .{ .variant = variant, .threads = threads };
66 }
67 }
68 }
69 return null;
70 }
71
72 pub fn histogramThreadsForCount(count: u64) u32 {
73 return geometry.threadsForExtent(count, histogram_thread_caps);
74 }
75
76 pub fn histogramThreadCandidatesForCount(count: u64) geometry.Thread1DCandidates {
77 return geometry.threadCandidatesForExtent(count, histogram_thread_caps);
78 }
79
80 test "histogram tuning resolves schedule winners through family targets" {
81 const allocator = std.testing.allocator;
82 const device: u64 = 0xfeed_dead_beef_2684;
83 const probe = Histogram{ .bins = 64, .count = 8192 };
84
85 const histogram_key = try histogramFamilyTuningKey(allocator, device, probe);
86 try std.testing.expectEqual(entry.operationFingerprint(.{ .indexing = .histogram }), histogram_key.operation_fingerprint);
87
88 const thread_candidates = histogramThreadCandidatesForCount(probe.count);
89 var winner = probe;
90 winner.variant = .shared_bins;
91 winner.threads = thread_candidates.slice()[0];
92 const winner_target = try naming.histogramFamilyTarget(allocator, winner);
93 defer allocator.free(winner_target);
94
95 const records = [_]tuning.FamilyTuningRecord{.{
96 .key = histogram_key,
97 .target = winner_target,
98 .winner_median_ns = 300,
99 .runner_up_median_ns = 700,
100 .sample_count = 30,
101 }};
102 const reader = tuning.FamilyTuningReader{
103 .device_fingerprint = device,
104 .table = .{ .records = records[0..] },
105 };
106
107 const resolved = (try resolveHistogramSchedule(allocator, reader, probe)) orelse
108 return error.TestExpectedSchedule;
109 try std.testing.expectEqual(HistogramVariant.shared_bins, resolved.variant);
110 try std.testing.expectEqual(winner.threads, resolved.threads);
111 }