lib/accy/src/choir/record/dispatch.zig

daab053ee43316e1809a84551d573ddd1e5bf3d2

  1 const std = @import("std");
  2 const choir_abi = @import("choir_abi");
  3 const reference = @import("root.zig").reference;
  4 
  5 pub const ScheduleResourceEstimate = struct {
  6     element_count: u64,
  7     element_size: u64,
  8     op_count: usize,
  9     external_input_value_count: usize = 0,
 10     external_operand_count: usize = 0,
 11     chain_operand_count: usize = 0,
 12     static_read_bytes: u64 = 0,
 13     static_write_bytes: u64 = 0,
 14     static_total_bytes: u64 = 0,
 15     estimated_element_ops: u64 = 0,
 16     static_bytes_complete: bool = true,
 17 
 18     pub fn elementOpsPerKiB(self: ScheduleResourceEstimate) u64 {
 19         if (self.static_total_bytes == 0) return 0;
 20         return (std.math.mul(u64, self.estimated_element_ops, 1024) catch
 21             std.math.maxInt(u64)) / self.static_total_bytes;
 22     }
 23 };
 24 
 25 pub const ScheduleWorkKind = enum {
 26     elementwise_single,
 27     elementwise_fusion,
 28     shape,
 29     dot_general,
 30     reduction,
 31     kernel_call,
 32     row_pipeline,
 33     iterate,
 34     flash_attention,
 35     scan,
 36 };
 37 
 38 pub const FusionClusterKind = enum {
 39     elementwise,
 40     dot_epilogue,
 41     reduction_input,
 42     row_pipeline,
 43     flash_attention,
 44 };
 45 
 46 pub const Cluster = struct {
 47     ops: []const reference.Operation,
 48     kind: FusionClusterKind,
 49 };
 50 
 51 pub const Fusion = struct {
 52     clusters: []const Cluster,
 53     elided: []const reference.Operation,
 54     fused_op_count: usize,
 55     max_cluster_len: usize,
 56 };
 57 
 58 pub const WorkItem = struct {
 59     id: usize,
 60     kind: ScheduleWorkKind,
 61     root: reference.Operation,
 62     ops: []const reference.Operation,
 63     output_value: reference.Value,
 64     dtype: choir_abi.DType,
 65     rank: usize,
 66     element_count: u64,
 67     resources: ScheduleResourceEstimate,
 68 };
 69 
 70 pub const Schedule = struct {
 71     work_items: []const WorkItem,
 72     single_work_count: usize,
 73     fusion_work_count: usize,
 74     kernel_call_work_count: usize,
 75     scheduled_op_count: usize,
 76     total_static_elements: u64,
 77 };
 78 
 79 pub const Record = struct { fusion: Fusion, schedule: Schedule };
 80 
 81 pub fn validate(allocator: std.mem.Allocator, value: Record) !void {
 82     const seen = try allocator.alloc(bool, value.schedule.work_items.len);
 83     defer allocator.free(seen);
 84     @memset(seen, false);
 85     for (value.schedule.work_items) |item| {
 86         if (item.id >= seen.len) return error.InvalidStageRecord;
 87         if (seen[item.id]) return error.InvalidStageRecord;
 88         seen[item.id] = true;
 89     }
 90     try validateCounters(value);
 91 }
 92 
 93 fn validateCounters(value: Record) !void {
 94     var fused: usize = 0;
 95     var longest: usize = 0;
 96     for (value.fusion.clusters) |cluster| {
 97         fused = std.math.add(usize, fused, cluster.ops.len) catch return error.InvalidStageRecord;
 98         longest = @max(longest, cluster.ops.len);
 99     }
100     if (fused != value.fusion.fused_op_count or longest != value.fusion.max_cluster_len) {
101         return error.InvalidStageRecord;
102     }
103     var single: usize = 0;
104     var fusion: usize = 0;
105     var calls: usize = 0;
106     var ops: usize = 0;
107     var elements: u64 = 0;
108     for (value.schedule.work_items) |item| {
109         switch (item.kind) {
110             .elementwise_single, .shape, .dot_general, .reduction, .iterate, .scan => single += 1,
111             .elementwise_fusion, .row_pipeline, .flash_attention => fusion += 1,
112             .kernel_call => calls += 1,
113         }
114         ops = std.math.add(usize, ops, item.ops.len) catch return error.InvalidStageRecord;
115         elements = std.math.add(u64, elements, item.element_count) catch return error.InvalidStageRecord;
116     }
117     const schedule = value.schedule;
118     if (single != schedule.single_work_count or fusion != schedule.fusion_work_count or
119         calls != schedule.kernel_call_work_count or ops != schedule.scheduled_op_count or
120         elements != schedule.total_static_elements) return error.InvalidStageRecord;
121 }