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 }