lib/accy/src/kernel/model/program/builder/surface/domain.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2
3 const core = @import("../../../core/root.zig");
4 const builder_schedule = @import("schedule.zig");
5 const builder = core.builder;
6 const domain_mod = core.domain;
7 const schedule_mod = core.schedule;
8
9 pub const Index1D = domain_mod.Index1D;
10 pub const VectorIndex1D = domain_mod.VectorIndex1D;
11 pub const DomainAxis = domain_mod.DomainAxis;
12 pub const Domain2D = domain_mod.Domain2D;
13 pub const Domain3D = domain_mod.Domain3D;
14 pub const Index2D = domain_mod.Index2D;
15 pub const Index3D = domain_mod.Index3D;
16
17 pub fn index1D(owner: anytype, name: []const u8, extent: u64, threads_per_block: u32) !Index1D {
18 return indexDimension(owner, domain_mod.domainAxis(name, extent, threads_per_block), .x);
19 }
20
21 pub fn vectorIndex1D(owner: anytype, name: []const u8, extent: u64, width: u32, threads_per_block: u32) !VectorIndex1D {
22 if (extent == 0 or width == 0 or threads_per_block == 0) return error.InvalidFactor;
23 const lane_count: u64 = @intCast(width);
24 if (extent % lane_count != 0) return error.ExtentNotDivisible;
25 const packet_extent = extent / lane_count;
26 const thread_packets: u64 = @intCast(threads_per_block);
27 const logical_thread_span = std.math.mul(u64, thread_packets, lane_count) catch return error.LaunchDimensionOverflow;
28 const axis_id = try owner.axis(name, extent);
29 var tiled_axis: ?schedule_mod.Split = null;
30 if (packet_extent <= thread_packets) {
31 try owner.vectorize(axis_id, width);
32 try owner.bind(axis_id, threadTarget(.x));
33 } else {
34 const tiled = try builder_schedule.tile(owner, axis_id, logical_thread_span);
35 tiled_axis = tiled;
36 try owner.vectorize(tiled.inner, width);
37 try owner.bind(tiled.outer, blockTarget(.x));
38 try owner.bind(tiled.inner, threadTarget(.x));
39 }
40
41 const packet = try owner.globalId(.x);
42 const width_value = try constantIndexFromU64(owner, lane_count);
43 const base = try owner.mul(packet, width_value);
44 const packet_bound = try constantIndexFromU64(owner, packet_extent);
45 const bound = try constantIndexFromU64(owner, extent);
46 return .{
47 .domain = axis_id,
48 .tile = tiled_axis,
49 .packet = packet,
50 .packet_bound = packet_bound,
51 .base = base,
52 .bound = bound,
53 .extent = extent,
54 .width = width,
55 .packets = packet_extent,
56 };
57 }
58
59 pub fn index2D(owner: anytype, domain: Domain2D) !Index2D {
60 return .{
61 .x = try indexDimension(owner, domain.x, .x),
62 .y = try indexDimension(owner, domain.y, .y),
63 };
64 }
65
66 pub fn index3D(owner: anytype, domain: Domain3D) !Index3D {
67 return .{
68 .x = try indexDimension(owner, domain.x, .x),
69 .y = try indexDimension(owner, domain.y, .y),
70 .z = try indexDimension(owner, domain.z, .z),
71 };
72 }
73
74 pub fn guardIndexDo(owner: anytype, index: Index1D, context: anytype, comptime body: anytype) !void {
75 var active_guard = try owner.guardIndex(index);
76 errdefer active_guard.abort();
77 try body(owner, index, context);
78 try active_guard.leave();
79 }
80
81 pub fn guardVectorIndexDo(owner: anytype, index: VectorIndex1D, context: anytype, comptime body: anytype) !void {
82 const in_bounds = try owner.compare(.lt, index.packet, index.packet_bound);
83 var active_guard = try owner.guard(in_bounds);
84 errdefer active_guard.abort();
85 try body(owner, index, context);
86 try active_guard.leave();
87 }
88
89 pub fn guardIndex2DDo(owner: anytype, index: Index2D, context: anytype, comptime body: anytype) !void {
90 var x_guard = try owner.guardIndex(index.x);
91 errdefer x_guard.abort();
92 var y_guard = try owner.guardIndex(index.y);
93 errdefer y_guard.abort();
94 try body(owner, index, context);
95 try y_guard.leave();
96 try x_guard.leave();
97 }
98
99 pub fn guardIndex3DDo(owner: anytype, index: Index3D, context: anytype, comptime body: anytype) !void {
100 var x_guard = try owner.guardIndex(index.x);
101 errdefer x_guard.abort();
102 var y_guard = try owner.guardIndex(index.y);
103 errdefer y_guard.abort();
104 var z_guard = try owner.guardIndex(index.z);
105 errdefer z_guard.abort();
106 try body(owner, index, context);
107 try z_guard.leave();
108 try y_guard.leave();
109 try x_guard.leave();
110 }
111
112 pub fn forEach1D(owner: anytype, name: []const u8, extent: u64, threads_per_block: u32, context: anytype, comptime body: anytype) !Index1D {
113 const index = try index1D(owner, name, extent, threads_per_block);
114 try guardIndexDo(owner, index, context, body);
115 return index;
116 }
117
118 pub fn forEachVector1D(owner: anytype, name: []const u8, extent: u64, width: u32, threads_per_block: u32, context: anytype, comptime body: anytype) !VectorIndex1D {
119 const index = try vectorIndex1D(owner, name, extent, width, threads_per_block);
120 try guardVectorIndexDo(owner, index, context, body);
121 return index;
122 }
123
124 pub fn forEach2D(owner: anytype, domain: Domain2D, context: anytype, comptime body: anytype) !Index2D {
125 const index = try index2D(owner, domain);
126 try guardIndex2DDo(owner, index, context, body);
127 return index;
128 }
129
130 pub fn forEach3D(owner: anytype, domain: Domain3D, context: anytype, comptime body: anytype) !Index3D {
131 const index = try index3D(owner, domain);
132 try guardIndex3DDo(owner, index, context, body);
133 return index;
134 }
135
136 fn indexDimension(owner: anytype, domain_axis: DomainAxis, dimension: builder.Dimension) !Index1D {
137 if (domain_axis.extent == 0 or domain_axis.threads_per_block == 0) return error.InvalidFactor;
138 const axis_id = try owner.axis(domain_axis.name, domain_axis.extent);
139 var tiled_axis: ?schedule_mod.Split = null;
140 if (domain_axis.extent <= domain_axis.threads_per_block) {
141 try owner.bind(axis_id, threadTarget(dimension));
142 } else {
143 const tiled = try builder_schedule.tile(owner, axis_id, domain_axis.threads_per_block);
144 tiled_axis = tiled;
145 try owner.bind(tiled.outer, blockTarget(dimension));
146 try owner.bind(tiled.inner, threadTarget(dimension));
147 }
148
149 const index = try owner.globalId(dimension);
150 const bound_value = try owner.constantIndex(std.math.cast(i64, domain_axis.extent) orelse return error.LaunchDimensionOverflow);
151 return .{
152 .domain = axis_id,
153 .tile = tiled_axis,
154 .index = index,
155 .bound = bound_value,
156 .extent = domain_axis.extent,
157 };
158 }
159
160 fn constantIndexFromU64(owner: anytype, value: u64) !builder.Value {
161 return owner.constantIndex(std.math.cast(i64, value) orelse return error.LaunchDimensionOverflow);
162 }
163
164 fn blockTarget(dimension: builder.Dimension) schedule_mod.BindTarget {
165 return switch (dimension) {
166 .x => .block_x,
167 .y => .block_y,
168 .z => .block_z,
169 };
170 }
171
172 fn threadTarget(dimension: builder.Dimension) schedule_mod.BindTarget {
173 return switch (dimension) {
174 .x => .thread_x,
175 .y => .thread_y,
176 .z => .thread_z,
177 };
178 }