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 }