lib/accy/src/kernel/model/program/builder/surface/builder.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const choir = @import("choir");
3
4 const core = @import("../../../core/root.zig");
5 const graph = @import("../../graph/root.zig");
6 const guard_mod = @import("../../guard/root.zig");
7
8 const builder = core.builder;
9 const builder_typed = @import("../typed/root.zig");
10 const domain_mod = core.domain;
11 const schedule_mod = core.schedule;
12 const surface_domain = @import("domain.zig");
13 const surface_flow = @import("flow.zig");
14 const surface_guard = @import("guard.zig");
15 const surface_primitive = @import("primitive.zig");
16 const surface_schedule = @import("schedule.zig");
17
18 const ir = choir.ir;
19
20 /// A caller instantiates this once per kernel program type to get the typed builder its kernels are
21 /// written with. The function returns a type that pairs a kernel program type with the dialects,
22 /// named groups of compiler operations, that each of its own compiler contexts, the objects that
23 /// own the dialects, types and operations of one compiler module, loads at setup. The shared kernel
24 /// program front end passes `.{}`, and the kernel execution path passes the target dialect
25 /// registrations and preloads `scf`. A builder that borrows a context the caller already owns skips
26 /// those requirements, so the owner prepares that context before building or freezing it.
27 pub fn Surface(
28 comptime ProgramType: type,
29 comptime context_requirements: builder.ContextRequirements,
30 ) type {
31 return struct {
32 pub const Index1D = domain_mod.Index1D;
33 pub const VectorIndex1D = domain_mod.VectorIndex1D;
34 pub const DomainAxis = domain_mod.DomainAxis;
35 pub const Domain2D = domain_mod.Domain2D;
36 pub const Domain3D = domain_mod.Domain3D;
37 pub const Index2D = domain_mod.Index2D;
38 pub const Index3D = domain_mod.Index3D;
39 pub const Vec2 = builder_typed.Vec2;
40 pub const Vec3 = builder_typed.Vec3;
41 pub const TypedValue = builder_typed.TypedValue;
42 pub const TypedVec2 = builder_typed.TypedVec2;
43 pub const TypedVec3 = builder_typed.TypedVec3;
44 pub const domainAxis = domain_mod.domainAxis;
45 pub const BufferView = builder_typed.BufferView;
46 pub const Program = ProgramType;
47 pub const Limits = core.Limits;
48 pub const Capacity = Builder.Capacity;
49
50 pub fn define(
51 allocator: std.mem.Allocator,
52 limits: Limits,
53 kernel_name: []const u8,
54 params: []const builder.Param,
55 context: anytype,
56 comptime body: anytype,
57 ) !Program {
58 var program_builder = try Builder.init(allocator, limits, kernel_name, params);
59 errdefer program_builder.deinit();
60 try body(&program_builder, context);
61 try program_builder.return_();
62 return program_builder.finish();
63 }
64
65 pub fn defineWithDiagnostic(
66 allocator: std.mem.Allocator,
67 limits: Limits,
68 kernel_name: []const u8,
69 params: []const builder.Param,
70 context: anytype,
71 comptime body: anytype,
72 diagnostic: *builder.VerificationDiagnostic,
73 ) !Program {
74 var program_builder = try Builder.init(allocator, limits, kernel_name, params);
75 errdefer program_builder.deinit();
76 try body(&program_builder, context);
77 try program_builder.return_();
78 return program_builder.finishWithDiagnostic(diagnostic);
79 }
80
81 pub const Guard = guard_mod.Guard(Builder);
82
83 pub const Builder = struct {
84 capacity: graph.Capacity,
85 body: builder.Builder,
86 schedule: schedule_mod.Schedule,
87 body_active: bool = true,
88 schedule_active: bool = true,
89
90 pub const Result: type = Program;
91 pub const Limits = core.Limits;
92 pub const Capacity = graph.Capacity;
93
94 pub fn init(
95 allocator: std.mem.Allocator,
96 limits: @This().Limits,
97 kernel_name: []const u8,
98 params: []const builder.Param,
99 ) !Builder {
100 const capacity = try graph.Capacity.derive(limits);
101 var body = try builder.Builder.init(
102 allocator,
103 limits.raw(),
104 kernel_name,
105 params,
106 context_requirements,
107 );
108 errdefer body.deinit();
109 var schedule = try schedule_mod.Schedule.init(allocator, limits.schedule);
110 errdefer schedule.deinit(allocator);
111
112 return .{
113 .capacity = capacity,
114 .body = body,
115 .schedule = schedule,
116 };
117 }
118
119 pub fn initBorrowing(
120 allocator: std.mem.Allocator,
121 limits: @This().Limits,
122 ir_ctx: *ir.Context,
123 kernel_name: []const u8,
124 params: []const builder.Param,
125 ) !Builder {
126 const capacity = try graph.Capacity.deriveBorrowing(limits);
127 var body = try builder.Builder.initBorrowing(
128 allocator,
129 limits.raw(),
130 ir_ctx,
131 kernel_name,
132 params,
133 );
134 errdefer body.deinit();
135 var schedule = try schedule_mod.Schedule.init(allocator, limits.schedule);
136 errdefer schedule.deinit(allocator);
137
138 return .{
139 .capacity = capacity,
140 .body = body,
141 .schedule = schedule,
142 };
143 }
144
145 pub fn deinit(self: *Builder) void {
146 const allocator = self.body.allocator;
147 if (self.body_active) self.body.deinit();
148 if (self.schedule_active) self.schedule.deinit(allocator);
149 self.* = undefined;
150 }
151
152 pub fn finish(self: *Builder) !Program {
153 return self.finishKernel(try self.body.finish());
154 }
155
156 pub fn finishWithDiagnostic(self: *Builder, diagnostic: *builder.VerificationDiagnostic) !Program {
157 return self.finishKernel(try self.body.finishWithDiagnostic(diagnostic));
158 }
159
160 fn finishKernel(self: *Builder, finished_kernel: builder.Kernel) !Program {
161 var kernel = finished_kernel;
162 errdefer kernel.deinit();
163
164 var schedule = self.schedule;
165 schedule.activate();
166 self.body_active = false;
167 self.schedule_active = false;
168
169 return Program.init(kernel, schedule, self.capacity);
170 }
171
172 pub const axis = surface_schedule.axis;
173 pub const split = surface_schedule.split;
174 pub const bind = surface_schedule.bind;
175 pub const vectorize = surface_schedule.vectorize;
176 pub const unroll = surface_schedule.unroll;
177
178 pub const index1D = surface_domain.index1D;
179 pub const vectorIndex1D = surface_domain.vectorIndex1D;
180 pub const index2D = surface_domain.index2D;
181 pub const index3D = surface_domain.index3D;
182 pub const guardIndexDo = surface_domain.guardIndexDo;
183 pub const guardVectorIndexDo = surface_domain.guardVectorIndexDo;
184 pub const guardIndex2DDo = surface_domain.guardIndex2DDo;
185 pub const guardIndex3DDo = surface_domain.guardIndex3DDo;
186 pub const forEach1D = surface_domain.forEach1D;
187 pub const forEachVector1D = surface_domain.forEachVector1D;
188 pub const forEach2D = surface_domain.forEach2D;
189 pub const forEach3D = surface_domain.forEach3D;
190
191 pub const guard = surface_guard.guard;
192 pub const guardIndex = surface_guard.guardIndex;
193 pub const guardDo = surface_guard.guardDo;
194
195 pub const argument = surface_primitive.argument;
196 pub const entryBlock = surface_primitive.entryBlock;
197 pub const insertionBlock = surface_primitive.insertionBlock;
198 pub const setInsertionBlock = surface_primitive.setInsertionBlock;
199 pub const enterBlock = surface_primitive.enterBlock;
200 pub const globalId = surface_primitive.globalId;
201 pub const threadId = surface_primitive.threadId;
202 pub const blockId = surface_primitive.blockId;
203 pub const blockDim = surface_primitive.blockDim;
204 pub const gridDim = surface_primitive.gridDim;
205 pub const laneId = surface_primitive.laneId;
206 pub const warpId = surface_primitive.warpId;
207 pub const constantInt = surface_primitive.constantInt;
208 pub const constantIndex = surface_primitive.constantIndex;
209 pub const constantFloat = surface_primitive.constantFloat;
210 pub const constantBool = surface_primitive.constantBool;
211 pub const add = surface_primitive.add;
212 pub const sub = surface_primitive.sub;
213 pub const mul = surface_primitive.mul;
214 pub const umulhi = surface_primitive.umulhi;
215 pub const div = surface_primitive.div;
216 pub const min = surface_primitive.min;
217 pub const max = surface_primitive.max;
218 pub const and_ = surface_primitive.and_;
219 pub const or_ = surface_primitive.or_;
220 pub const xor = surface_primitive.xor;
221 pub const not = surface_primitive.not;
222 pub const popcount = surface_primitive.popcount;
223 pub const shl = surface_primitive.shl;
224 pub const shr = surface_primitive.shr;
225 pub const ushr = surface_primitive.ushr;
226 pub const bitcast = surface_primitive.bitcast;
227 pub const neg = surface_primitive.neg;
228 pub const abs = surface_primitive.abs;
229 pub const splatVector = surface_primitive.splatVector;
230 pub const shuffleVector = surface_primitive.shuffleVector;
231 pub const extractLane = surface_primitive.extractLane;
232 pub const insertLane = surface_primitive.insertLane;
233 pub const packVector = surface_primitive.packVector;
234 pub const sqrt = surface_primitive.sqrt;
235 pub const exp = surface_primitive.exp;
236 pub const log = surface_primitive.log;
237 pub const tanh = surface_primitive.tanh;
238 pub const sin = surface_primitive.sin;
239 pub const cos = surface_primitive.cos;
240 pub const tan = surface_primitive.tan;
241 pub const floor = surface_primitive.floor;
242 pub const round = surface_primitive.round;
243 pub const trunc = surface_primitive.trunc;
244 pub const tf32Round = surface_primitive.tf32Round;
245 pub const pow = surface_primitive.pow;
246 pub const atan2 = surface_primitive.atan2;
247 pub const fma = surface_primitive.fma;
248 pub const compare = surface_primitive.compare;
249 pub const select = surface_primitive.select;
250 pub const cast = surface_primitive.cast;
251 pub const castIndex = surface_primitive.castIndex;
252 pub const load = surface_primitive.load;
253 pub const loadVector = surface_primitive.loadVector;
254 pub const linearIndex = surface_primitive.linearIndex;
255 pub const loadIndex = surface_primitive.loadIndex;
256 pub const loadVectorIndex = surface_primitive.loadVectorIndex;
257 pub const sharedBuffer = surface_primitive.sharedBuffer;
258 pub const dynamicSharedBuffer = surface_primitive.dynamicSharedBuffer;
259 pub const store = surface_primitive.store;
260 pub const storeIndex = surface_primitive.storeIndex;
261 pub const atomicRmw = surface_primitive.atomicRmw;
262 pub const atomicCas = surface_primitive.atomicCas;
263 pub const barrier = surface_primitive.barrier;
264 pub const activeMask = surface_primitive.activeMask;
265 pub const syncWarp = surface_primitive.syncWarp;
266 pub const allSync = surface_primitive.allSync;
267 pub const anySync = surface_primitive.anySync;
268 pub const ballotSync = surface_primitive.ballotSync;
269 pub const mmaSync = surface_primitive.mmaSync;
270 pub const fence = surface_primitive.fence;
271 pub const asyncCopyShared = surface_primitive.asyncCopyShared;
272 pub const asyncCopyCommit = surface_primitive.asyncCopyCommit;
273 pub const asyncCopyWait = surface_primitive.asyncCopyWait;
274 pub const if_ = surface_primitive.if_;
275 pub const for_ = surface_primitive.for_;
276 pub const forScope = surface_primitive.forScope;
277 pub const while_ = surface_primitive.while_;
278 pub const whileScope = surface_primitive.whileScope;
279 pub const yield_ = surface_primitive.yield_;
280 pub const return_ = surface_primitive.return_;
281
282 pub const typedArgument = builder_typed.typedArgument;
283 pub const typedValue = builder_typed.typedValue;
284 pub const castValue = builder_typed.castValue;
285 pub const bufferArgument = builder_typed.bufferArgument;
286 pub const bufferView = builder_typed.bufferView;
287 pub const globalIdValue = builder_typed.globalIdValue;
288 pub const threadIdValue = builder_typed.threadIdValue;
289 pub const blockIdValue = builder_typed.blockIdValue;
290 pub const blockDimValue = builder_typed.blockDimValue;
291 pub const gridDimValue = builder_typed.gridDimValue;
292 pub const laneIdValue = builder_typed.laneIdValue;
293 pub const warpIdValue = builder_typed.warpIdValue;
294 pub const constantValue = builder_typed.constantValue;
295 pub const linearIndexValue = builder_typed.linearIndexValue;
296 pub const sharedBufferView = builder_typed.sharedBufferView;
297 pub const vec2 = builder_typed.vec2;
298 pub const typedVec2 = builder_typed.typedVec2;
299 pub const vec3 = builder_typed.vec3;
300 pub const typedVec3 = builder_typed.typedVec3;
301 pub const splat2 = builder_typed.splat2;
302 pub const typedSplat2 = builder_typed.typedSplat2;
303 pub const splat3 = builder_typed.splat3;
304 pub const typedSplat3 = builder_typed.typedSplat3;
305 pub const activeMaskValue = builder_typed.activeMaskValue;
306 pub const allSyncValue = builder_typed.allSyncValue;
307 pub const anySyncValue = builder_typed.anySyncValue;
308 pub const ballotSyncValue = builder_typed.ballotSyncValue;
309 pub const warpReduce = builder_typed.warpReduce;
310 pub const warpScan = builder_typed.warpScan;
311 pub const shuffleSync = builder_typed.shuffleSync;
312
313 pub const forDo = surface_flow.forDo;
314 pub const forValueDo = surface_flow.forValueDo;
315 pub const fold = surface_flow.fold;
316 pub const whileLoop = surface_flow.whileLoop;
317 pub const forRangeDo = surface_flow.forRangeDo;
318 pub const forRangeValueDo = surface_flow.forRangeValueDo;
319 pub const foldRange = surface_flow.foldRange;
320 };
321 };
322 }
323
324 test "Program Builder finish transfers preacquired body and Schedule storage" {
325 const generated = Surface(graph.Program, .{});
326 var failing = std.testing.FailingAllocator.init(std.testing.allocator, .{});
327 var program_builder = try generated.Builder.init(
328 failing.allocator(),
329 generated.Limits.standard,
330 "program_transfer",
331 &.{},
332 );
333 errdefer program_builder.deinit();
334 const body_storage = program_builder.body.storage.storage;
335 const schedule_storage = program_builder.schedule.storage;
336 try program_builder.return_();
337
338 failing.fail_index = failing.alloc_index;
339 failing.resize_fail_index = failing.resize_index;
340 var program = try program_builder.finish();
341 program_builder.deinit();
342 defer program.deinit();
343
344 try std.testing.expectEqual(body_storage, program.storage.kernel.storage.storage);
345 try std.testing.expectEqual(schedule_storage, program.storage.schedule.storage);
346 try std.testing.expect(!failing.has_induced_failure);
347 }