lib/accy/src/kernel/model/core/builder.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const choir_abi = @import("choir_abi");
3 const testing = std.testing;
4
5 const alloc_phase = @import("alloc_phase");
6 const choir = @import("choir");
7 const gpu = choir.dialects.gpu;
8
9 const ir = choir.ir;
10 const dialects = choir.dialects;
11 const ArithDialect = dialects.ArithDialect;
12 const BuiltinDialect = dialects.BuiltinDialect;
13 const FuncDialect = dialects.FuncDialect;
14 const GpuDialect = gpu.GpuDialect;
15 const MemrefDialect = dialects.MemrefDialect;
16 const ScfDialect = dialects.ScfDialect;
17
18 pub const Dimension = gpu.Dimension;
19 pub const AddressSpace = dialects.AddressSpace;
20 pub const AtomicRmwKind = dialects.AtomicRmwKind;
21 pub const MemoryOrder = gpu.MemoryOrder;
22 pub const Scope = gpu.Scope;
23 pub const MmaShape = gpu.MmaShape;
24 pub const ShuffleMode = gpu.ShuffleMode;
25 pub const WarpOpKind = gpu.WarpOpKind;
26 pub const BackendError = choir.backends.interface.BackendError;
27 pub const FingerprintError = choir.OperationFingerprintError;
28
29 pub const WarpScanMode = enum {
30 inclusive,
31 exclusive,
32 };
33
34 const DType = choir_abi.DType;
35
36 pub const Compare = dialects.arith.CmpPredicate;
37 pub const VerificationDiagnostic = choir.backends.contract.VerificationDiagnostic;
38
39 /// A caller passes this to a kernel builder to say which dialects, named groups of compiler
40 /// operations, are needed by the kernel's own compiler context, the object that owns the dialects,
41 /// types and operations of one compiler module. The builder reads the value once while setting up a
42 /// compiler context it owns and before that context is frozen. `registrations` lists dialects to
43 /// register and then load, and `preload` lists names of dialects already registered with the
44 /// compiler, which are loaded by name. The context copies each name it keeps, so the caller's
45 /// slices only need to live through setup.
46 pub const ContextRequirements = struct {
47 registrations: []const choir.extensions.DialectRegistration = &.{},
48 preload: []const []const u8 = &.{},
49
50 pub fn prepare(self: ContextRequirements, ctx: *ir.Context) !void {
51 const extension = choir.extensions.PackageExtension{
52 .name = "accy-kernel-context",
53 .dialects = self.registrations,
54 };
55 try extension.registerContext(ctx);
56 for (self.registrations) |registration| {
57 _ = try ctx.getOrLoadDialect(registration.name);
58 }
59 for (self.preload) |name| _ = try ctx.getOrLoadDialect(name);
60 }
61 };
62
63 pub const Type = struct {
64 value: ir.Type,
65
66 fn raw(self: Type) ir.Type {
67 return self.value;
68 }
69
70 pub fn eql(self: Type, other: Type) bool {
71 return self.value.eql(other.value);
72 }
73 };
74
75 pub const Value = struct {
76 ptr: *ir.Value,
77
78 fn raw(self: Value) *ir.Value {
79 return self.ptr;
80 }
81
82 pub fn valueType(self: Value) Type {
83 return .{ .value = self.ptr.type };
84 }
85
86 pub fn eql(self: Value, other: Value) bool {
87 return self.ptr == other.ptr;
88 }
89 };
90
91 pub const Block = struct {
92 ptr: *ir.Block,
93
94 fn raw(self: Block) *ir.Block {
95 return self.ptr;
96 }
97
98 pub fn argument(self: Block, index: usize) ?Value {
99 return wrapOptionalValue(self.ptr.getArgument(index));
100 }
101 };
102
103 pub const If = struct {
104 op: ScfDialect.IfOp,
105
106 pub fn thenBlock(self: If) Block {
107 return wrapBlock(self.op.getThenBlock());
108 }
109
110 pub fn elseBlock(self: If) ?Block {
111 const block = self.op.getElseBlock() orelse return null;
112 return wrapBlock(block);
113 }
114
115 pub fn result(self: *const If, index: usize) ?Value {
116 return wrapOptionalValue(self.op.getResult(index));
117 }
118
119 pub fn resultCount(self: If) usize {
120 return self.op.getNumResults();
121 }
122
123 fn raw(self: If) ScfDialect.IfOp {
124 return self.op;
125 }
126 };
127
128 pub const For = struct {
129 op: ScfDialect.ForOp,
130
131 pub fn bodyBlock(self: For) Block {
132 return wrapBlock(self.op.getBodyBlock());
133 }
134
135 pub fn inductionVar(self: For) Value {
136 return wrapValue(self.op.getInductionVar());
137 }
138
139 pub fn iterArg(self: For, index: usize) ?Value {
140 const args = self.op.getIterArgs();
141 if (index >= args.len) return null;
142 return wrapValue(args[index]);
143 }
144
145 pub fn result(self: *const For, index: usize) ?Value {
146 return wrapOptionalValue(self.op.getResult(index));
147 }
148
149 fn raw(self: For) ScfDialect.ForOp {
150 return self.op;
151 }
152 };
153
154 pub const ForScope = struct {
155 builder: *Builder,
156 loop: For,
157 previous: Block,
158 active: bool = true,
159
160 pub fn inductionVar(self: ForScope) Value {
161 return self.loop.inductionVar();
162 }
163
164 pub fn iterArg(self: ForScope, index: usize) ?Value {
165 return self.loop.iterArg(index);
166 }
167
168 pub fn result(self: *const ForScope, index: usize) ?Value {
169 return self.loop.result(index);
170 }
171
172 pub fn leave(self: *ForScope, values: []const Value) !void {
173 if (!self.active) return;
174 errdefer self.builder.setInsertionBlock(self.previous);
175 try self.builder.yield_(values);
176 self.builder.setInsertionBlock(self.previous);
177 self.active = false;
178 }
179
180 pub fn abort(self: *ForScope) void {
181 if (!self.active) return;
182 self.builder.setInsertionBlock(self.previous);
183 self.active = false;
184 }
185 };
186
187 pub const While = struct {
188 op: ScfDialect.WhileOp,
189
190 pub fn beforeBlock(self: While) Block {
191 return wrapBlock(self.op.getBeforeBlock());
192 }
193
194 pub fn afterBlock(self: While) Block {
195 return wrapBlock(self.op.getAfterBlock());
196 }
197
198 pub fn result(self: *const While, index: usize) ?Value {
199 return wrapOptionalValue(self.op.op.getResult(index));
200 }
201
202 fn raw(self: While) ScfDialect.WhileOp {
203 return self.op;
204 }
205 };
206
207 pub const WhileScope = struct {
208 builder: *Builder,
209 loop: While,
210 previous: Block,
211 phase: Phase = .before,
212
213 const Phase = enum { before, after, done };
214
215 pub fn beforeArg(self: WhileScope, index: usize) ?Value {
216 return self.loop.beforeBlock().argument(index);
217 }
218
219 pub fn afterArg(self: WhileScope, index: usize) ?Value {
220 return self.loop.afterBlock().argument(index);
221 }
222
223 pub fn result(self: *const WhileScope, index: usize) ?Value {
224 return self.loop.result(index);
225 }
226
227 pub fn condition(self: *WhileScope, cond: Value, args: []const Value) !void {
228 if (self.phase != .before) return error.WhileScopeMisused;
229 try self.builder.condition_(cond, args);
230 self.builder.setInsertionBlock(self.loop.afterBlock());
231 self.phase = .after;
232 }
233
234 pub fn leave(self: *WhileScope, values: []const Value) !void {
235 if (self.phase != .after) return error.WhileScopeMisused;
236 errdefer self.builder.setInsertionBlock(self.previous);
237 try self.builder.yield_(values);
238 self.builder.setInsertionBlock(self.previous);
239 self.phase = .done;
240 }
241
242 pub fn abort(self: *WhileScope) void {
243 if (self.phase == .done) return;
244 self.builder.setInsertionBlock(self.previous);
245 self.phase = .done;
246 }
247 };
248
249 const State = struct {
250 ctx: *ir.Context,
251 module: *ir.Operation,
252 func: FuncDialect.FuncOp,
253 params: []Param,
254 decoded: ?choir.bytecode.DecodedModule = null,
255 };
256
257 pub const Buffer = struct {
258 dtype: DType,
259 size: ?u64 = null,
260 space: AddressSpace = .device,
261
262 pub fn eql(self: Buffer, other: Buffer) bool {
263 return self.dtype == other.dtype and self.size == other.size and self.space == other.space;
264 }
265 };
266
267 pub const Param = union(enum) {
268 scalar: DType,
269 buffer: Buffer,
270
271 pub fn getType(self: Param, ctx: *ir.Context) !ir.Type {
272 return switch (self) {
273 .scalar => |dtype| scalarType(ctx, dtype),
274 .buffer => |buf| bufferType(ctx, buf),
275 };
276 }
277
278 pub fn eql(self: Param, other: Param) bool {
279 return switch (self) {
280 .scalar => |dtype| switch (other) {
281 .scalar => |other_dtype| dtype == other_dtype,
282 .buffer => false,
283 },
284 .buffer => |spec| switch (other) {
285 .scalar => false,
286 .buffer => |other_spec| spec.eql(other_spec),
287 },
288 };
289 }
290 };
291
292 pub fn paramsEql(lhs: []const Param, rhs: []const Param) bool {
293 if (lhs.len != rhs.len) return false;
294 for (lhs, rhs) |left, right| {
295 if (!left.eql(right)) return false;
296 }
297 return true;
298 }
299
300 pub fn scalar(dtype: DType) Param {
301 return .{ .scalar = dtype };
302 }
303
304 pub fn buffer(dtype: DType, size: u64) Param {
305 return .{ .buffer = .{ .dtype = dtype, .size = size } };
306 }
307
308 pub fn dynamicBuffer(dtype: DType) Param {
309 return .{ .buffer = .{ .dtype = dtype } };
310 }
311
312 pub const Storage = struct {
313 phase: alloc_phase.capacity.Phase,
314 capacity: Capacity,
315 storage: [*]u8,
316 context: *ir.Context,
317 state: *State,
318 params: []Param,
319 values: []*ir.Value,
320 types: []ir.Type,
321 owns_context: bool,
322 parameters_high_water: usize,
323 kernel_name_bytes_high_water: usize,
324 temporary_values_high_water: usize,
325 temporary_types_high_water: usize,
326
327 pub const Limits = @import("limits.zig").RawLimits;
328
329 pub const Capacity = struct {
330 context: ir.Context.Capacity,
331 parameters: usize,
332 kernel_name_bytes: usize,
333 temporary_values: usize,
334 temporary_types: usize,
335 context_offset: usize,
336 state_offset: usize,
337 parameters_offset: usize,
338 values_offset: usize,
339 types_offset: usize,
340 storage_bytes: usize,
341 storage_alignment: std.mem.Alignment,
342 borrowed_acquired_bytes: usize,
343 owned_acquired_bytes: usize,
344
345 pub fn derive(limits: Limits) error{CapacityOverflow}!Capacity {
346 const context = try ir.Context.Capacity.derive(limits.context);
347 var cursor: usize = 0;
348 var alignment: usize = 1;
349 const context_offset = try placeSlice(ir.Context, 1, &cursor, &alignment);
350 const state_offset = try placeSlice(State, 1, &cursor, &alignment);
351 const parameters_offset = try placeSlice(Param, limits.parameters, &cursor, &alignment);
352 const values_offset = try placeSlice(*ir.Value, limits.temporary_values, &cursor, &alignment);
353 const types_offset = try placeSlice(ir.Type, limits.temporary_types, &cursor, &alignment);
354 if (cursor == 0) cursor = 1;
355 const owned_acquired_bytes = std.math.add(
356 usize,
357 cursor,
358 context.storage_bytes,
359 ) catch return error.CapacityOverflow;
360 return .{
361 .context = context,
362 .parameters = limits.parameters,
363 .kernel_name_bytes = limits.kernel_name_bytes,
364 .temporary_values = limits.temporary_values,
365 .temporary_types = limits.temporary_types,
366 .context_offset = context_offset,
367 .state_offset = state_offset,
368 .parameters_offset = parameters_offset,
369 .values_offset = values_offset,
370 .types_offset = types_offset,
371 .storage_bytes = cursor,
372 .storage_alignment = .fromByteUnits(alignment),
373 .borrowed_acquired_bytes = cursor,
374 .owned_acquired_bytes = owned_acquired_bytes,
375 };
376 }
377 };
378
379 pub const Exhaustion = error{
380 ContextCapacityExceeded,
381 ParameterCapacityExceeded,
382 KernelNameCapacityExceeded,
383 TemporaryValueCapacityExceeded,
384 TemporaryTypeCapacityExceeded,
385 };
386
387 pub const Usage = struct {
388 parameters: usize,
389 kernel_name_bytes: usize,
390 temporary_values: usize,
391 temporary_types: usize,
392 };
393
394 pub const claim: alloc_phase.capacity.Declaration = .{
395 .source = .{
396 .id = "accy.kernel_raw_builder_storage",
397 .kind = .phase_static,
398 .limit_source = .caller,
399 .storage = .{
400 .covered = &.{
401 .{
402 .id = "stable_builder_state_and_optional_owned_choir_context_handle",
403 .lifetime = .steady,
404 .detail = "stable Builder state and optional owned Choir Context handle",
405 },
406 .{
407 .id = "retained_kernel_parameter_schema",
408 .lifetime = .transferred,
409 .detail = "retained kernel parameter schema",
410 },
411 .{
412 .id = "bounded_value_and_type_conversion_scratch",
413 .lifetime = .initialization,
414 .detail = "bounded value and type conversion scratch",
415 },
416 .{
417 .id = "caller_bounded_kernel_name_admitted_into_choir_context_storage",
418 .lifetime = .transferred,
419 .detail = "caller-bounded kernel name admitted into Choir Context storage",
420 },
421 .{
422 .id = "owned_choir_context_storage_and_all_retained_ir",
423 .lifetime = .transferred,
424 .detail = "owned Choir Context storage and all retained IR",
425 },
426 },
427 .excluded = &.{
428 "borrowed Choir Context storage supplied by initBorrowing",
429 "caller allocator implementation state and caller-owned diagnostic output",
430 "program Schedule, Snapshot, replay, backend, and oracle owners",
431 },
432 },
433 .capacity = .{
434 .inputs = &.{
435 alloc_phase.capacity.bindInput(Limits, "context_attributes_payload_bytes", "context.attributes.payload_bytes"),
436 alloc_phase.capacity.bindInput(Limits, "context_attributes_table_bytes", "context.attributes.table_bytes"),
437 alloc_phase.capacity.bindInput(Limits, "context_configuration_interface_bytes", "context.configuration.interface_bytes"),
438 alloc_phase.capacity.bindInput(Limits, "context_configuration_name_bytes", "context.configuration.name_bytes"),
439 alloc_phase.capacity.bindInput(Limits, "context_configuration_table_bytes", "context.configuration.table_bytes"),
440 alloc_phase.capacity.bindInput(Limits, "context_configuration_transaction_bytes", "context.configuration.transaction_bytes"),
441 alloc_phase.capacity.bindInput(Limits, "context_diagnostics_handler_bytes", "context.diagnostics.handler_bytes"),
442 alloc_phase.capacity.bindInput(Limits, "context_diagnostics_payload_bytes", "context.diagnostics.payload_bytes"),
443 alloc_phase.capacity.bindInput(Limits, "context_operations_nested_bytes", "context.operations.nested_bytes"),
444 alloc_phase.capacity.bindInput(Limits, "context_operations_storage_bytes", "context.operations.storage_bytes"),
445 alloc_phase.capacity.bindInput(Limits, "context_transient_bytes", "context.transient_bytes"),
446 alloc_phase.capacity.bindInput(Limits, "context_types_key_bytes", "context.types.key_bytes"),
447 alloc_phase.capacity.bindInput(Limits, "context_types_payload_bytes", "context.types.payload_bytes"),
448 alloc_phase.capacity.bindInput(Limits, "context_types_table_bytes", "context.types.table_bytes"),
449 alloc_phase.capacity.bindInput(Limits, "kernel_name_bytes", "kernel_name_bytes"),
450 alloc_phase.capacity.bindInput(Limits, "parameters", "parameters"),
451 alloc_phase.capacity.bindInput(Limits, "temporary_types", "temporary_types"),
452 alloc_phase.capacity.bindInput(Limits, "temporary_values", "temporary_values"),
453 },
454 .type_selectors = &.{},
455 .nodes = &.{
456 .{ .input = 0 },
457 .{ .input = 1 },
458 .{ .input = 2 },
459 .{ .input = 3 },
460 .{ .input = 4 },
461 .{ .input = 5 },
462 .{ .input = 6 },
463 .{ .input = 7 },
464 .{ .input = 8 },
465 .{ .input = 9 },
466 .{ .input = 10 },
467 .{ .input = 11 },
468 .{ .input = 12 },
469 .{ .input = 13 },
470 .{ .input = 14 },
471 .{ .input = 15 },
472 .{ .input = 16 },
473 .{ .input = 17 },
474 .{ .add = .{ .left = 0, .right = 1 } },
475 .{ .add = .{ .left = 18, .right = 2 } },
476 .{ .add = .{ .left = 19, .right = 3 } },
477 .{ .add = .{ .left = 20, .right = 4 } },
478 .{ .add = .{ .left = 21, .right = 5 } },
479 .{ .add = .{ .left = 22, .right = 6 } },
480 .{ .add = .{ .left = 23, .right = 7 } },
481 .{ .add = .{ .left = 24, .right = 8 } },
482 .{ .add = .{ .left = 25, .right = 9 } },
483 .{ .add = .{ .left = 26, .right = 10 } },
484 .{ .add = .{ .left = 27, .right = 11 } },
485 .{ .add = .{ .left = 28, .right = 12 } },
486 .{ .add = .{ .left = 29, .right = 13 } },
487 .{ .add = .{ .left = 30, .right = 14 } },
488 .{ .add = .{ .left = 31, .right = 15 } },
489 .{ .add = .{ .left = 32, .right = 16 } },
490 .{ .add = .{ .left = 33, .right = 17 } },
491 },
492 .assertions = &.{.{
493 .scope = .closure_total,
494 .measure = .retained,
495 .relation = .upper_bound,
496 .expression = 34,
497 }},
498 },
499 .overload = .{
500 .kind = .reject_before_mutation,
501 .detail = "name, parameter, scratch, Context budget, arithmetic, and fixed Context exhaustion reject before publishing partial Builder state",
502 },
503 .risks = .{
504 .transitive = .{
505 .status = .witnessed,
506 .detail = "Builder methods use only preacquired local scratch and the independently bounded Choir Context",
507 },
508 .foreign = .{
509 .status = .excluded,
510 .detail = "kernel construction crosses no foreign or operating-system boundary",
511 },
512 },
513 .dependencies = &.{"choir.context"},
514 .obligations = &.{
515 .{ .key = "accy_kernel_raw_builder_capacity", .role = .capacity_model },
516 .{ .key = "accy_kernel_raw_builder_boundary", .role = .overload },
517 .{ .key = "accy_kernel_raw_builder_context_boundary", .role = .overload },
518 .{ .key = "accy_kernel_raw_builder_sealed_construction", .role = .transitive_risk },
519 .{ .key = "accy_kernel_raw_builder_transfer", .role = .foreign_risk },
520 },
521 },
522 .bindings = .{
523 .owner = @This(),
524 .seal = .{
525 .family = alloc_phase.capacity.selector(@This().activate),
526 .premise = .{
527 .class = .checked_semantic_fact,
528 .authority = .checker,
529 },
530 },
531 .teardown = .{
532 .family = alloc_phase.capacity.selector(@This().deinit),
533 .premise = .{
534 .class = .checked_semantic_fact,
535 .authority = .checker,
536 },
537 },
538 },
539 };
540
541 pub fn init(allocator: std.mem.Allocator, limits: Limits) !Storage {
542 var storage = try acquire(allocator, limits);
543 errdefer releaseStorage(&storage, allocator);
544 storage.context.* = try ir.Context.init(allocator, limits.context);
545 storage.owns_context = true;
546 return storage;
547 }
548
549 pub fn initBorrowing(
550 allocator: std.mem.Allocator,
551 limits: Limits,
552 context: *ir.Context,
553 ) !Storage {
554 var storage = try acquire(allocator, limits);
555 storage.context = context;
556 return storage;
557 }
558
559 fn acquire(allocator: std.mem.Allocator, limits: Limits) !Storage {
560 const capacity = try Capacity.derive(limits);
561 const storage = allocator.rawAlloc(
562 capacity.storage_bytes,
563 capacity.storage_alignment,
564 @returnAddress(),
565 ) orelse return error.OutOfMemory;
566 return .{
567 .phase = .initialization,
568 .capacity = capacity,
569 .storage = storage,
570 .context = &typedSlice(ir.Context, storage, capacity.context_offset, 1)[0],
571 .state = &typedSlice(State, storage, capacity.state_offset, 1)[0],
572 .params = typedSlice(Param, storage, capacity.parameters_offset, capacity.parameters),
573 .values = typedSlice(*ir.Value, storage, capacity.values_offset, capacity.temporary_values),
574 .types = typedSlice(ir.Type, storage, capacity.types_offset, capacity.temporary_types),
575 .owns_context = false,
576 .parameters_high_water = 0,
577 .kernel_name_bytes_high_water = 0,
578 .temporary_values_high_water = 0,
579 .temporary_types_high_water = 0,
580 };
581 }
582
583 pub fn activate(self: *Storage) void {
584 std.debug.assert(self.phase == .initialization);
585 if (self.owns_context) self.context.activate();
586 self.phase = .steady;
587 }
588
589 pub fn require(self: *Storage, usage: Usage) Exhaustion!void {
590 if (usage.parameters > self.capacity.parameters) return error.ParameterCapacityExceeded;
591 if (usage.kernel_name_bytes > self.capacity.kernel_name_bytes) return error.KernelNameCapacityExceeded;
592 if (usage.temporary_values > self.capacity.temporary_values) return error.TemporaryValueCapacityExceeded;
593 if (usage.temporary_types > self.capacity.temporary_types) return error.TemporaryTypeCapacityExceeded;
594 self.parameters_high_water = @max(self.parameters_high_water, usage.parameters);
595 self.kernel_name_bytes_high_water = @max(
596 self.kernel_name_bytes_high_water,
597 usage.kernel_name_bytes,
598 );
599 self.temporary_values_high_water = @max(
600 self.temporary_values_high_water,
601 usage.temporary_values,
602 );
603 self.temporary_types_high_water = @max(
604 self.temporary_types_high_water,
605 usage.temporary_types,
606 );
607 }
608
609 pub fn capacityUsage(self: *const Storage) Usage {
610 return .{
611 .parameters = self.parameters_high_water,
612 .kernel_name_bytes = self.kernel_name_bytes_high_water,
613 .temporary_values = self.temporary_values_high_water,
614 .temporary_types = self.temporary_types_high_water,
615 };
616 }
617
618 pub fn deinit(self: *Storage, allocator: std.mem.Allocator) void {
619 std.debug.assert(self.phase != .teardown);
620 if (self.owns_context) self.context.deinit(allocator);
621 releaseStorage(self, allocator);
622 }
623 };
624
625 comptime {
626 alloc_phase.capacity.requireAllocatorRejectingOwnerShape(Storage);
627 }
628
629 fn releaseStorage(storage_owner: *Storage, allocator: std.mem.Allocator) void {
630 const storage = storage_owner.storage;
631 const capacity = storage_owner.capacity;
632 storage_owner.phase = .teardown;
633 storage_owner.* = undefined;
634 allocator.rawFree(
635 storage[0..capacity.storage_bytes],
636 capacity.storage_alignment,
637 @returnAddress(),
638 );
639 }
640
641 pub const Kernel = struct {
642 allocator: std.mem.Allocator,
643 capacity: Builder.Capacity,
644 storage: Storage,
645 state: *State,
646
647 /// A caller uses this to wrap a kernel module decoded from bytes as a ready kernel, to avoid
648 /// building it again. The kernel takes ownership of the decoded module, its tables, and the
649 /// storage passed in on success, and the caller keeps all three on failure because
650 /// `storage_value` is passed by value. The caller decodes the module through the operation
651 /// allocator of the storage's compiler context before the call, and the call asserts this
652 /// requirement. `function` must be a kernel with a body and a name inside the decoded module,
653 /// and its arguments must match `parameters` in count and type, and otherwise the call returns
654 /// `error.InvalidKernel`. The call also returns a storage capacity error when the name or
655 /// parameters exceed the storage, or a verification error when the module fails to verify. The
656 /// parameters are copied into the kernel's storage, so the caller's slice only needs to live
657 /// through the call.
658 pub fn fromDecoded(
659 allocator: std.mem.Allocator,
660 storage_value: Storage,
661 decoded: choir.bytecode.DecodedModule,
662 function: FuncDialect.FuncOp,
663 parameters: []const Param,
664 ) !Kernel {
665 const module_op = decoded.module;
666 const retained = ir.context.operationAllocator(storage_value.context);
667 std.debug.assert(decoded.tables.allocator.ptr == retained.ptr);
668 std.debug.assert(decoded.tables.allocator.vtable == retained.vtable);
669 std.debug.assert(storage_value.phase == .initialization);
670 std.debug.assert(module_op.context == storage_value.context);
671 if (!module_op.isProperAncestor(function.op) or !function.isKernel() or
672 !function.hasBody() or function.getNumArguments() != parameters.len)
673 {
674 return error.InvalidKernel;
675 }
676 const name = function.getName() orelse return error.InvalidKernel;
677 var storage = storage_value;
678 try storage.require(.{
679 .parameters = parameters.len,
680 .kernel_name_bytes = name.len,
681 .temporary_values = 0,
682 .temporary_types = 0,
683 });
684 for (parameters, function.getArguments()) |parameter, argument| {
685 if (!argument.type.eql(try parameter.getType(storage.context))) {
686 return error.InvalidKernel;
687 }
688 }
689 try ir.verifyOperation(module_op, ir.verify.default_options);
690 const params_copy = storage.params[0..parameters.len];
691 @memcpy(params_copy, parameters);
692 storage.state.* = .{
693 .ctx = storage.context,
694 .module = module_op,
695 .func = function,
696 .params = params_copy,
697 .decoded = decoded,
698 };
699 storage.activate();
700 return .{
701 .allocator = allocator,
702 .capacity = storage.capacity,
703 .storage = storage,
704 .state = storage.state,
705 };
706 }
707
708 pub fn deinit(self: *Kernel) void {
709 destroyState(self.state, self.allocator, &self.storage);
710 self.* = undefined;
711 }
712
713 pub fn verify(self: *Kernel) !void {
714 try ir.verifyOperation(self.state.module, ir.verify.default_options);
715 }
716
717 pub fn verifyWithDiagnostic(self: *Kernel, diagnostic: *VerificationDiagnostic) !void {
718 try choir.backends.contract.verifyModuleWithDiagnostic("accy/kernel/verify", self.state.module, diagnostic);
719 }
720
721 pub fn fingerprint(self: *const Kernel, allocator: std.mem.Allocator) FingerprintError!u64 {
722 return choir.operationFingerprint(allocator, self.state.module);
723 }
724
725 pub fn func(self: *const Kernel) FuncDialect.FuncOp {
726 return self.state.func;
727 }
728
729 pub fn module(self: *const Kernel) *ir.Operation {
730 return self.state.module;
731 }
732
733 pub fn params(self: *const Kernel) []const Param {
734 return self.state.params;
735 }
736
737 pub fn capacityUsage(self: *const Kernel) Builder.Usage {
738 return self.storage.capacityUsage();
739 }
740
741 pub fn acquiredBytes(self: *const Kernel) usize {
742 return acquiredBytesForStorage(&self.storage);
743 }
744
745 pub fn contextCapacityUsage(self: *const Kernel) ir.Context.Usage {
746 return self.state.ctx.capacityUsage();
747 }
748 };
749
750 pub const Builder = struct {
751 allocator: std.mem.Allocator,
752 capacity: Capacity,
753 storage: Storage,
754 state: *State,
755 insertion_block: *ir.Block,
756 loc: ir.Location,
757 finished: bool = false,
758 terminated: bool = false,
759
760 pub const Limits = @import("limits.zig").RawLimits;
761 pub const Capacity = Storage.Capacity;
762 pub const Usage = Storage.Usage;
763
764 pub fn init(
765 backing: std.mem.Allocator,
766 limits: Limits,
767 kernel_name: []const u8,
768 params: []const Param,
769 requirements: ContextRequirements,
770 ) !Builder {
771 try validateInputs(limits, kernel_name, params);
772 var storage = try Storage.init(backing, limits);
773 errdefer storage.deinit(backing);
774 try dialects.registerChoirDialect(storage.context);
775 try requirements.prepare(storage.context);
776 return initIn(backing, storage, storage.context, kernel_name, params);
777 }
778
779 pub fn initBorrowing(
780 backing: std.mem.Allocator,
781 limits: Limits,
782 ctx: *ir.Context,
783 kernel_name: []const u8,
784 params: []const Param,
785 ) !Builder {
786 try validateInputs(limits, kernel_name, params);
787 try requireContextBudget(ctx, limits.context);
788 var storage = try Storage.initBorrowing(backing, limits, ctx);
789 errdefer storage.deinit(backing);
790 return initIn(backing, storage, ctx, kernel_name, params);
791 }
792
793 fn initIn(
794 backing: std.mem.Allocator,
795 storage_value: Storage,
796 ctx: *ir.Context,
797 kernel_name: []const u8,
798 params: []const Param,
799 ) !Builder {
800 var storage = storage_value;
801 try storage.require(.{
802 .parameters = params.len,
803 .kernel_name_bytes = kernel_name.len,
804 .temporary_values = 0,
805 .temporary_types = params.len,
806 });
807 const state = storage.state;
808 const owned_params = storage.params[0..params.len];
809 @memcpy(owned_params, params);
810
811 const loc = ir.Location.getUnknown();
812 const module = try BuiltinDialect.ModuleOp.create(ctx, loc);
813 errdefer module.op.erase();
814
815 const input_types = storage.types[0..params.len];
816 for (params, 0..) |param, i| {
817 input_types[i] = try param.getType(ctx);
818 }
819
820 var func = try FuncDialect.FuncOp.createKernel(ctx, loc, kernel_name, input_types);
821 errdefer if (func.op.getBlock() == null) func.op.erase();
822 try module.getBodyBlock().addOperation(func.op);
823 state.* = .{
824 .ctx = ctx,
825 .module = module.op,
826 .func = func,
827 .params = owned_params,
828 };
829
830 return .{
831 .allocator = backing,
832 .capacity = storage.capacity,
833 .storage = storage,
834 .state = state,
835 .insertion_block = func.getEntryBlock(),
836 .loc = loc,
837 };
838 }
839
840 pub fn deinit(self: *Builder) void {
841 if (!self.finished) {
842 destroyState(self.state, self.allocator, &self.storage);
843 }
844 self.* = undefined;
845 }
846
847 pub fn capacityUsage(self: *const Builder) Usage {
848 return self.storage.capacityUsage();
849 }
850
851 pub fn acquiredBytes(self: *const Builder) usize {
852 return acquiredBytesForStorage(&self.storage);
853 }
854
855 pub fn contextCapacityUsage(self: *const Builder) ir.Context.Usage {
856 return self.state.ctx.capacityUsage();
857 }
858
859 pub fn argument(self: *Builder, index: usize) Value {
860 return wrapValue(self.state.func.getArgument(index));
861 }
862
863 pub fn entryBlock(self: *Builder) Block {
864 return wrapBlock(self.state.func.getEntryBlock());
865 }
866
867 pub fn insertionBlock(self: *Builder) Block {
868 return wrapBlock(self.insertion_block);
869 }
870
871 pub fn setInsertionBlock(self: *Builder, block: Block) void {
872 self.insertion_block = block.ptr;
873 }
874
875 pub fn enterBlock(self: *Builder, block: Block) !InsertionScope {
876 try self.ensureOpen();
877 const previous = self.insertionBlock();
878 self.setInsertionBlock(block);
879 return .{
880 .builder = self,
881 .previous = previous,
882 };
883 }
884
885 pub fn globalId(self: *Builder, dim: Dimension) !Value {
886 try self.ensureOpen();
887 var op = try GpuDialect.GlobalIdxOp.create(self.state.ctx, self.loc, dim);
888 try self.addOperation(op.op);
889 return wrapValue(op.getResult());
890 }
891
892 pub fn threadId(self: *Builder, dim: Dimension) !Value {
893 try self.ensureOpen();
894 var op = try GpuDialect.ThreadIdxOp.create(self.state.ctx, self.loc, dim);
895 try self.addOperation(op.op);
896 return wrapValue(op.getResult());
897 }
898
899 pub fn blockId(self: *Builder, dim: Dimension) !Value {
900 try self.ensureOpen();
901 var op = try GpuDialect.BlockIdxOp.create(self.state.ctx, self.loc, dim);
902 try self.addOperation(op.op);
903 return wrapValue(op.getResult());
904 }
905
906 pub fn blockDim(self: *Builder, dim: Dimension) !Value {
907 try self.ensureOpen();
908 var op = try GpuDialect.BlockDimOp.create(self.state.ctx, self.loc, dim);
909 try self.addOperation(op.op);
910 return wrapValue(op.getResult());
911 }
912
913 pub fn gridDim(self: *Builder, dim: Dimension) !Value {
914 try self.ensureOpen();
915 var op = try GpuDialect.GridDimOp.create(self.state.ctx, self.loc, dim);
916 try self.addOperation(op.op);
917 return wrapValue(op.getResult());
918 }
919
920 pub fn laneId(self: *Builder) !Value {
921 try self.ensureOpen();
922 var op = try GpuDialect.LaneIdOp.create(self.state.ctx, self.loc);
923 try self.addOperation(op.op);
924 return wrapValue(op.getResult());
925 }
926
927 pub fn warpId(self: *Builder) !Value {
928 try self.ensureOpen();
929 var op = try GpuDialect.WarpIdOp.create(self.state.ctx, self.loc);
930 try self.addOperation(op.op);
931 return wrapValue(op.getResult());
932 }
933
934 pub fn constantInt(self: *Builder, dtype: DType, value: i64) !Value {
935 try self.ensureOpen();
936 const ty = try scalarType(self.state.ctx, dtype);
937 var op = try ArithDialect.ConstantOp.createInt(self.state.ctx, self.loc, ty, value);
938 try self.addOperation(op.op);
939 return wrapValue(op.getResult());
940 }
941
942 pub fn constantIndex(self: *Builder, value: i64) !Value {
943 try self.ensureOpen();
944 const ty = try ArithDialect.getIndexType(self.state.ctx);
945 var op = try ArithDialect.ConstantOp.createInt(self.state.ctx, self.loc, ty, value);
946 try self.addOperation(op.op);
947 return wrapValue(op.getResult());
948 }
949
950 pub fn constantFloat(self: *Builder, dtype: DType, value: f64) !Value {
951 try self.ensureOpen();
952 const ty = try scalarType(self.state.ctx, dtype);
953 var op = try ArithDialect.ConstantOp.createFloat(self.state.ctx, self.loc, ty, value);
954 try self.addOperation(op.op);
955 return wrapValue(op.getResult());
956 }
957
958 pub fn constantBool(self: *Builder, value: bool) !Value {
959 try self.ensureOpen();
960 var op = try ArithDialect.ConstantOp.createBool(self.state.ctx, self.loc, value);
961 try self.addOperation(op.op);
962 return wrapValue(op.getResult());
963 }
964
965 pub fn add(self: *Builder, lhs: Value, rhs: Value) !Value {
966 try self.ensureOpen();
967 var op = try ArithDialect.AddOp.create(self.state.ctx, self.loc, lhs.ptr, rhs.ptr);
968 try self.addOperation(op.op);
969 return wrapValue(op.getResult());
970 }
971
972 pub fn sub(self: *Builder, lhs: Value, rhs: Value) !Value {
973 try self.ensureOpen();
974 var op = try ArithDialect.SubOp.create(self.state.ctx, self.loc, lhs.ptr, rhs.ptr);
975 try self.addOperation(op.op);
976 return wrapValue(op.getResult());
977 }
978
979 pub fn mul(self: *Builder, lhs: Value, rhs: Value) !Value {
980 try self.ensureOpen();
981 var op = try ArithDialect.MulOp.create(self.state.ctx, self.loc, lhs.ptr, rhs.ptr);
982 try self.addOperation(op.op);
983 return wrapValue(op.getResult());
984 }
985
986 pub fn div(self: *Builder, lhs: Value, rhs: Value) !Value {
987 try self.ensureOpen();
988 var op = try ArithDialect.DivOp.create(self.state.ctx, self.loc, lhs.ptr, rhs.ptr);
989 try self.addOperation(op.op);
990 return wrapValue(op.getResult());
991 }
992
993 pub fn min(self: *Builder, lhs: Value, rhs: Value) !Value {
994 try self.ensureOpen();
995 var op = try ArithDialect.MinOp.create(self.state.ctx, self.loc, lhs.ptr, rhs.ptr);
996 try self.addOperation(op.op);
997 return wrapValue(op.getResult());
998 }
999
1000 pub fn max(self: *Builder, lhs: Value, rhs: Value) !Value {
1001 try self.ensureOpen();
1002 var op = try ArithDialect.MaxOp.create(self.state.ctx, self.loc, lhs.ptr, rhs.ptr);
1003 try self.addOperation(op.op);
1004 return wrapValue(op.getResult());
1005 }
1006
1007 pub fn and_(self: *Builder, lhs: Value, rhs: Value) !Value {
1008 try self.ensureOpen();
1009 var op = try ArithDialect.AndOp.create(self.state.ctx, self.loc, lhs.ptr, rhs.ptr);
1010 try self.addOperation(op.op);
1011 return wrapValue(op.getResult());
1012 }
1013
1014 pub fn or_(self: *Builder, lhs: Value, rhs: Value) !Value {
1015 try self.ensureOpen();
1016 var op = try ArithDialect.OrOp.create(self.state.ctx, self.loc, lhs.ptr, rhs.ptr);
1017 try self.addOperation(op.op);
1018 return wrapValue(op.getResult());
1019 }
1020
1021 pub fn xor(self: *Builder, lhs: Value, rhs: Value) !Value {
1022 try self.ensureOpen();
1023 var op = try ArithDialect.XorOp.create(self.state.ctx, self.loc, lhs.ptr, rhs.ptr);
1024 try self.addOperation(op.op);
1025 return wrapValue(op.getResult());
1026 }
1027
1028 pub fn not(self: *Builder, input: Value) !Value {
1029 try self.ensureOpen();
1030 var op = try ArithDialect.NotOp.create(self.state.ctx, self.loc, input.ptr);
1031 try self.addOperation(op.op);
1032 return wrapValue(op.getResult());
1033 }
1034
1035 pub fn popcount(self: *Builder, input: Value) !Value {
1036 try self.ensureOpen();
1037 var op = try ArithDialect.PopCountOp.create(self.state.ctx, self.loc, input.ptr);
1038 try self.addOperation(op.op);
1039 return wrapValue(op.getResult());
1040 }
1041
1042 pub fn shl(self: *Builder, value: Value, shift: Value) !Value {
1043 try self.ensureOpen();
1044 var op = try ArithDialect.ShlOp.create(self.state.ctx, self.loc, value.ptr, shift.ptr);
1045 try self.addOperation(op.op);
1046 return wrapValue(op.getResult());
1047 }
1048
1049 pub fn shr(self: *Builder, value: Value, shift: Value) !Value {
1050 try self.ensureOpen();
1051 var op = try ArithDialect.ShrOp.create(self.state.ctx, self.loc, value.ptr, shift.ptr);
1052 try self.addOperation(op.op);
1053 return wrapValue(op.getResult());
1054 }
1055
1056 pub fn ushr(self: *Builder, value: Value, shift: Value) !Value {
1057 try self.ensureOpen();
1058 var op = try ArithDialect.UshrOp.create(self.state.ctx, self.loc, value.ptr, shift.ptr);
1059 try self.addOperation(op.op);
1060 return wrapValue(op.getResult());
1061 }
1062
1063 pub fn umulhi(self: *Builder, lhs: Value, rhs: Value) !Value {
1064 try self.ensureOpen();
1065 var op = try ArithDialect.UmulhiOp.create(self.state.ctx, self.loc, lhs.ptr, rhs.ptr);
1066 try self.addOperation(op.op);
1067 return wrapValue(op.getResult());
1068 }
1069
1070 pub fn bitcast(self: *Builder, input: Value, dtype: DType) !Value {
1071 try self.ensureOpen();
1072 const ty = try scalarType(self.state.ctx, dtype);
1073 var op = try ArithDialect.BitcastOp.create(self.state.ctx, self.loc, input.ptr, ty);
1074 try self.addOperation(op.op);
1075 return wrapValue(op.getResult());
1076 }
1077
1078 pub fn neg(self: *Builder, input: Value) !Value {
1079 try self.ensureOpen();
1080 var op = try ArithDialect.NegOp.create(self.state.ctx, self.loc, input.ptr);
1081 try self.addOperation(op.op);
1082 return wrapValue(op.getResult());
1083 }
1084
1085 pub fn abs(self: *Builder, input: Value) !Value {
1086 try self.ensureOpen();
1087 var op = try ArithDialect.AbsOp.create(self.state.ctx, self.loc, input.ptr);
1088 try self.addOperation(op.op);
1089 return wrapValue(op.getResult());
1090 }
1091
1092 pub fn sqrt(self: *Builder, input: Value) !Value {
1093 try self.ensureOpen();
1094 var op = try ArithDialect.SqrtOp.create(self.state.ctx, self.loc, input.ptr);
1095 try self.addOperation(op.op);
1096 return wrapValue(op.getResult());
1097 }
1098
1099 pub fn exp(self: *Builder, input: Value) !Value {
1100 try self.ensureOpen();
1101 var op = try ArithDialect.ExpOp.create(self.state.ctx, self.loc, input.ptr);
1102 try self.addOperation(op.op);
1103 return wrapValue(op.getResult());
1104 }
1105
1106 pub fn log(self: *Builder, input: Value) !Value {
1107 try self.ensureOpen();
1108 var op = try ArithDialect.LogOp.create(self.state.ctx, self.loc, input.ptr);
1109 try self.addOperation(op.op);
1110 return wrapValue(op.getResult());
1111 }
1112
1113 pub fn tanh(self: *Builder, input: Value) !Value {
1114 try self.ensureOpen();
1115 var op = try ArithDialect.TanhOp.create(self.state.ctx, self.loc, input.ptr);
1116 try self.addOperation(op.op);
1117 return wrapValue(op.getResult());
1118 }
1119
1120 pub fn sin(self: *Builder, input: Value) !Value {
1121 try self.ensureOpen();
1122 var op = try ArithDialect.SinOp.create(self.state.ctx, self.loc, input.ptr);
1123 try self.addOperation(op.op);
1124 return wrapValue(op.getResult());
1125 }
1126
1127 pub fn cos(self: *Builder, input: Value) !Value {
1128 try self.ensureOpen();
1129 var op = try ArithDialect.CosOp.create(self.state.ctx, self.loc, input.ptr);
1130 try self.addOperation(op.op);
1131 return wrapValue(op.getResult());
1132 }
1133
1134 pub fn tan(self: *Builder, input: Value) !Value {
1135 try self.ensureOpen();
1136 var op = try ArithDialect.TanOp.create(self.state.ctx, self.loc, input.ptr);
1137 try self.addOperation(op.op);
1138 return wrapValue(op.getResult());
1139 }
1140
1141 pub fn floor(self: *Builder, input: Value) !Value {
1142 try self.ensureOpen();
1143 var op = try ArithDialect.FloorOp.create(self.state.ctx, self.loc, input.ptr);
1144 try self.addOperation(op.op);
1145 return wrapValue(op.getResult());
1146 }
1147
1148 pub fn round(self: *Builder, input: Value) !Value {
1149 try self.ensureOpen();
1150 var op = try ArithDialect.RoundOp.create(self.state.ctx, self.loc, input.ptr);
1151 try self.addOperation(op.op);
1152 return wrapValue(op.getResult());
1153 }
1154
1155 pub fn trunc(self: *Builder, input: Value) !Value {
1156 try self.ensureOpen();
1157 var op = try ArithDialect.TruncOp.create(self.state.ctx, self.loc, input.ptr);
1158 try self.addOperation(op.op);
1159 return wrapValue(op.getResult());
1160 }
1161
1162 pub fn tf32Round(self: *Builder, input: Value) !Value {
1163 try self.ensureOpen();
1164 var op = try ArithDialect.Tf32RoundOp.create(self.state.ctx, self.loc, input.ptr);
1165 try self.addOperation(op.op);
1166 return wrapValue(op.getResult());
1167 }
1168
1169 pub fn pow(self: *Builder, base: Value, exponent: Value) !Value {
1170 try self.ensureOpen();
1171 var op = try ArithDialect.PowOp.create(self.state.ctx, self.loc, base.ptr, exponent.ptr);
1172 try self.addOperation(op.op);
1173 return wrapValue(op.getResult());
1174 }
1175
1176 pub fn atan2(self: *Builder, y: Value, x: Value) !Value {
1177 try self.ensureOpen();
1178 var op = try ArithDialect.Atan2Op.create(self.state.ctx, self.loc, y.ptr, x.ptr);
1179 try self.addOperation(op.op);
1180 return wrapValue(op.getResult());
1181 }
1182
1183 pub fn fma(self: *Builder, a: Value, b: Value, c: Value) !Value {
1184 try self.ensureOpen();
1185 var op = try ArithDialect.FmaOp.create(self.state.ctx, self.loc, a.ptr, b.ptr, c.ptr);
1186 try self.addOperation(op.op);
1187 return wrapValue(op.getResult());
1188 }
1189
1190 pub fn splatVector(self: *Builder, input: Value, width: u32) !Value {
1191 try self.ensureOpen();
1192 const result_ty = try scalarVectorType(self.state.ctx, input.ptr.type, width);
1193 var op = try ArithDialect.SplatOp.create(self.state.ctx, self.loc, input.ptr, result_ty);
1194 try self.addOperation(op.op);
1195 return wrapValue(op.getResult());
1196 }
1197
1198 pub fn shuffleVector(self: *Builder, input: Value, indices: []const i64) !Value {
1199 try self.ensureOpen();
1200 var op = try ArithDialect.VecShuffleOp.create(self.state.ctx, self.loc, input.ptr, input.ptr.type, indices);
1201 try self.addOperation(op.op);
1202 return wrapValue(op.getResult());
1203 }
1204
1205 pub fn compare(self: *Builder, predicate: Compare, lhs: Value, rhs: Value) !Value {
1206 try self.ensureOpen();
1207 const lhs_is_vector = try valueIsVector(lhs.ptr);
1208 const rhs_is_vector = try valueIsVector(rhs.ptr);
1209 if (lhs_is_vector != rhs_is_vector) return error.UnsupportedVectorType;
1210 if (lhs_is_vector) {
1211 var op = try ArithDialect.VecCmpOp.create(self.state.ctx, self.loc, predicate, lhs.ptr, rhs.ptr, lhs.ptr.type);
1212 try self.addOperation(op.op);
1213 return wrapValue(op.getResult());
1214 }
1215 var op = try ArithDialect.CmpOp.create(self.state.ctx, self.loc, predicate, lhs.ptr, rhs.ptr);
1216 try self.addOperation(op.op);
1217 return wrapValue(op.getResult());
1218 }
1219
1220 pub fn select(self: *Builder, condition: Value, true_value: Value, false_value: Value) !Value {
1221 try self.ensureOpen();
1222 var op = try ArithDialect.SelectOp.create(self.state.ctx, self.loc, condition.ptr, true_value.ptr, false_value.ptr);
1223 try self.addOperation(op.op);
1224 return wrapValue(op.getResult());
1225 }
1226
1227 pub fn cast(self: *Builder, input: Value, dtype: DType) !Value {
1228 try self.ensureOpen();
1229 const ty = try scalarType(self.state.ctx, dtype);
1230 var op = try ArithDialect.CastOp.create(self.state.ctx, self.loc, input.ptr, ty);
1231 try self.addOperation(op.op);
1232 return wrapValue(op.getResult());
1233 }
1234
1235 pub fn castIndex(self: *Builder, input: Value) !Value {
1236 try self.ensureOpen();
1237 const ty = try ArithDialect.getIndexType(self.state.ctx);
1238 var op = try ArithDialect.CastOp.create(self.state.ctx, self.loc, input.ptr, ty);
1239 try self.addOperation(op.op);
1240 return wrapValue(op.getResult());
1241 }
1242
1243 pub fn load(self: *Builder, memref: Value, index: Value) !Value {
1244 try self.ensureOpen();
1245 const elem_ty = try memrefElementType(self.state.ctx, memref.ptr.type);
1246 var op = try MemrefDialect.LoadOp.create(self.state.ctx, self.loc, memref.ptr, index.ptr, elem_ty);
1247 try self.addOperation(op.op);
1248 return wrapValue(op.getResult());
1249 }
1250
1251 pub fn loadVector(self: *Builder, memref: Value, index: Value, width: u32) !Value {
1252 try self.ensureOpen();
1253 const result_ty = try memrefVectorType(self.state.ctx, memref.ptr.type, width);
1254 var op = try MemrefDialect.LoadOp.create(self.state.ctx, self.loc, memref.ptr, index.ptr, result_ty);
1255 try self.addOperation(op.op);
1256 return wrapValue(op.getResult());
1257 }
1258
1259 pub fn extractLane(self: *Builder, vector: Value, lane: u32, dtype: DType) !Value {
1260 try self.ensureOpen();
1261 const result_ty = try scalarType(self.state.ctx, dtype);
1262 var op = try ArithDialect.ExtractOp.create(self.state.ctx, self.loc, vector.ptr, @intCast(lane), result_ty);
1263 try self.addOperation(op.op);
1264 return wrapValue(op.getResult());
1265 }
1266
1267 pub fn insertLane(self: *Builder, vector: Value, lane_value: Value, lane: u32) !Value {
1268 try self.ensureOpen();
1269 var op = try ArithDialect.InsertOp.create(self.state.ctx, self.loc, vector.ptr, lane_value.ptr, @intCast(lane));
1270 try self.addOperation(op.op);
1271 return wrapValue(op.getResult());
1272 }
1273
1274 pub fn packVector(self: *Builder, lanes: [4]Value) !Value {
1275 var vector = try self.splatVector(lanes[0], 4);
1276 for (lanes[1..], 1..) |lane_value, lane| {
1277 vector = try self.insertLane(vector, lane_value, @intCast(lane));
1278 }
1279 return vector;
1280 }
1281
1282 pub fn sharedBuffer(self: *Builder, dtype: DType, size: u64) !Value {
1283 try self.ensureOpen();
1284 if (size == 0) return error.InvalidSize;
1285 const elem_ty = try scalarType(self.state.ctx, dtype);
1286 const memref_ty = try MemrefDialect.getMemrefType1D(self.state.ctx, size, elem_ty, .shared);
1287 var op = try MemrefDialect.AllocOp.createStatic(self.state.ctx, self.loc, memref_ty);
1288 try self.addOperation(op.op);
1289 return wrapValue(op.getResult());
1290 }
1291
1292 pub fn dynamicSharedBuffer(self: *Builder, dtype: DType, size: u64, byte_offset: u64) !Value {
1293 try self.ensureOpen();
1294 if (size == 0) return error.InvalidSize;
1295 const elem_ty = try scalarType(self.state.ctx, dtype);
1296 const memref_ty = try MemrefDialect.getMemrefType1D(self.state.ctx, size, elem_ty, .shared);
1297 const dynamic_size = try self.constantIndex(std.math.cast(i64, size) orelse return error.IntegerOutOfBounds);
1298 var op = try MemrefDialect.AllocOp.createDynamic(self.state.ctx, self.loc, dynamic_size.ptr, memref_ty);
1299 var op_owned = true;
1300 errdefer if (op_owned) op.op.erase();
1301 try op.op.setAttr(
1302 gpu.attr_names.dynamic_shared_byte_offset,
1303 try self.state.ctx.getI64Attr(std.math.cast(i64, byte_offset) orelse return error.IntegerOutOfBounds),
1304 );
1305 op_owned = false;
1306 try self.addOperation(op.op);
1307 return wrapValue(op.getResult());
1308 }
1309
1310 pub fn store(self: *Builder, value: Value, memref: Value, index: Value) !void {
1311 try self.ensureOpen();
1312 const op = try MemrefDialect.StoreOp.create(self.state.ctx, self.loc, value.ptr, memref.ptr, index.ptr);
1313 try self.addOperation(op.op);
1314 }
1315
1316 pub fn atomicRmw(
1317 self: *Builder,
1318 kind: dialects.AtomicRmwKind,
1319 value: Value,
1320 memref: Value,
1321 index: Value,
1322 ) !Value {
1323 try self.ensureOpen();
1324 const elem_ty = try memrefElementType(self.state.ctx, memref.ptr.type);
1325 var op = try MemrefDialect.AtomicRmwOp.create(
1326 self.state.ctx,
1327 self.loc,
1328 kind,
1329 value.ptr,
1330 memref.ptr,
1331 index.ptr,
1332 elem_ty,
1333 );
1334 try self.addOperation(op.op);
1335 return wrapValue(op.getResult());
1336 }
1337
1338 pub fn atomicCas(
1339 self: *Builder,
1340 expected: Value,
1341 desired: Value,
1342 memref: Value,
1343 index: Value,
1344 ) !Value {
1345 try self.ensureOpen();
1346 const elem_ty = try memrefElementType(self.state.ctx, memref.ptr.type);
1347 var op = try MemrefDialect.AtomicCasOp.create(
1348 self.state.ctx,
1349 self.loc,
1350 expected.ptr,
1351 desired.ptr,
1352 memref.ptr,
1353 index.ptr,
1354 elem_ty,
1355 );
1356 try self.addOperation(op.op);
1357 return wrapValue(op.getResult());
1358 }
1359
1360 pub fn barrier(self: *Builder, scope: Scope) !void {
1361 try self.ensureOpen();
1362 const op = try GpuDialect.BarrierOp.create(self.state.ctx, self.loc, scope);
1363 try self.addOperation(op.op);
1364 }
1365
1366 pub fn activeMask(self: *Builder) !Value {
1367 try self.ensureOpen();
1368 var op = try GpuDialect.ActiveMaskOp.create(self.state.ctx, self.loc);
1369 try self.addOperation(op.op);
1370 return wrapValue(op.getResult());
1371 }
1372
1373 pub fn syncWarp(self: *Builder) !void {
1374 try self.ensureOpen();
1375 const mask = try self.fullWarpMask();
1376 const op = try GpuDialect.SyncWarpOp.create(self.state.ctx, self.loc, mask.ptr);
1377 try self.addOperation(op.op);
1378 }
1379
1380 pub fn allSync(self: *Builder, predicate: Value) !Value {
1381 try self.ensureOpen();
1382 const mask = try self.fullWarpMask();
1383 var op = try GpuDialect.AllSyncOp.create(self.state.ctx, self.loc, mask.ptr, predicate.ptr);
1384 try self.addOperation(op.op);
1385 return wrapValue(op.getResult());
1386 }
1387
1388 pub fn anySync(self: *Builder, predicate: Value) !Value {
1389 try self.ensureOpen();
1390 const mask = try self.fullWarpMask();
1391 var op = try GpuDialect.AnySyncOp.create(self.state.ctx, self.loc, mask.ptr, predicate.ptr);
1392 try self.addOperation(op.op);
1393 return wrapValue(op.getResult());
1394 }
1395
1396 pub fn ballotSync(self: *Builder, predicate: Value) !Value {
1397 try self.ensureOpen();
1398 const mask = try self.fullWarpMask();
1399 var op = try GpuDialect.BallotSyncOp.create(self.state.ctx, self.loc, mask.ptr, predicate.ptr);
1400 try self.addOperation(op.op);
1401 return wrapValue(op.getResult());
1402 }
1403
1404 pub fn warpReduce(self: *Builder, op_kind: WarpOpKind, value: Value) !Value {
1405 try self.ensureOpen();
1406 const mask = try self.fullWarpMask();
1407 var op = try GpuDialect.WarpReduceOp.create(self.state.ctx, self.loc, op_kind, mask.ptr, value.ptr);
1408 try self.addOperation(op.op);
1409 return wrapValue(op.getResult());
1410 }
1411
1412 pub fn warpScan(self: *Builder, op_kind: WarpOpKind, mode: WarpScanMode, value: Value) !Value {
1413 try self.ensureOpen();
1414 const mask = try self.fullWarpMask();
1415 var op = try GpuDialect.WarpScanOp.create(self.state.ctx, self.loc, op_kind, mode == .inclusive, mask.ptr, value.ptr);
1416 try self.addOperation(op.op);
1417 return wrapValue(op.getResult());
1418 }
1419
1420 pub fn shuffleSync(self: *Builder, mode: ShuffleMode, value: Value, lane_or_delta: Value) !Value {
1421 try self.ensureOpen();
1422 const mask = try self.fullWarpMask();
1423 var op = try GpuDialect.ShflSyncOp.create(self.state.ctx, self.loc, mode, mask.ptr, value.ptr, lane_or_delta.ptr);
1424 try self.addOperation(op.op);
1425 return wrapValue(op.getResult());
1426 }
1427
1428 pub fn mmaSync(self: *Builder, shape: MmaShape, a: [4]Value, b: [2]Value, c: [4]Value) ![4]Value {
1429 try self.ensureOpen();
1430 var op = try GpuDialect.MmaSyncOp.create(
1431 self.state.ctx,
1432 self.loc,
1433 .{ a[0].ptr, a[1].ptr, a[2].ptr, a[3].ptr },
1434 .{ b[0].ptr, b[1].ptr },
1435 .{ c[0].ptr, c[1].ptr, c[2].ptr, c[3].ptr },
1436 shape,
1437 );
1438 try self.addOperation(op.op);
1439 return .{
1440 wrapValue(op.getD(0)),
1441 wrapValue(op.getD(1)),
1442 wrapValue(op.getD(2)),
1443 wrapValue(op.getD(3)),
1444 };
1445 }
1446
1447 pub fn fence(self: *Builder, scope: Scope) !void {
1448 try self.ensureOpen();
1449 const op = try GpuDialect.FenceOp.create(self.state.ctx, self.loc, scope, .seq_cst);
1450 try self.addOperation(op.op);
1451 }
1452
1453 pub fn asyncCopyShared(self: *Builder, dst: Value, dst_index: Value, src: Value, src_index: Value, bytes: u32) !void {
1454 try self.ensureOpen();
1455 const op = try GpuDialect.CpAsyncSharedOp.create(self.state.ctx, self.loc, dst.ptr, dst_index.ptr, src.ptr, src_index.ptr, bytes);
1456 try self.addOperation(op.op);
1457 }
1458
1459 pub fn asyncCopyCommit(self: *Builder) !void {
1460 try self.ensureOpen();
1461 const op = try GpuDialect.CpAsyncCommitOp.create(self.state.ctx, self.loc);
1462 try self.addOperation(op.op);
1463 }
1464
1465 pub fn asyncCopyWait(self: *Builder, groups: u32) !void {
1466 try self.ensureOpen();
1467 const op = try GpuDialect.CpAsyncWaitOp.create(self.state.ctx, self.loc, groups);
1468 try self.addOperation(op.op);
1469 }
1470
1471 fn fullWarpMask(self: *Builder) !Value {
1472 return self.constantInt(.i32, -1);
1473 }
1474
1475 pub fn if_(self: *Builder, condition: Value, result_types: []const Type) !If {
1476 try self.ensureOpen();
1477 const raw_types = try self.rawTypes(result_types);
1478 const op = try ScfDialect.IfOp.create(self.state.ctx, self.loc, condition.ptr, raw_types);
1479 try self.addOperation(op.op);
1480 return .{ .op = op };
1481 }
1482
1483 pub fn for_(self: *Builder, lower: Value, upper: Value, step: Value, init_args: []const Value, result_types: []const Type) !For {
1484 try self.ensureOpen();
1485 const raw_init_args = try self.rawValues(init_args);
1486 const raw_result_types = try self.rawTypes(result_types);
1487 const op = try ScfDialect.ForOp.create(self.state.ctx, self.loc, lower.ptr, upper.ptr, step.ptr, raw_init_args, raw_result_types);
1488 try self.addOperation(op.op);
1489 return .{ .op = op };
1490 }
1491
1492 pub fn forScope(self: *Builder, lower: Value, upper: Value, step: Value, init_args: []const Value, result_types: []const Type) !ForScope {
1493 const loop = try self.for_(lower, upper, step, init_args, result_types);
1494 const previous = self.insertionBlock();
1495 self.setInsertionBlock(loop.bodyBlock());
1496 return .{
1497 .builder = self,
1498 .loop = loop,
1499 .previous = previous,
1500 };
1501 }
1502
1503 pub fn while_(self: *Builder, init_args: []const Value, result_types: []const Type) !While {
1504 try self.ensureOpen();
1505 const raw_init_args = try self.rawValues(init_args);
1506 const raw_result_types = try self.rawTypes(result_types);
1507 const op = try ScfDialect.WhileOp.create(self.state.ctx, self.loc, raw_init_args, raw_result_types);
1508 try self.addOperation(op.op);
1509 return .{ .op = op };
1510 }
1511
1512 pub fn whileScope(self: *Builder, init_args: []const Value, result_types: []const Type) !WhileScope {
1513 const loop = try self.while_(init_args, result_types);
1514 const previous = self.insertionBlock();
1515 self.setInsertionBlock(loop.beforeBlock());
1516 return .{
1517 .builder = self,
1518 .loop = loop,
1519 .previous = previous,
1520 };
1521 }
1522
1523 pub fn condition_(self: *Builder, cond: Value, args: []const Value) !void {
1524 try self.ensureOpen();
1525 const raw_args = try self.rawValues(args);
1526 const op = try ScfDialect.ConditionOp.create(self.state.ctx, self.loc, cond.ptr, raw_args);
1527 try self.addOperation(op.op);
1528 }
1529
1530 pub fn forDo(self: *Builder, lower: Value, upper: Value, step: Value, context: anytype, comptime body: anytype) !For {
1531 var scope = try self.forScope(lower, upper, step, &.{}, &.{});
1532 errdefer scope.abort();
1533 try body(self, scope.inductionVar(), context);
1534 try scope.leave(&.{});
1535 return scope.loop;
1536 }
1537
1538 pub fn fold(self: *Builder, lower: Value, upper: Value, step: Value, initial: Value, context: anytype, comptime body: anytype) !Value {
1539 var scope = try self.forScope(lower, upper, step, &.{initial}, &.{initial.valueType()});
1540 errdefer scope.abort();
1541 const next = try body(self, scope.inductionVar(), scope.iterArg(0).?, context);
1542 try scope.leave(&.{next});
1543 return scope.result(0).?;
1544 }
1545
1546 pub fn yield_(self: *Builder, values: []const Value) !void {
1547 try self.ensureOpen();
1548 const raw_values = try self.rawValues(values);
1549 const op = try ScfDialect.YieldOp.create(self.state.ctx, self.loc, raw_values);
1550 try self.addOperation(op.op);
1551 }
1552
1553 pub fn return_(self: *Builder) !void {
1554 try self.ensureOpen();
1555 const op = try FuncDialect.ReturnOp.create(self.state.ctx, self.loc, &.{});
1556 try self.addOperation(op.op);
1557 self.terminated = true;
1558 }
1559
1560 pub fn finish(self: *Builder) !Kernel {
1561 return self.finishChecking(null);
1562 }
1563
1564 pub fn finishWithDiagnostic(self: *Builder, diagnostic: *VerificationDiagnostic) !Kernel {
1565 return self.finishChecking(diagnostic);
1566 }
1567
1568 fn finishChecking(self: *Builder, diagnostic: ?*VerificationDiagnostic) !Kernel {
1569 if (!self.terminated) return error.MissingTerminator;
1570 if (diagnostic) |captured| {
1571 try choir.backends.contract.verifyModuleWithDiagnostic("accy/kernel/finish", self.state.module, captured);
1572 } else {
1573 try ir.verifyOperation(self.state.module, ir.verify.default_options);
1574 }
1575 if (self.storage.owns_context) {
1576 try ir.dialects.loadDialectSpec(self.state.ctx, GpuDialect.spec);
1577 }
1578 self.storage.activate();
1579 self.finished = true;
1580 return .{
1581 .allocator = self.allocator,
1582 .capacity = self.capacity,
1583 .storage = self.storage,
1584 .state = self.state,
1585 };
1586 }
1587
1588 fn addOperation(self: *Builder, op: *ir.Operation) !void {
1589 errdefer if (op.getBlock() == null) op.erase();
1590 try self.insertion_block.addOperation(op);
1591 }
1592
1593 fn ensureOpen(self: *const Builder) !void {
1594 if (self.finished) return error.AlreadyFinished;
1595 if (self.terminated) return error.AlreadyTerminated;
1596 }
1597
1598 fn rawValues(self: *Builder, values: []const Value) ![]*ir.Value {
1599 try self.storage.require(.{
1600 .parameters = 0,
1601 .kernel_name_bytes = 0,
1602 .temporary_values = values.len,
1603 .temporary_types = 0,
1604 });
1605 const out = self.storage.values[0..values.len];
1606 for (values, out) |value, *raw| raw.* = value.ptr;
1607 return out;
1608 }
1609
1610 fn rawTypes(self: *Builder, types: []const Type) ![]ir.Type {
1611 try self.storage.require(.{
1612 .parameters = 0,
1613 .kernel_name_bytes = 0,
1614 .temporary_values = 0,
1615 .temporary_types = types.len,
1616 });
1617 const out = self.storage.types[0..types.len];
1618 for (types, out) |ty, *raw| raw.* = ty.value;
1619 return out;
1620 }
1621 };
1622
1623 pub const InsertionScope = struct {
1624 builder: *Builder,
1625 previous: Block,
1626 active: bool = true,
1627
1628 pub fn leave(self: *InsertionScope) void {
1629 if (!self.active) return;
1630 self.builder.setInsertionBlock(self.previous);
1631 self.active = false;
1632 }
1633 };
1634
1635 fn wrapValue(value: *ir.Value) Value {
1636 return .{ .ptr = value };
1637 }
1638
1639 fn wrapOptionalValue(value: ?*ir.Value) ?Value {
1640 return wrapValue(value orelse return null);
1641 }
1642
1643 fn wrapBlock(block: *ir.Block) Block {
1644 return .{ .ptr = block };
1645 }
1646
1647 fn validateInputs(limits: Builder.Limits, kernel_name: []const u8, params: []const Param) !void {
1648 if (params.len > limits.parameters) return error.ParameterCapacityExceeded;
1649 if (params.len > limits.temporary_types) return error.TemporaryTypeCapacityExceeded;
1650 if (kernel_name.len > limits.kernel_name_bytes) return error.KernelNameCapacityExceeded;
1651 }
1652
1653 fn requireContextBudget(ctx: *const ir.Context, limits: ir.Context.Limits) !void {
1654 if (limits.maximum_alignment.toByteUnits() > ctx.capacity.storage_alignment.toByteUnits()) {
1655 return error.ContextCapacityExceeded;
1656 }
1657 const usage = ctx.capacityUsage();
1658 try requireSegmentBudget(usage.configuration_tables, limits.configuration.table_bytes);
1659 try requireSegmentBudget(usage.configuration_names, limits.configuration.name_bytes);
1660 try requireSegmentBudget(usage.configuration_interfaces, limits.configuration.interface_bytes);
1661 try requireSegmentBudget(usage.configuration_transactions, limits.configuration.transaction_bytes);
1662 try requireSegmentBudget(usage.type_tables, limits.types.table_bytes);
1663 try requireSegmentBudget(usage.type_keys, limits.types.key_bytes);
1664 try requireSegmentBudget(usage.type_payloads, limits.types.payload_bytes);
1665 try requireSegmentBudget(usage.attribute_tables, limits.attributes.table_bytes);
1666 try requireSegmentBudget(usage.attribute_payloads, limits.attributes.payload_bytes);
1667 try requireSegmentBudget(usage.operation_storage, limits.operations.storage_bytes);
1668 try requireSegmentBudget(usage.operation_nested, limits.operations.nested_bytes);
1669 try requireSegmentBudget(usage.diagnostic_handlers, limits.diagnostics.handler_bytes);
1670 try requireSegmentBudget(usage.diagnostic_payloads, limits.diagnostics.payload_bytes);
1671 try requireSegmentBudget(usage.transient, limits.transient_bytes);
1672 }
1673
1674 fn requireSegmentBudget(usage: anytype, bytes: usize) !void {
1675 if (usage.remainingBytes() < bytes) return error.ContextCapacityExceeded;
1676 }
1677
1678 fn destroyState(state: *State, allocator: std.mem.Allocator, storage: *Storage) void {
1679 state.module.erase();
1680 if (state.decoded) |*decoded| decoded.deinit();
1681 storage.deinit(allocator);
1682 }
1683
1684 fn acquiredBytesForStorage(storage: *const Storage) usize {
1685 return if (storage.owns_context)
1686 storage.capacity.owned_acquired_bytes
1687 else
1688 storage.capacity.borrowed_acquired_bytes;
1689 }
1690
1691 fn placeSlice(
1692 comptime T: type,
1693 count: usize,
1694 cursor: *usize,
1695 allocation_alignment: *usize,
1696 ) error{CapacityOverflow}!usize {
1697 const bytes = std.math.mul(usize, @sizeOf(T), count) catch return error.CapacityOverflow;
1698 const alignment = @alignOf(T);
1699 const mask = std.math.sub(usize, alignment, 1) catch unreachable;
1700 const padded = std.math.add(usize, cursor.*, mask) catch return error.CapacityOverflow;
1701 const offset = padded & ~mask;
1702 cursor.* = std.math.add(usize, offset, bytes) catch return error.CapacityOverflow;
1703 allocation_alignment.* = @max(allocation_alignment.*, alignment);
1704 return offset;
1705 }
1706
1707 fn typedSlice(
1708 comptime T: type,
1709 storage: [*]u8,
1710 offset: usize,
1711 count: usize,
1712 ) []T {
1713 const pointer: [*]T = @ptrCast(@alignCast(storage + offset));
1714 return pointer[0..count];
1715 }
1716
1717 fn bufferType(ctx: *ir.Context, buf: Buffer) !ir.Type {
1718 const elem = try scalarType(ctx, buf.dtype);
1719 if (buf.size) |size| {
1720 return MemrefDialect.getMemrefType1D(ctx, size, elem, buf.space);
1721 }
1722 return MemrefDialect.getMemrefTypeDynamic(ctx, elem, buf.space);
1723 }
1724
1725 fn scalarType(ctx: *ir.Context, dtype: DType) !ir.Type {
1726 return ArithDialect.getScalarType(ctx, switch (dtype) {
1727 .i1 => .bool,
1728 .i8 => .i8,
1729 .i16 => .i16,
1730 .i32 => .i32,
1731 .i64 => .i64,
1732 .u8 => .u8,
1733 .u16 => .u16,
1734 .u32 => .u32,
1735 .u64 => .u64,
1736 .f16 => .f16,
1737 .bf16 => .bf16,
1738 .f32 => .f32,
1739 .f64 => .f64,
1740 .key => .i32,
1741 });
1742 }
1743
1744 fn memrefElementType(ctx: *ir.Context, ty: ir.Type) !ir.Type {
1745 const payload = try ctx.getTypeParamPayload(ty, MemrefDialect.MemrefTypePayload) orelse return error.ExpectedMemref;
1746 return payload.element_type orelse error.ExpectedMemref;
1747 }
1748
1749 fn memrefVectorType(ctx: *ir.Context, ty: ir.Type, width: u32) !ir.Type {
1750 const elem_ty = try memrefElementType(ctx, ty);
1751 const elem_name = elem_ty.getDialectTypeName() orelse return error.ExpectedMemref;
1752 return (try ArithDialect.getVecType(ctx, width, elem_name)) orelse error.UnsupportedVectorType;
1753 }
1754
1755 fn scalarVectorType(ctx: *ir.Context, ty: ir.Type, width: u32) !ir.Type {
1756 const elem_name = ty.getDialectTypeName() orelse return error.UnsupportedVectorType;
1757 return (try ArithDialect.getVecType(ctx, width, elem_name)) orelse error.UnsupportedVectorType;
1758 }
1759
1760 fn valueIsVector(value: *ir.Value) !bool {
1761 const type_name = value.type.getDialectTypeName() orelse return error.UnsupportedType;
1762 return dialects.arith.parseVectorTypeName(type_name) != null;
1763 }
1764
1765 const KernelResourceCounts = struct {
1766 operations: usize,
1767
1768 fn capture(ctx: *const ir.Context) KernelResourceCounts {
1769 return .{
1770 .operations = ctx.operationCount(),
1771 };
1772 }
1773
1774 fn expectEqual(self: KernelResourceCounts, ctx: *const ir.Context) !void {
1775 try testing.expectEqual(self.operations, ctx.operationCount());
1776 }
1777 };
1778
1779 test "Raw Builder capacity follows an independent aligned byte model" {
1780 comptime {
1781 @stardustClaim(
1782 @import("alloc_phase").capacity.witness(Storage, "accy_kernel_raw_builder_capacity"),
1783 null,
1784 null,
1785 null,
1786 null,
1787 null,
1788 null,
1789 );
1790 }
1791
1792 const limits = Builder.Limits{
1793 .context = ir.Context.Limits.standard,
1794 .parameters = 3,
1795 .kernel_name_bytes = 5,
1796 .temporary_values = 7,
1797 .temporary_types = 11,
1798 };
1799 const capacity = try Builder.Capacity.derive(limits);
1800 const context_capacity = try ir.Context.Capacity.derive(limits.context);
1801 var expected: usize = 0;
1802 expected = std.mem.alignForward(usize, expected, @alignOf(ir.Context)) + @sizeOf(ir.Context);
1803 expected = std.mem.alignForward(usize, expected, @alignOf(State)) + @sizeOf(State);
1804 expected = std.mem.alignForward(usize, expected, @alignOf(Param)) + limits.parameters * @sizeOf(Param);
1805 expected = std.mem.alignForward(usize, expected, @alignOf(*ir.Value)) + limits.temporary_values * @sizeOf(*ir.Value);
1806 expected = std.mem.alignForward(usize, expected, @alignOf(ir.Type)) + limits.temporary_types * @sizeOf(ir.Type);
1807 try testing.expectEqual(expected, capacity.storage_bytes);
1808 try testing.expectEqual(expected, capacity.borrowed_acquired_bytes);
1809 try testing.expectEqual(expected + context_capacity.storage_bytes, capacity.owned_acquired_bytes);
1810
1811 var overflow = limits;
1812 overflow.parameters = std.math.maxInt(usize);
1813 try testing.expectError(error.CapacityOverflow, Builder.Capacity.derive(overflow));
1814 }
1815
1816 test "Raw Builder admits parameter and scratch maxima and rejects max plus one before mutation" {
1817 comptime {
1818 @stardustClaim(
1819 @import("alloc_phase").capacity.witness(Storage, "accy_kernel_raw_builder_boundary"),
1820 null,
1821 null,
1822 null,
1823 null,
1824 null,
1825 null,
1826 );
1827 }
1828
1829 var ctx = try ir.Context.init(testing.allocator, ir.Context.Limits.testing);
1830 defer ctx.deinit(testing.allocator);
1831 try dialects.registerChoirDialect(&ctx);
1832
1833 var limits = Builder.Limits.borrowed_testing;
1834 limits.parameters = 2;
1835 limits.kernel_name_bytes = 4;
1836 limits.temporary_values = 2;
1837 limits.temporary_types = 2;
1838 const params = [_]Param{ dynamicBuffer(.f32), dynamicBuffer(.f32) };
1839 const before_rejections = KernelResourceCounts.capture(&ctx);
1840 try testing.expectError(
1841 error.KernelNameCapacityExceeded,
1842 Builder.initBorrowing(testing.allocator, limits, &ctx, "names", ¶ms),
1843 );
1844 try testing.expectError(
1845 error.ParameterCapacityExceeded,
1846 Builder.initBorrowing(testing.allocator, limits, &ctx, "name", &.{ params[0], params[1], params[0] }),
1847 );
1848 try before_rejections.expectEqual(&ctx);
1849
1850 var builder = try Builder.initBorrowing(testing.allocator, limits, &ctx, "name", ¶ms);
1851 defer builder.deinit();
1852 const values = [_]Value{ builder.argument(0), builder.argument(1) };
1853 const types = [_]Type{ values[0].valueType(), values[1].valueType() };
1854 _ = try builder.rawValues(&values);
1855 _ = try builder.rawTypes(&types);
1856 const exact_usage = builder.capacityUsage();
1857 try testing.expectEqual(@as(usize, 2), exact_usage.temporary_values);
1858 try testing.expectEqual(@as(usize, 2), exact_usage.temporary_types);
1859
1860 const before = try choir.operationFingerprint(testing.allocator, builder.state.module);
1861 try testing.expectError(
1862 error.TemporaryValueCapacityExceeded,
1863 builder.rawValues(&.{ values[0], values[1], values[0] }),
1864 );
1865 try testing.expectError(
1866 error.TemporaryTypeCapacityExceeded,
1867 builder.rawTypes(&.{ types[0], types[1], types[0] }),
1868 );
1869 try testing.expectEqual(before, try choir.operationFingerprint(testing.allocator, builder.state.module));
1870 try testing.expectEqualDeep(exact_usage, builder.capacityUsage());
1871 }
1872
1873 test "borrowed Raw Builder verifies explicit remaining Context capacity" {
1874 comptime {
1875 @stardustClaim(
1876 @import("alloc_phase").capacity.witness(Storage, "accy_kernel_raw_builder_context_boundary"),
1877 null,
1878 null,
1879 null,
1880 null,
1881 null,
1882 null,
1883 );
1884 }
1885
1886 var ctx = try ir.Context.init(testing.allocator, ir.Context.Limits.standard);
1887 defer ctx.deinit(testing.allocator);
1888 try dialects.registerChoirDialect(&ctx);
1889 const usage = ctx.capacityUsage();
1890 var context_limits = emptyContextLimits();
1891 context_limits.configuration.table_bytes = usage.configuration_tables.remainingBytes();
1892 try requireContextBudget(&ctx, context_limits);
1893 context_limits.configuration.table_bytes += 1;
1894 try testing.expectError(
1895 error.ContextCapacityExceeded,
1896 requireContextBudget(&ctx, context_limits),
1897 );
1898
1899 var limits = Builder.Limits.borrowed_standard;
1900 limits.context = context_limits;
1901 limits.parameters = 0;
1902 limits.temporary_values = 0;
1903 limits.temporary_types = 0;
1904 var failing = testing.FailingAllocator.init(testing.allocator, .{});
1905 const before = KernelResourceCounts.capture(&ctx);
1906 try testing.expectError(
1907 error.ContextCapacityExceeded,
1908 Builder.initBorrowing(failing.allocator(), limits, &ctx, "", &.{}),
1909 );
1910 try before.expectEqual(&ctx);
1911 try testing.expectEqual(@as(usize, 0), failing.alloc_index);
1912 }
1913
1914 test "Raw Builder construction after initialization makes no backing allocator calls" {
1915 comptime {
1916 @stardustClaim(
1917 @import("alloc_phase").capacity.witness(Storage, "accy_kernel_raw_builder_sealed_construction"),
1918 null,
1919 null,
1920 null,
1921 null,
1922 null,
1923 null,
1924 );
1925 }
1926
1927 var failing = testing.FailingAllocator.init(testing.allocator, .{});
1928 var builder = try Builder.init(
1929 failing.allocator(),
1930 Builder.Limits.standard,
1931 "sealed_raw_builder",
1932 &.{dynamicBuffer(.f32)},
1933 .{},
1934 );
1935 defer builder.deinit();
1936 failing.fail_index = failing.alloc_index;
1937 failing.resize_fail_index = failing.resize_index;
1938
1939 const index = try builder.globalId(.x);
1940 const value = try builder.load(builder.argument(0), index);
1941 const zero = try builder.constantFloat(.f32, 0);
1942 _ = try builder.compare(.gt, value, zero);
1943 try builder.return_();
1944
1945 try testing.expect(!failing.has_induced_failure);
1946 }
1947
1948 test "Raw Builder finish transfers owned state without acquisition" {
1949 comptime {
1950 @stardustClaim(
1951 @import("alloc_phase").capacity.witness(Storage, "accy_kernel_raw_builder_transfer"),
1952 null,
1953 null,
1954 null,
1955 null,
1956 null,
1957 null,
1958 );
1959 }
1960
1961 var failing = testing.FailingAllocator.init(testing.allocator, .{});
1962 var kernel_name = [_]u8{ 't', 'r', 'a', 'n', 's', 'f', 'e', 'r' };
1963 var params = [_]Param{dynamicBuffer(.f32)};
1964 var builder = try Builder.init(
1965 failing.allocator(),
1966 Builder.Limits.standard,
1967 &kernel_name,
1968 ¶ms,
1969 .{},
1970 );
1971 errdefer builder.deinit();
1972 const storage_pointer = builder.storage.storage;
1973 const state_pointer = builder.state;
1974 kernel_name[0] = 'x';
1975 params[0] = scalar(.i32);
1976 try builder.return_();
1977
1978 failing.fail_index = failing.alloc_index;
1979 failing.resize_fail_index = failing.resize_index;
1980 var kernel = try builder.finish();
1981 builder.deinit();
1982 defer kernel.deinit();
1983
1984 try testing.expectEqual(storage_pointer, kernel.storage.storage);
1985 try testing.expectEqual(state_pointer, kernel.state);
1986 try testing.expectEqual(kernel.capacity.owned_acquired_bytes, kernel.acquiredBytes());
1987 try testing.expectEqualStrings("transfer", kernel.func().getName().?);
1988 try testing.expect(paramsEql(kernel.params(), &.{dynamicBuffer(.f32)}));
1989 try testing.expect(!failing.has_induced_failure);
1990 }
1991
1992 fn emptyContextLimits() ir.Context.Limits {
1993 return .{
1994 .maximum_alignment = .@"64",
1995 .configuration = .{
1996 .table_bytes = 0,
1997 .name_bytes = 0,
1998 .interface_bytes = 0,
1999 .transaction_bytes = 0,
2000 },
2001 .types = .{
2002 .table_bytes = 0,
2003 .key_bytes = 0,
2004 .payload_bytes = 0,
2005 },
2006 .attributes = .{
2007 .table_bytes = 0,
2008 .payload_bytes = 0,
2009 },
2010 .operations = .{
2011 .storage_bytes = 0,
2012 .nested_bytes = 0,
2013 },
2014 .diagnostics = .{
2015 .handler_bytes = 0,
2016 .payload_bytes = 0,
2017 },
2018 .transient_bytes = 0,
2019 };
2020 }
2021
2022 test "Builder borrows a reusable Choir context" {
2023 const allocator = testing.allocator;
2024 var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
2025 defer ctx.deinit(allocator);
2026 try dialects.registerChoirDialect(&ctx);
2027 const baseline = KernelResourceCounts.capture(&ctx);
2028
2029 var abandoned = try Builder.initBorrowing(allocator, Builder.Limits.borrowed_testing, &ctx, "borrowed_abandoned", &.{});
2030 try testing.expectEqual(&ctx, abandoned.state.ctx);
2031 try testing.expect(!abandoned.storage.owns_context);
2032 abandoned.deinit();
2033 try baseline.expectEqual(&ctx);
2034
2035 var rejected = try Builder.initBorrowing(allocator, Builder.Limits.borrowed_testing, &ctx, "borrowed_rejected", &.{});
2036 var rejected_live = true;
2037 errdefer if (rejected_live) rejected.deinit();
2038 const zero = try rejected.constantFloat(.f32, 0);
2039 const predicate = try rejected.compare(.gt, zero, zero);
2040 _ = try rejected.if_(predicate, &.{zero.valueType()});
2041 try rejected.return_();
2042 try testing.expectError(error.ScfIfYieldMissing, rejected.finish());
2043 rejected.deinit();
2044 rejected_live = false;
2045 try baseline.expectEqual(&ctx);
2046
2047 var builder = try Builder.initBorrowing(allocator, Builder.Limits.borrowed_testing, &ctx, "borrowed_copy_f32", &.{
2048 dynamicBuffer(.f32),
2049 dynamicBuffer(.f32),
2050 });
2051 errdefer builder.deinit();
2052 const src = builder.argument(0);
2053 const dst = builder.argument(1);
2054 const index = try builder.globalId(.x);
2055 const value = try builder.load(src, index);
2056 try builder.store(value, dst, index);
2057 try builder.return_();
2058
2059 var kernel = try builder.finish();
2060 var kernel_live = true;
2061 errdefer if (kernel_live) kernel.deinit();
2062 try testing.expectEqual(&ctx, kernel.module().context);
2063 try kernel.verify();
2064 kernel.deinit();
2065 kernel_live = false;
2066 try baseline.expectEqual(&ctx);
2067 }
2068
2069 fn checkBorrowedBuilderAllocationFailures(allocator: std.mem.Allocator) !void {
2070 var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
2071 defer ctx.deinit(allocator);
2072 try dialects.registerChoirDialect(&ctx);
2073 const baseline = KernelResourceCounts.capture(&ctx);
2074 defer baseline.expectEqual(&ctx) catch unreachable;
2075
2076 var builder = try Builder.initBorrowing(allocator, Builder.Limits.borrowed_testing, &ctx, "borrowed_failure_f32", &.{
2077 dynamicBuffer(.f32),
2078 dynamicBuffer(.f32),
2079 });
2080 defer builder.deinit();
2081
2082 const src = builder.argument(0);
2083 const dst = builder.argument(1);
2084 const index = try builder.globalId(.x);
2085 const value = try builder.load(src, index);
2086 const zero = try builder.constantFloat(.f32, 0);
2087 _ = try builder.dynamicSharedBuffer(.f32, 16, 32);
2088 _ = try builder.atomicRmw(.add, value, dst, index);
2089 try builder.barrier(.block);
2090 try builder.fence(.device);
2091 const predicate = try builder.compare(.gt, value, zero);
2092 var branch = try builder.if_(predicate, &.{value.valueType()});
2093 {
2094 var scope = try builder.enterBlock(branch.thenBlock());
2095 defer scope.leave();
2096 try builder.yield_(&.{value});
2097 }
2098 {
2099 var scope = try builder.enterBlock(branch.elseBlock().?);
2100 defer scope.leave();
2101 try builder.yield_(&.{zero});
2102 }
2103 try builder.store(branch.result(0).?, dst, index);
2104 try builder.return_();
2105
2106 var kernel = try builder.finish();
2107 defer kernel.deinit();
2108 try kernel.verify();
2109 }
2110
2111 test "borrowed Builder restores context resources on allocation failure" {
2112 try @import("../../../fixture/root.zig").checkAllAllocationFailures(
2113 checkBorrowedBuilderAllocationFailures,
2114 .{},
2115 );
2116 }
2117
2118 test "Builder creates a Choir-backed copy kernel" {
2119 var b = try Builder.init(testing.allocator, Builder.Limits.testing, "copy_f32", &.{
2120 dynamicBuffer(.f32),
2121 dynamicBuffer(.f32),
2122 }, .{});
2123 errdefer b.deinit();
2124
2125 const src = b.argument(0);
2126 const dst = b.argument(1);
2127 const i = try b.globalId(.x);
2128 const x = try b.load(src, i);
2129 try b.store(x, dst, i);
2130 try b.return_();
2131
2132 var kernel = try b.finish();
2133 defer kernel.deinit();
2134
2135 try testing.expectEqualStrings(FuncDialect.FuncOp.operation_name, kernel.func().op.name.name);
2136 try testing.expect(kernel.func().isKernel());
2137 try kernel.verify();
2138 try testing.expectEqual(@as(usize, 2), kernel.params().len);
2139 try testing.expect((Param{ .buffer = .{ .dtype = .f32 } }).eql(kernel.params()[0]));
2140 try testing.expect((Param{ .buffer = .{ .dtype = .f32 } }).eql(kernel.params()[1]));
2141 }
2142
2143 test "Builder creates while loops with carried state" {
2144 var b = try Builder.init(testing.allocator, Builder.Limits.testing, "while_countdown_i32", &.{
2145 dynamicBuffer(.i32),
2146 dynamicBuffer(.i32),
2147 }, .{});
2148 errdefer b.deinit();
2149
2150 const src = b.argument(0);
2151 const dst = b.argument(1);
2152 const i = try b.globalId(.x);
2153 const start = try b.load(src, i);
2154 const zero = try b.constantInt(.i32, 0);
2155 const one = try b.constantInt(.i32, 1);
2156
2157 var scope = try b.whileScope(&.{ start, zero }, &.{ start.valueType(), zero.valueType() });
2158 errdefer scope.abort();
2159 const remaining = scope.beforeArg(0).?;
2160 const total = scope.beforeArg(1).?;
2161 const proceed = try b.compare(.gt, remaining, zero);
2162 try scope.condition(proceed, &.{ remaining, total });
2163 const after_remaining = scope.afterArg(0).?;
2164 const after_total = scope.afterArg(1).?;
2165 const next_total = try b.add(after_total, after_remaining);
2166 const next_remaining = try b.sub(after_remaining, one);
2167 try scope.leave(&.{ next_remaining, next_total });
2168
2169 try b.store(scope.result(1).?, dst, i);
2170 try b.return_();
2171
2172 var kernel = try b.finish();
2173 defer kernel.deinit();
2174 try kernel.verify();
2175 }
2176
2177 test "Builder creates vector memref arithmetic" {
2178 var b = try Builder.init(testing.allocator, Builder.Limits.testing, "vec4_add_f32", &.{
2179 dynamicBuffer(.f32),
2180 dynamicBuffer(.f32),
2181 dynamicBuffer(.f32),
2182 }, .{});
2183 errdefer b.deinit();
2184
2185 const dst = b.argument(0);
2186 const lhs_mem = b.argument(1);
2187 const rhs_mem = b.argument(2);
2188 const zero = try b.constantIndex(0);
2189 const lhs = try b.loadVector(lhs_mem, zero, 4);
2190 const rhs = try b.loadVector(rhs_mem, zero, 4);
2191 const sum = try b.add(lhs, rhs);
2192 try b.store(sum, dst, zero);
2193 try b.return_();
2194
2195 var kernel = try b.finish();
2196 defer kernel.deinit();
2197
2198 try testing.expectEqualStrings("arith.vec4xf32", sum.ptr.type.getDialectTypeName().?);
2199 try testing.expectEqualStrings(FuncDialect.FuncOp.operation_name, kernel.func().op.name.name);
2200 try testing.expectEqual(@as(usize, 3), kernel.params().len);
2201 }
2202
2203 test "Builder compare selects scalar or vector compare operation" {
2204 var b = try Builder.init(testing.allocator, Builder.Limits.testing, "cmp_surface", &.{
2205 dynamicBuffer(.u32),
2206 dynamicBuffer(.u32),
2207 }, .{});
2208 defer b.deinit();
2209
2210 const zero = try b.constantIndex(0);
2211 const one = try b.constantIndex(1);
2212 const scalar_cmp = try b.compare(.lt, zero, one);
2213 const scalar_def: *ir.Operation = @ptrCast(@alignCast(scalar_cmp.ptr.getDefiningOp() orelse return error.ExpectedDefiningOp));
2214 try testing.expectEqualStrings(ArithDialect.CmpOp.operation_name, scalar_def.name.name);
2215 try testing.expectEqualStrings("arith.bool", scalar_cmp.ptr.type.getDialectTypeName().?);
2216
2217 const lhs = try b.loadVector(b.argument(0), zero, 4);
2218 const rhs = try b.loadVector(b.argument(1), zero, 4);
2219 const vector = try b.compare(.ult, lhs, rhs);
2220 const vector_def: *ir.Operation = @ptrCast(@alignCast(vector.ptr.getDefiningOp() orelse return error.ExpectedDefiningOp));
2221 try testing.expectEqualStrings(ArithDialect.VecCmpOp.operation_name, vector_def.name.name);
2222 try testing.expectEqualStrings("arith.vec4xu32", vector.ptr.type.getDialectTypeName().?);
2223 try testing.expectError(error.UnsupportedVectorType, b.compare(.eq, zero, lhs));
2224 }
2225
2226 test "Builder creates vector splat arithmetic" {
2227 var b = try Builder.init(testing.allocator, Builder.Limits.testing, "vec4_splat_u32", &.{
2228 dynamicBuffer(.u32),
2229 scalar(.u32),
2230 }, .{});
2231 errdefer b.deinit();
2232
2233 const dst = b.argument(0);
2234 const bias = b.argument(1);
2235 const zero = try b.constantIndex(0);
2236 const splat = try b.splatVector(bias, 4);
2237 try b.store(splat, dst, zero);
2238 try b.return_();
2239
2240 var kernel = try b.finish();
2241 defer kernel.deinit();
2242
2243 try testing.expectEqualStrings("arith.vec4xu32", splat.ptr.type.getDialectTypeName().?);
2244 try testing.expectEqualStrings(FuncDialect.FuncOp.operation_name, kernel.func().op.name.name);
2245 try testing.expectEqual(@as(usize, 2), kernel.params().len);
2246 }
2247
2248 test "Builder creates vector shuffle arithmetic" {
2249 var b = try Builder.init(testing.allocator, Builder.Limits.testing, "vec4_shuffle_u32", &.{
2250 dynamicBuffer(.u32),
2251 dynamicBuffer(.u32),
2252 }, .{});
2253 errdefer b.deinit();
2254
2255 const dst = b.argument(0);
2256 const src_mem = b.argument(1);
2257 const zero = try b.constantIndex(0);
2258 const src = try b.loadVector(src_mem, zero, 4);
2259 const indices = [_]i64{ 3, 2, 1, 0 };
2260 const shuffled = try b.shuffleVector(src, indices[0..]);
2261 try b.store(shuffled, dst, zero);
2262 try b.return_();
2263
2264 var kernel = try b.finish();
2265 defer kernel.deinit();
2266
2267 try testing.expectEqualStrings("arith.vec4xu32", shuffled.ptr.type.getDialectTypeName().?);
2268 try testing.expectEqualStrings(FuncDialect.FuncOp.operation_name, kernel.func().op.name.name);
2269 try testing.expectEqual(@as(usize, 2), kernel.params().len);
2270 }
2271
2272 test "Builder loads core dialects lazily" {
2273 var b = try Builder.init(testing.allocator, Builder.Limits.testing, "lazy_context", &.{
2274 dynamicBuffer(.f32),
2275 }, .{});
2276 defer b.deinit();
2277
2278 try testing.expect(b.state.ctx.isDialectLoaded("builtin"));
2279 try testing.expect(b.state.ctx.isDialectLoaded("arith"));
2280 try testing.expect(b.state.ctx.isDialectLoaded("memref"));
2281 try testing.expect(b.state.ctx.isDialectLoaded("func"));
2282 try testing.expect(!b.state.ctx.isDialectLoaded("tile"));
2283 try testing.expect(!b.state.ctx.isDialectLoaded("aarch64"));
2284 try testing.expect(!b.state.ctx.isDialectLoaded("x86_64"));
2285 try testing.expect(!b.state.ctx.isDialectLoaded("rc"));
2286 try testing.expect(!b.state.ctx.isDialectLoaded("scf"));
2287 try testing.expect(!b.state.ctx.isDialectLoaded("gpu"));
2288 try testing.expect(!b.state.ctx.isBackendDialect("spirv"));
2289 try testing.expect(!b.state.ctx.isBackendDialect("nvptx"));
2290 try testing.expectError(error.UnknownDialect, b.state.ctx.getOrLoadDialect("nvptx"));
2291 }
2292
2293 test "kernel context requirements preload only the selected dialects" {
2294 var names = [_][]const u8{"scf"};
2295 var builder = try Builder.init(testing.allocator, Builder.Limits.testing, "host", &.{}, .{
2296 .preload = &names,
2297 });
2298 errdefer builder.deinit();
2299 names[0] = "nvptx";
2300 try builder.return_();
2301 var kernel = try builder.finish();
2302 defer kernel.deinit();
2303
2304 const ctx = kernel.module().getContext();
2305 try testing.expect(ctx.isFrozen());
2306 try testing.expect(ctx.isDialectLoaded("gpu"));
2307 try testing.expect(ctx.isDialectLoaded("scf"));
2308 try testing.expect(!ctx.isDialectLoaded("nvptx"));
2309 try testing.expect(!ctx.isDialectLoaded("spirv"));
2310 try testing.expect(!ctx.isBackendDialect("nvptx"));
2311 try testing.expect(!ctx.isBackendDialect("spirv"));
2312 try testing.expectError(error.ContextFrozen, ctx.getOrLoadDialect("nvptx"));
2313 try kernel.verify();
2314 }
2315
2316 test "kernel context requirements release owned storage when preparation fails" {
2317 try testing.expectError(error.UnknownDialect, Builder.init(
2318 testing.allocator,
2319 Builder.Limits.testing,
2320 "missing_dialect",
2321 &.{},
2322 .{ .preload = &.{"unavailable"} },
2323 ));
2324 }
2325
2326 test "Builder supports structured control flow through Choir scf ops" {
2327 var b = try Builder.init(testing.allocator, Builder.Limits.testing, "clamp_positive", &.{
2328 dynamicBuffer(.f32),
2329 dynamicBuffer(.f32),
2330 }, .{});
2331 errdefer b.deinit();
2332
2333 const src = b.argument(0);
2334 const dst = b.argument(1);
2335 const i = try b.globalId(.x);
2336 const x = try b.load(src, i);
2337 const zero = try b.constantFloat(.f32, 0.0);
2338 const pred = try b.compare(.gt, x, zero);
2339 var if_op = try b.if_(pred, &.{x.valueType()});
2340
2341 {
2342 var then_scope = try b.enterBlock(if_op.thenBlock());
2343 defer then_scope.leave();
2344 try b.yield_(&.{x});
2345 }
2346
2347 {
2348 var else_scope = try b.enterBlock(if_op.elseBlock().?);
2349 defer else_scope.leave();
2350 try b.yield_(&.{zero});
2351 }
2352
2353 const clamped = if_op.result(0).?;
2354 try b.store(clamped, dst, i);
2355 try b.return_();
2356
2357 var kernel = try b.finish();
2358 defer kernel.deinit();
2359
2360 try kernel.verify();
2361 }
2362
2363 test "Builder scopes loop callbacks over Choir scf.for" {
2364 var b = try Builder.init(testing.allocator, Builder.Limits.testing, "fill_loop", &.{
2365 dynamicBuffer(.i32),
2366 }, .{});
2367 errdefer b.deinit();
2368
2369 const dst = b.argument(0);
2370 const lower = try b.constantIndex(0);
2371 const upper = try b.constantIndex(4);
2372 const step = try b.constantIndex(1);
2373 _ = try b.forDo(lower, upper, step, dst, struct {
2374 fn body(k: *Builder, index: Value, target: Value) !void {
2375 const one = try k.constantInt(.i32, 1);
2376 try k.store(one, target, index);
2377 }
2378 }.body);
2379
2380 try b.return_();
2381
2382 var kernel = try b.finish();
2383 defer kernel.deinit();
2384
2385 try kernel.verify();
2386 }
2387
2388 test "Builder finishWithDiagnostic localizes Choir verifier failures" {
2389 var b = try Builder.init(testing.allocator, Builder.Limits.testing, "bad_if", &.{}, .{});
2390 defer b.deinit();
2391
2392 const zero = try b.constantFloat(.f32, 0.0);
2393 const pred = try b.compare(.gt, zero, zero);
2394 _ = try b.if_(pred, &.{zero.valueType()});
2395 try b.return_();
2396
2397 var diagnostic: VerificationDiagnostic = .{};
2398 try testing.expectError(error.VerificationFailed, b.finishWithDiagnostic(&diagnostic));
2399 try testing.expect(diagnostic.hasText());
2400 try testing.expect(std.mem.indexOf(u8, diagnostic.text(), "accy/kernel/finish") != null);
2401 try testing.expect(std.mem.indexOf(u8, diagnostic.text(), "scf.if") != null);
2402 }
2403
2404 test "Kernel fingerprint reuses Choir stable hashing" {
2405 var first_builder = try Builder.init(testing.allocator, Builder.Limits.testing, "fingerprint_copy_i32", &.{
2406 dynamicBuffer(.i32),
2407 dynamicBuffer(.i32),
2408 }, .{});
2409 errdefer first_builder.deinit();
2410 const first_src = first_builder.argument(0);
2411 const first_dst = first_builder.argument(1);
2412 const first_index = try first_builder.globalId(.x);
2413 const first_value = try first_builder.load(first_src, first_index);
2414 try first_builder.store(first_value, first_dst, first_index);
2415 try first_builder.return_();
2416 var first = try first_builder.finish();
2417 defer first.deinit();
2418
2419 var second_builder = try Builder.init(testing.allocator, Builder.Limits.testing, "fingerprint_copy_i32", &.{
2420 dynamicBuffer(.i32),
2421 dynamicBuffer(.i32),
2422 }, .{});
2423 errdefer second_builder.deinit();
2424 const second_src = second_builder.argument(0);
2425 const second_dst = second_builder.argument(1);
2426 const second_index = try second_builder.globalId(.x);
2427 const second_value = try second_builder.load(second_src, second_index);
2428 try second_builder.store(second_value, second_dst, second_index);
2429 try second_builder.return_();
2430 var second = try second_builder.finish();
2431 defer second.deinit();
2432
2433 var renamed_builder = try Builder.init(testing.allocator, Builder.Limits.testing, "fingerprint_copy_renamed_i32", &.{
2434 dynamicBuffer(.i32),
2435 dynamicBuffer(.i32),
2436 }, .{});
2437 errdefer renamed_builder.deinit();
2438 const renamed_src = renamed_builder.argument(0);
2439 const renamed_dst = renamed_builder.argument(1);
2440 const renamed_index = try renamed_builder.globalId(.x);
2441 const renamed_value = try renamed_builder.load(renamed_src, renamed_index);
2442 try renamed_builder.store(renamed_value, renamed_dst, renamed_index);
2443 try renamed_builder.return_();
2444 var renamed = try renamed_builder.finish();
2445 defer renamed.deinit();
2446
2447 try testing.expectEqual(try first.fingerprint(testing.allocator), try second.fingerprint(testing.allocator));
2448 try testing.expect(try first.fingerprint(testing.allocator) != try renamed.fingerprint(testing.allocator));
2449 }
2450
2451 test "Builder supports u32 dynamic buffers" {
2452 var b = try Builder.init(testing.allocator, Builder.Limits.testing, "copy_u32", &.{
2453 dynamicBuffer(.u32),
2454 dynamicBuffer(.u32),
2455 }, .{});
2456 errdefer b.deinit();
2457
2458 const src = b.argument(0);
2459 const dst = b.argument(1);
2460 const i = try b.globalId(.x);
2461 const x = try b.load(src, i);
2462 try b.store(x, dst, i);
2463 try b.return_();
2464
2465 var kernel = try b.finish();
2466 defer kernel.deinit();
2467
2468 try testing.expect(paramsEql(kernel.params(), &.{
2469 dynamicBuffer(.u32),
2470 dynamicBuffer(.u32),
2471 }));
2472 try kernel.verify();
2473 }
2474
2475 test "Builder packs scalar lanes into a vector store" {
2476 var b = try Builder.init(testing.allocator, Builder.Limits.testing, "pack_v4", &.{
2477 dynamicBuffer(.f32),
2478 dynamicBuffer(.f32),
2479 }, .{});
2480 errdefer b.deinit();
2481
2482 const src = b.argument(0);
2483 const dst = b.argument(1);
2484 const i = try b.globalId(.x);
2485 const four = try b.constantIndex(4);
2486 const base = try b.mul(i, four);
2487 const loaded = try b.loadVector(src, base, 4);
2488 var lanes: [4]Value = undefined;
2489 for (&lanes, 0..) |*lane, lane_index| {
2490 lane.* = try b.extractLane(loaded, @intCast(lane_index), .f32);
2491 }
2492 const packed_value = try b.packVector(lanes);
2493 try b.store(packed_value, dst, base);
2494 try b.return_();
2495
2496 var kernel = try b.finish();
2497 defer kernel.deinit();
2498 try kernel.verify();
2499 }