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 }