lib/accy/src/preparation/kernelization/lowering/shape.zig

daab053ee43316e1809a84551d573ddd1e5bf3d2

  1 const std = @import("std");
  2 const choir_abi = @import("choir_abi");
  3 const generated_abi = @import("abi.zig");
  4 const generated_builder = @import("builder.zig");
  5 const common = @import("common.zig");
  6 const elementwise = @import("elementwise.zig");
  7 const generated_guard = @import("guard.zig");
  8 const generated_name = @import("name.zig");
  9 const generated_schedule = @import("schedule.zig");
 10 
 11 const ir = common.ir;
 12 const dialect_mod = common.dialect_mod;
 13 const kernel_root = common.kernel_root;
 14 const bufferization = common.bufferization;
 15 const kernelization_model = @import("../model/root.zig");
 16 const schedule_planning = common.schedule_planning;
 17 const i64_attr_list_stack_capacity = common.i64_attr_list_stack_capacity;
 18 const bufferSlotById = common.bufferSlotById;
 19 const readI64ListAttrBounded = common.readI64ListAttrBounded;
 20 const readIntegerAttr = common.readIntegerAttr;
 21 const mapKernelBuildError = common.mapKernelBuildError;
 22 const isName = common.isName;
 23 
 24 const ShapeKernel = kernelization_model.ShapeKernel;
 25 const LoweredKernel = kernelization_model.LoweredKernel;
 26 
 27 pub fn lower(
 28     allocator: std.mem.Allocator,
 29     ir_ctx: *ir.Context,
 30     outline: kernelization_model.KernelOutline,
 31     work: schedule_planning.ScheduleWorkItem,
 32     buffer_plan: *const bufferization.BufferPlanAnalysis,
 33 ) common.LoweringError!LoweredKernel {
 34     if (work.kind != .shape) return error.UnsupportedOperation;
 35     if (work.ops.len != 1) return error.UnsupportedOperation;
 36     const kind = kernelForOperation(work.ops[0]) orelse return error.UnsupportedOperation;
 37     const expected_inputs: usize = switch (kind) {
 38         .iota => 0,
 39         .broadcast_in_dim, .reshape, .transpose, .slice => 1,
 40         .pad, .gather => 2,
 41         .scatter => 3,
 42         .concatenate => work.ops[0].getNumOperands(),
 43     };
 44     if (outline.inputCount() != expected_inputs) return error.UnsupportedOperation;
 45     if (!executableShapeDType(work.dtype)) return error.CapabilityMismatch;
 46 
 47     const output_slot = bufferSlotById(buffer_plan, outline.output_slot_id) orelse return error.InvalidArtifact;
 48     if (output_slot.dtype != work.dtype or output_slot.element_count == null) return error.CapabilityMismatch;
 49 
 50     const input_slots = allocator.alloc(bufferization.BufferSlot, outline.inputCount()) catch return error.OutOfMemory;
 51     defer allocator.free(input_slots);
 52     for (outline.input_slot_ids, 0..) |slot_id, index| {
 53         const input_slot = bufferSlotById(buffer_plan, slot_id) orelse return error.InvalidArtifact;
 54         input_slots[index] = input_slot.*;
 55         const indices_input = (kind == .gather or kind == .scatter) and index == 1;
 56         const expected_dtype = if (indices_input) choir_abi.DType.i32 else work.dtype;
 57         if (input_slot.dtype != expected_dtype or input_slot.element_count == null) return error.CapabilityMismatch;
 58     }
 59 
 60     switch (kind) {
 61         .iota => {},
 62         .reshape, .transpose => if (input_slots[0].element_count.? != output_slot.element_count.?) return error.InvalidArtifact,
 63         .broadcast_in_dim, .slice => {},
 64         .pad => {
 65             if (input_slots[1].dims.len != 0) return error.CapabilityMismatch;
 66             if (input_slots[1].element_count.? != 1) return error.CapabilityMismatch;
 67         },
 68         .concatenate => try validateConcatenateSlots(work.ops[0], input_slots, output_slot.*),
 69         .gather => try validateGatherSlots(work.ops[0], input_slots, output_slot.*),
 70         .scatter => try validateScatterSlots(work.ops[0], input_slots, output_slot.*),
 71     }
 72 
 73     const entry_name = try generated_name.shape(allocator, kind, work.id);
 74     errdefer allocator.free(entry_name);
 75 
 76     var abi = switch (kind) {
 77         .gather => try generated_abi.gather(allocator, work.dtype),
 78         .scatter => try generated_abi.scatter(allocator, work.dtype),
 79         else => try generated_abi.flat(allocator, work.dtype, outline.inputCount()),
 80     };
 81     defer abi.deinit(allocator);
 82 
 83     const inputs = allocator.alloc(kernel_root.Value, outline.inputCount()) catch return error.OutOfMemory;
 84     defer allocator.free(inputs);
 85 
 86     return generated_builder.withLaunch(allocator, ir_ctx, work.id, entry_name, abi.params(), generated_schedule.flat(), .{
 87         .abi = abi,
 88         .kind = kind,
 89         .op = work.ops[0],
 90         .input_slots = input_slots,
 91         .output_slot = output_slot.*,
 92         .inputs = inputs,
 93     }, emitShapeBody);
 94 }
 95 
 96 fn executableShapeDType(dtype: choir_abi.DType) bool {
 97     return switch (dtype) {
 98         .i1, .f32, .f64, .i8, .i16, .i32, .u8, .u16, .u32, .i64, .u64, .f16, .bf16, .key => true,
 99     };
100 }
101 
102 fn emitShapeBody(logical: anytype, ctx: anytype) !void {
103     for (ctx.inputs, 0..) |*input, index| {
104         input.* = ctx.abi.input(logical, index);
105     }
106     const logical_index = try logical.index1D("i", ctx.output_slot.element_count.?);
107     try generated_guard.countIndexDo(logical, logical_index, ctx.abi.count(logical), .{
108         .kind = ctx.kind,
109         .op = ctx.op,
110         .input_slots = ctx.input_slots,
111         .output_slot = ctx.output_slot,
112         .inputs = ctx.inputs,
113         .out = ctx.abi.output(logical),
114     }, struct {
115         fn emit(guarded: anytype, index: kernel_root.Index1D, body_ctx: anytype) !void {
116             const value = try shapeKernelOutputValue(
117                 guarded,
118                 body_ctx.kind,
119                 index.index,
120                 body_ctx.op,
121                 body_ctx.input_slots,
122                 body_ctx.output_slot,
123                 body_ctx.inputs,
124             );
125             try guarded.storeIndex(value, body_ctx.out, index);
126         }
127     }.emit);
128 }
129 
130 fn shapeKernelOutputValue(
131     builder: anytype,
132     kind: ShapeKernel,
133     output_index: kernel_root.Value,
134     op: *ir.Operation,
135     input_slots: []const bufferization.BufferSlot,
136     output_slot: bufferization.BufferSlot,
137     inputs: []const kernel_root.Value,
138 ) common.LoweringError!kernel_root.Value {
139     return switch (kind) {
140         .broadcast_in_dim => {
141             const input_index = try elementwise.broadcastInDimSourceIndex(
142                 builder,
143                 output_index,
144                 op,
145                 &input_slots[0],
146             );
147             return builder.loadIndex(inputs[0], input_index) catch |err|
148                 return mapKernelBuildError(err);
149         },
150         .iota => iotaOutputValue(builder, output_index, op, output_slot),
151         .reshape, .transpose, .slice => {
152             const input_index = try shapeKernelInputIndex(builder, kind, output_index, op, input_slots[0], output_slot);
153             return builder.loadIndex(inputs[0], input_index) catch |err| return mapKernelBuildError(err);
154         },
155         .pad => {
156             const mapping = try padSourceMapping(builder, output_index, op, input_slots[0], output_slot);
157             const zero = builder.constantIndex(0) catch |err| return mapKernelBuildError(err);
158             const safe_input_index = builder.select(mapping.in_input, mapping.input_index, zero) catch |err| return mapKernelBuildError(err);
159             const source_value = builder.loadIndex(inputs[0], safe_input_index) catch |err| return mapKernelBuildError(err);
160             const padding_value = builder.loadIndex(inputs[1], zero) catch |err| return mapKernelBuildError(err);
161             return builder.select(mapping.in_input, source_value, padding_value) catch |err| return mapKernelBuildError(err);
162         },
163         .concatenate => concatenateOutputValue(builder, output_index, op, input_slots, output_slot, inputs),
164         .gather => {
165             const source_index = try gatherSourceIndex(builder, output_index, op, input_slots[0], input_slots[1], inputs[1]);
166             return builder.loadIndex(inputs[0], source_index) catch |err| return mapKernelBuildError(err);
167         },
168         .scatter => scatterOutputValue(builder, output_index, op, input_slots, inputs),
169     };
170 }
171 
172 fn scatterOutputValue(
173     builder: anytype,
174     output_index: kernel_root.Value,
175     op: *ir.Operation,
176     input_slots: []const bufferization.BufferSlot,
177     inputs: []const kernel_root.Value,
178 ) common.LoweringError!kernel_root.Value {
179     const data_slot = input_slots[0];
180     const indices_slot = input_slots[1];
181     const axis = try axisForOperation(op, "axis", data_slot.dims.len);
182     const axis_size = data_slot.dims[axis];
183     if (axis_size <= 0) return error.CapabilityMismatch;
184     const inner = try dimsProductI64(data_slot.dims[axis + 1 ..]);
185     const update_count = indices_slot.dims[0];
186     if (update_count <= 0) return error.CapabilityMismatch;
187     const axis_block = std.math.mul(i64, axis_size, inner) catch return error.CapabilityMismatch;
188 
189     const inner_value = builder.constantIndex(inner) catch |err| return mapKernelBuildError(err);
190     const axis_block_value = builder.constantIndex(axis_block) catch |err| return mapKernelBuildError(err);
191 
192     const outer_coord = builder.div(output_index, axis_block_value) catch |err| return mapKernelBuildError(err);
193     const outer_consumed = builder.mul(outer_coord, axis_block_value) catch |err| return mapKernelBuildError(err);
194     const axis_rem = builder.sub(output_index, outer_consumed) catch |err| return mapKernelBuildError(err);
195     const axis_coord = builder.div(axis_rem, inner_value) catch |err| return mapKernelBuildError(err);
196     const axis_consumed = builder.mul(axis_coord, inner_value) catch |err| return mapKernelBuildError(err);
197     const within = builder.sub(axis_rem, axis_consumed) catch |err| return mapKernelBuildError(err);
198 
199     const update_block = std.math.mul(i64, update_count, inner) catch return error.CapabilityMismatch;
200     const update_block_value = builder.constantIndex(update_block) catch |err| return mapKernelBuildError(err);
201     const update_base = builder.mul(outer_coord, update_block_value) catch |err| return mapKernelBuildError(err);
202 
203     const fold_lower = builder.constantIndex(0) catch |err| return mapKernelBuildError(err);
204     const fold_upper = builder.constantIndex(update_count) catch |err| return mapKernelBuildError(err);
205     const fold_step = builder.constantIndex(1) catch |err| return mapKernelBuildError(err);
206     const initial = builder.loadIndex(inputs[0], output_index) catch |err| return mapKernelBuildError(err);
207 
208     return builder.fold(fold_lower, fold_upper, fold_step, initial, .{
209         .axis_coord = axis_coord,
210         .within = within,
211         .update_base = update_base,
212         .inner_value = inner_value,
213         .indices = inputs[1],
214         .updates = inputs[2],
215     }, struct {
216         fn apply(fold_inner: anytype, update_position: kernel_root.Value, current: kernel_root.Value, fold_ctx: anytype) !kernel_root.Value {
217             const loaded = try fold_inner.loadIndex(fold_ctx.indices, update_position);
218             const target = try fold_inner.castIndex(loaded);
219             const match = try fold_inner.compare(.eq, target, fold_ctx.axis_coord);
220             const update_row = try fold_inner.mul(update_position, fold_ctx.inner_value);
221             const update_offset = try fold_inner.add(fold_ctx.update_base, update_row);
222             const update_index = try fold_inner.add(update_offset, fold_ctx.within);
223             const candidate = try fold_inner.loadIndex(fold_ctx.updates, update_index);
224             return fold_inner.select(match, candidate, current);
225         }
226     }.apply) catch |err| return mapKernelBuildError(err);
227 }
228 
229 fn validateScatterSlots(
230     op: *ir.Operation,
231     input_slots: []const bufferization.BufferSlot,
232     output_slot: bufferization.BufferSlot,
233 ) common.LoweringError!void {
234     const data_slot = input_slots[0];
235     const indices_slot = input_slots[1];
236     const updates_slot = input_slots[2];
237     if (data_slot.dims.len == 0) return error.CapabilityMismatch;
238     if (indices_slot.dims.len != 1) return error.CapabilityMismatch;
239     const axis = try axisForOperation(op, "axis", data_slot.dims.len);
240     if (!std.mem.eql(i64, output_slot.dims, data_slot.dims)) return error.InvalidArtifact;
241     if (updates_slot.dims.len != data_slot.dims.len) return error.InvalidArtifact;
242     if (!std.mem.eql(i64, updates_slot.dims[0..axis], data_slot.dims[0..axis])) return error.InvalidArtifact;
243     if (updates_slot.dims[axis] != indices_slot.dims[0]) return error.InvalidArtifact;
244     if (!std.mem.eql(i64, updates_slot.dims[axis + 1 ..], data_slot.dims[axis + 1 ..])) return error.InvalidArtifact;
245 }
246 
247 fn gatherSourceIndex(
248     builder: anytype,
249     output_index: kernel_root.Value,
250     op: *ir.Operation,
251     data_slot: bufferization.BufferSlot,
252     indices_slot: bufferization.BufferSlot,
253     indices_buffer: kernel_root.Value,
254 ) common.LoweringError!kernel_root.Value {
255     const axis = try axisForOperation(op, "axis", data_slot.dims.len);
256     const axis_size = data_slot.dims[axis];
257     if (axis_size <= 0) return error.CapabilityMismatch;
258     const inner = try dimsProductI64(data_slot.dims[axis + 1 ..]);
259     const gathered = try dimsProductI64(indices_slot.dims);
260     if (gathered <= 0) return error.CapabilityMismatch;
261     const block = std.math.mul(i64, gathered, inner) catch return error.CapabilityMismatch;
262 
263     const inner_value = builder.constantIndex(inner) catch |err| return mapKernelBuildError(err);
264     const block_value = builder.constantIndex(block) catch |err| return mapKernelBuildError(err);
265 
266     const outer_coord = builder.div(output_index, block_value) catch |err| return mapKernelBuildError(err);
267     const outer_consumed = builder.mul(outer_coord, block_value) catch |err| return mapKernelBuildError(err);
268     const block_rem = builder.sub(output_index, outer_consumed) catch |err| return mapKernelBuildError(err);
269     const position = builder.div(block_rem, inner_value) catch |err| return mapKernelBuildError(err);
270     const position_consumed = builder.mul(position, inner_value) catch |err| return mapKernelBuildError(err);
271     const within = builder.sub(block_rem, position_consumed) catch |err| return mapKernelBuildError(err);
272 
273     const loaded = builder.loadIndex(indices_buffer, position) catch |err| return mapKernelBuildError(err);
274     const zero_i32 = builder.constantInt(.i32, 0) catch |err| return mapKernelBuildError(err);
275     const limit_value = @min(axis_size - 1, @as(i64, std.math.maxInt(i32)));
276     const limit_i32 = builder.constantInt(.i32, limit_value) catch |err| return mapKernelBuildError(err);
277     const lower_clamped_i32 = builder.max(loaded, zero_i32) catch |err| return mapKernelBuildError(err);
278     const clamped_i32 = builder.min(lower_clamped_i32, limit_i32) catch |err| return mapKernelBuildError(err);
279     const clamped = builder.castIndex(clamped_i32) catch |err| return mapKernelBuildError(err);
280 
281     const axis_stride = std.math.mul(i64, axis_size, inner) catch return error.CapabilityMismatch;
282     const axis_stride_value = builder.constantIndex(axis_stride) catch |err| return mapKernelBuildError(err);
283     const outer_offset = builder.mul(outer_coord, axis_stride_value) catch |err| return mapKernelBuildError(err);
284     const gathered_offset = builder.mul(clamped, inner_value) catch |err| return mapKernelBuildError(err);
285     const partial = builder.add(outer_offset, gathered_offset) catch |err| return mapKernelBuildError(err);
286     return builder.add(partial, within) catch |err| return mapKernelBuildError(err);
287 }
288 
289 fn validateGatherSlots(
290     op: *ir.Operation,
291     input_slots: []const bufferization.BufferSlot,
292     output_slot: bufferization.BufferSlot,
293 ) common.LoweringError!void {
294     const data_slot = input_slots[0];
295     const indices_slot = input_slots[1];
296     if (data_slot.dims.len == 0) return error.CapabilityMismatch;
297     if (indices_slot.dims.len == 0) return error.CapabilityMismatch;
298     const axis = try axisForOperation(op, "axis", data_slot.dims.len);
299     if (output_slot.dims.len != data_slot.dims.len - 1 + indices_slot.dims.len) return error.InvalidArtifact;
300     if (!std.mem.eql(i64, output_slot.dims[0..axis], data_slot.dims[0..axis])) return error.InvalidArtifact;
301     if (!std.mem.eql(i64, output_slot.dims[axis .. axis + indices_slot.dims.len], indices_slot.dims)) return error.InvalidArtifact;
302     if (!std.mem.eql(i64, output_slot.dims[axis + indices_slot.dims.len ..], data_slot.dims[axis + 1 ..])) return error.InvalidArtifact;
303 }
304 
305 fn dimsProductI64(dims: []const i64) common.LoweringError!i64 {
306     var product: i64 = 1;
307     for (dims) |dim| {
308         if (dim <= 0) return error.CapabilityMismatch;
309         product = std.math.mul(i64, product, dim) catch return error.CapabilityMismatch;
310     }
311     return product;
312 }
313 
314 fn shapeKernelInputIndex(
315     builder: anytype,
316     kind: ShapeKernel,
317     output_index: kernel_root.Value,
318     op: *ir.Operation,
319     input_slot: bufferization.BufferSlot,
320     output_slot: bufferization.BufferSlot,
321 ) common.LoweringError!kernel_root.Value {
322     return switch (kind) {
323         .reshape => output_index,
324         .transpose => transposeSourceIndex(builder, output_index, op, input_slot, output_slot),
325         .slice => sliceSourceIndex(builder, output_index, op, input_slot, output_slot),
326         .broadcast_in_dim, .iota, .pad, .concatenate, .gather, .scatter => unreachable,
327     };
328 }
329 
330 fn validateConcatenateSlots(
331     op: *ir.Operation,
332     input_slots: []const bufferization.BufferSlot,
333     output_slot: bufferization.BufferSlot,
334 ) common.LoweringError!void {
335     if (input_slots.len == 0) return error.UnsupportedOperation;
336     if (output_slot.row_major_strides == null) return error.CapabilityMismatch;
337     const axis = try axisForOperation(op, "dimension", output_slot.dims.len);
338     if (output_slot.dims[axis] <= 0) return error.CapabilityMismatch;
339 
340     var axis_total: i64 = 0;
341     for (input_slots) |input_slot| {
342         if (input_slot.dims.len != output_slot.dims.len) return error.InvalidArtifact;
343         if (input_slot.row_major_strides == null) return error.CapabilityMismatch;
344         if (input_slot.dims[axis] <= 0) return error.CapabilityMismatch;
345         axis_total = std.math.add(i64, axis_total, input_slot.dims[axis]) catch return error.CapabilityMismatch;
346         for (input_slot.dims, 0..) |dim, dim_index| {
347             if (dim_index == axis) continue;
348             if (dim != output_slot.dims[dim_index]) return error.InvalidArtifact;
349         }
350     }
351     if (axis_total != output_slot.dims[axis]) return error.InvalidArtifact;
352 }
353 
354 fn iotaOutputValue(
355     builder: anytype,
356     output_index: kernel_root.Value,
357     op: *ir.Operation,
358     output_slot: bufferization.BufferSlot,
359 ) common.LoweringError!kernel_root.Value {
360     const axis = try axisForOperation(op, "iota_dimension", output_slot.dims.len);
361     const coord = try linearCoordinate(builder, output_index, output_slot, axis);
362     return builder.cast(coord, output_slot.dtype) catch |err| return mapKernelBuildError(err);
363 }
364 
365 fn concatenateOutputValue(
366     builder: anytype,
367     output_index: kernel_root.Value,
368     op: *ir.Operation,
369     input_slots: []const bufferization.BufferSlot,
370     output_slot: bufferization.BufferSlot,
371     inputs: []const kernel_root.Value,
372 ) common.LoweringError!kernel_root.Value {
373     if (input_slots.len == 0 or input_slots.len != inputs.len) return error.InvalidArtifact;
374     const axis = try axisForOperation(op, "dimension", output_slot.dims.len);
375     const axis_coord = try linearCoordinate(builder, output_index, output_slot, axis);
376     const zero = builder.constantIndex(0) catch |err| return mapKernelBuildError(err);
377 
378     var axis_start: i64 = 0;
379     var selected_value: ?kernel_root.Value = null;
380     for (input_slots, inputs) |input_slot, input| {
381         const axis_end = std.math.add(i64, axis_start, input_slot.dims[axis]) catch return error.CapabilityMismatch;
382         const start_value = builder.constantIndex(axis_start) catch |err| return mapKernelBuildError(err);
383         const end_value = builder.constantIndex(axis_end) catch |err| return mapKernelBuildError(err);
384         const at_or_after_start = builder.compare(.ge, axis_coord, start_value) catch |err| return mapKernelBuildError(err);
385         const before_end = builder.compare(.lt, axis_coord, end_value) catch |err| return mapKernelBuildError(err);
386         const in_segment = builder.and_(at_or_after_start, before_end) catch |err| return mapKernelBuildError(err);
387         const source_index = try concatenateSourceIndex(builder, output_index, output_slot, input_slot, axis, axis_start);
388         const safe_source_index = builder.select(in_segment, source_index, zero) catch |err| return mapKernelBuildError(err);
389         const input_value = builder.loadIndex(input, safe_source_index) catch |err| return mapKernelBuildError(err);
390         selected_value = if (selected_value) |current|
391             builder.select(in_segment, input_value, current) catch |err| return mapKernelBuildError(err)
392         else
393             input_value;
394         axis_start = axis_end;
395     }
396     return selected_value orelse error.InvalidArtifact;
397 }
398 
399 fn concatenateSourceIndex(
400     builder: anytype,
401     output_index: kernel_root.Value,
402     output_slot: bufferization.BufferSlot,
403     input_slot: bufferization.BufferSlot,
404     axis: usize,
405     axis_start: i64,
406 ) common.LoweringError!kernel_root.Value {
407     const input_strides = input_slot.row_major_strides orelse return error.CapabilityMismatch;
408     var source_index = builder.constantIndex(0) catch |err| return mapKernelBuildError(err);
409     for (input_slot.dims, 0..) |_, dim_index| {
410         var coord = try linearCoordinate(builder, output_index, output_slot, dim_index);
411         if (dim_index == axis and axis_start != 0) {
412             const start_value = builder.constantIndex(axis_start) catch |err| return mapKernelBuildError(err);
413             coord = builder.sub(coord, start_value) catch |err| return mapKernelBuildError(err);
414         }
415         const stride = std.math.cast(i64, input_strides[dim_index]) orelse return error.CapabilityMismatch;
416         const term = if (stride == 1) coord else blk: {
417             const stride_value = builder.constantIndex(stride) catch |err| return mapKernelBuildError(err);
418             break :blk builder.mul(coord, stride_value) catch |err| return mapKernelBuildError(err);
419         };
420         source_index = builder.add(source_index, term) catch |err| return mapKernelBuildError(err);
421     }
422     return source_index;
423 }
424 
425 fn linearCoordinate(
426     builder: anytype,
427     linear_index: kernel_root.Value,
428     slot: bufferization.BufferSlot,
429     dim_index: usize,
430 ) common.LoweringError!kernel_root.Value {
431     const strides = slot.row_major_strides orelse return error.CapabilityMismatch;
432     if (dim_index >= slot.dims.len or dim_index >= strides.len) return error.InvalidArtifact;
433     const dim = slot.dims[dim_index];
434     if (dim <= 0) return error.CapabilityMismatch;
435     const stride = std.math.cast(i64, strides[dim_index]) orelse return error.CapabilityMismatch;
436     const quotient = if (stride == 1) linear_index else blk: {
437         const stride_value = builder.constantIndex(stride) catch |err| return mapKernelBuildError(err);
438         break :blk builder.div(linear_index, stride_value) catch |err| return mapKernelBuildError(err);
439     };
440     const dim_value = builder.constantIndex(dim) catch |err| return mapKernelBuildError(err);
441     const whole_dim_count = builder.div(quotient, dim_value) catch |err| return mapKernelBuildError(err);
442     const consumed = builder.mul(whole_dim_count, dim_value) catch |err| return mapKernelBuildError(err);
443     return builder.sub(quotient, consumed) catch |err| return mapKernelBuildError(err);
444 }
445 
446 fn axisForOperation(
447     op: *ir.Operation,
448     attr_name: []const u8,
449     rank: usize,
450 ) common.LoweringError!usize {
451     if (rank == 0) return error.CapabilityMismatch;
452     const axis_value = try readIntegerAttr(op, attr_name);
453     if (axis_value < 0) return error.InvalidArtifact;
454     const axis = std.math.cast(usize, axis_value) orelse return error.InvalidArtifact;
455     if (axis >= rank) return error.InvalidArtifact;
456     return axis;
457 }
458 
459 const PadSourceMapping = struct {
460     in_input: kernel_root.Value,
461     input_index: kernel_root.Value,
462 };
463 
464 fn padSourceMapping(
465     builder: anytype,
466     output_index: kernel_root.Value,
467     op: *ir.Operation,
468     input_slot: bufferization.BufferSlot,
469     output_slot: bufferization.BufferSlot,
470 ) common.LoweringError!PadSourceMapping {
471     var low_buffer: [i64_attr_list_stack_capacity]i64 = undefined;
472     var high_buffer: [i64_attr_list_stack_capacity]i64 = undefined;
473     var interior_buffer: [i64_attr_list_stack_capacity]i64 = undefined;
474     const edge_low = try readI64ListAttrBounded(op, "edge_low", dialect_mod.AccyDialect.PadOp.dialectAttrName("edge_low"), &low_buffer);
475     const edge_high = try readI64ListAttrBounded(op, "edge_high", dialect_mod.AccyDialect.PadOp.dialectAttrName("edge_high"), &high_buffer);
476     const interior = try readI64ListAttrBounded(op, "interior", dialect_mod.AccyDialect.PadOp.dialectAttrName("interior"), &interior_buffer);
477     if (edge_low.len != input_slot.dims.len or edge_high.len != input_slot.dims.len or interior.len != input_slot.dims.len) {
478         return error.InvalidArtifact;
479     }
480     if (output_slot.dims.len != input_slot.dims.len) return error.InvalidArtifact;
481     const input_strides = input_slot.row_major_strides orelse return error.CapabilityMismatch;
482     const output_strides = output_slot.row_major_strides orelse return error.CapabilityMismatch;
483     if (input_strides.len != input_slot.dims.len or output_strides.len != output_slot.dims.len) return error.InvalidArtifact;
484 
485     var source_index = builder.constantIndex(0) catch |err| return mapKernelBuildError(err);
486     var in_input = builder.constantBool(true) catch |err| return mapKernelBuildError(err);
487     for (output_slot.dims, 0..) |dim, dim_index| {
488         if (dim <= 0 or input_slot.dims[dim_index] <= 0) return error.CapabilityMismatch;
489         const low = edge_low[dim_index];
490         const high = edge_high[dim_index];
491         const gap = interior[dim_index];
492         if (gap != 0) return error.CapabilityMismatch;
493         const expected_dim = low + input_slot.dims[dim_index] + high;
494         if (dim != expected_dim) return error.InvalidArtifact;
495 
496         const coord = try linearCoordinate(builder, output_index, output_slot, dim_index);
497 
498         const low_value = builder.constantIndex(low) catch |err| return mapKernelBuildError(err);
499         const upper_value = builder.constantIndex(low + input_slot.dims[dim_index]) catch |err| return mapKernelBuildError(err);
500         const at_or_after_low = builder.compare(.ge, coord, low_value) catch |err| return mapKernelBuildError(err);
501         const before_upper = builder.compare(.lt, coord, upper_value) catch |err| return mapKernelBuildError(err);
502         const dim_in_input = builder.and_(at_or_after_low, before_upper) catch |err| return mapKernelBuildError(err);
503         in_input = builder.and_(in_input, dim_in_input) catch |err| return mapKernelBuildError(err);
504 
505         const source_coord = if (low == 0) coord else builder.sub(coord, low_value) catch |err| return mapKernelBuildError(err);
506         const input_stride = std.math.cast(i64, input_strides[dim_index]) orelse return error.CapabilityMismatch;
507         const term = if (input_stride == 1) source_coord else blk: {
508             const stride_value = builder.constantIndex(input_stride) catch |err| return mapKernelBuildError(err);
509             break :blk builder.mul(source_coord, stride_value) catch |err| return mapKernelBuildError(err);
510         };
511         source_index = builder.add(source_index, term) catch |err| return mapKernelBuildError(err);
512     }
513     return .{
514         .in_input = in_input,
515         .input_index = source_index,
516     };
517 }
518 
519 fn transposeSourceIndex(
520     builder: anytype,
521     output_index: kernel_root.Value,
522     op: *ir.Operation,
523     input_slot: bufferization.BufferSlot,
524     output_slot: bufferization.BufferSlot,
525 ) common.LoweringError!kernel_root.Value {
526     var permutation_buffer: [i64_attr_list_stack_capacity]i64 = undefined;
527     const permutation = try readI64ListAttrBounded(op, "permutation", dialect_mod.AccyDialect.TransposeOp.dialectAttrName("permutation"), &permutation_buffer);
528     if (permutation.len != input_slot.dims.len or permutation.len != output_slot.dims.len) return error.InvalidArtifact;
529     try validatePermutation(permutation);
530     const input_strides = input_slot.row_major_strides orelse return error.CapabilityMismatch;
531     const output_strides = output_slot.row_major_strides orelse return error.CapabilityMismatch;
532     if (input_strides.len != input_slot.dims.len or output_strides.len != output_slot.dims.len) return error.InvalidArtifact;
533 
534     var source_index = builder.constantIndex(0) catch |err| return mapKernelBuildError(err);
535     for (output_slot.dims, 0..) |dim, output_dim| {
536         if (dim <= 0) return error.CapabilityMismatch;
537         const coord = try linearCoordinate(builder, output_index, output_slot, output_dim);
538 
539         const input_dim = std.math.cast(usize, permutation[output_dim]) orelse return error.InvalidArtifact;
540         if (input_slot.dims[input_dim] != dim) return error.InvalidArtifact;
541         const input_stride = std.math.cast(i64, input_strides[input_dim]) orelse return error.CapabilityMismatch;
542         const term = if (input_stride == 1) coord else blk: {
543             const stride_value = builder.constantIndex(input_stride) catch |err| return mapKernelBuildError(err);
544             break :blk builder.mul(coord, stride_value) catch |err| return mapKernelBuildError(err);
545         };
546         source_index = builder.add(source_index, term) catch |err| return mapKernelBuildError(err);
547     }
548     return source_index;
549 }
550 
551 fn validatePermutation(permutation: []const i64) common.LoweringError!void {
552     for (permutation, 0..) |dim, index| {
553         const dim_index = std.math.cast(usize, dim) orelse return error.InvalidArtifact;
554         if (dim_index >= permutation.len) return error.InvalidArtifact;
555         for (permutation[0..index]) |previous| {
556             const previous_index = std.math.cast(usize, previous) orelse return error.InvalidArtifact;
557             if (previous_index == dim_index) return error.InvalidArtifact;
558         }
559     }
560 }
561 
562 fn sliceSourceIndex(
563     builder: anytype,
564     output_index: kernel_root.Value,
565     op: *ir.Operation,
566     input_slot: bufferization.BufferSlot,
567     output_slot: bufferization.BufferSlot,
568 ) common.LoweringError!kernel_root.Value {
569     var starts_buffer: [i64_attr_list_stack_capacity]i64 = undefined;
570     var limits_buffer: [i64_attr_list_stack_capacity]i64 = undefined;
571     var strides_buffer: [i64_attr_list_stack_capacity]i64 = undefined;
572     const starts = try readI64ListAttrBounded(op, "starts", dialect_mod.AccyDialect.SliceOp.dialectAttrName("starts"), &starts_buffer);
573     const limits = try readI64ListAttrBounded(op, "limits", dialect_mod.AccyDialect.SliceOp.dialectAttrName("limits"), &limits_buffer);
574     const strides = try readI64ListAttrBounded(op, "strides", dialect_mod.AccyDialect.SliceOp.dialectAttrName("strides"), &strides_buffer);
575     if (starts.len != input_slot.dims.len or limits.len != input_slot.dims.len or strides.len != input_slot.dims.len) {
576         return error.InvalidArtifact;
577     }
578     if (output_slot.dims.len != input_slot.dims.len) return error.InvalidArtifact;
579     const input_strides = input_slot.row_major_strides orelse return error.CapabilityMismatch;
580     const output_strides = output_slot.row_major_strides orelse return error.CapabilityMismatch;
581     if (input_strides.len != input_slot.dims.len or output_strides.len != output_slot.dims.len) return error.InvalidArtifact;
582 
583     var source_index = builder.constantIndex(0) catch |err| return mapKernelBuildError(err);
584     for (output_slot.dims, 0..) |dim, dim_index| {
585         if (dim <= 0 or input_slot.dims[dim_index] <= 0) return error.CapabilityMismatch;
586         const start = starts[dim_index];
587         const limit = limits[dim_index];
588         const stride = strides[dim_index];
589         if (stride <= 0 or start < 0 or limit < start or limit > input_slot.dims[dim_index]) return error.InvalidArtifact;
590         const span = limit - start;
591         const expected_dim = @divTrunc(span + stride - 1, stride);
592         if (dim != expected_dim) return error.InvalidArtifact;
593 
594         const coord = try linearCoordinate(builder, output_index, output_slot, dim_index);
595 
596         const stepped_coord = if (stride == 1) coord else blk: {
597             const stride_value = builder.constantIndex(stride) catch |err| return mapKernelBuildError(err);
598             break :blk builder.mul(coord, stride_value) catch |err| return mapKernelBuildError(err);
599         };
600         const source_coord = if (start == 0) stepped_coord else blk: {
601             const start_value = builder.constantIndex(start) catch |err| return mapKernelBuildError(err);
602             break :blk builder.add(start_value, stepped_coord) catch |err| return mapKernelBuildError(err);
603         };
604         const input_stride = std.math.cast(i64, input_strides[dim_index]) orelse return error.CapabilityMismatch;
605         const term = if (input_stride == 1) source_coord else blk: {
606             const stride_value = builder.constantIndex(input_stride) catch |err| return mapKernelBuildError(err);
607             break :blk builder.mul(source_coord, stride_value) catch |err| return mapKernelBuildError(err);
608         };
609         source_index = builder.add(source_index, term) catch |err| return mapKernelBuildError(err);
610     }
611     return source_index;
612 }
613 
614 pub fn kernelForOperation(op: *ir.Operation) ?ShapeKernel {
615     if (isName(op.name.name, dialect_mod.AccyDialect.BroadcastInDimOp.operation_name)) {
616         return .broadcast_in_dim;
617     }
618     if (std.mem.eql(u8, op.name.name, dialect_mod.AccyDialect.IotaOp.operation_name)) return .iota;
619     if (std.mem.eql(u8, op.name.name, dialect_mod.AccyDialect.ReshapeOp.operation_name)) return .reshape;
620     if (std.mem.eql(u8, op.name.name, dialect_mod.AccyDialect.TransposeOp.operation_name)) return .transpose;
621     if (std.mem.eql(u8, op.name.name, dialect_mod.AccyDialect.SliceOp.operation_name)) return .slice;
622     if (std.mem.eql(u8, op.name.name, dialect_mod.AccyDialect.PadOp.operation_name)) return .pad;
623     if (std.mem.eql(u8, op.name.name, dialect_mod.AccyDialect.ConcatenateOp.operation_name)) return .concatenate;
624     if (std.mem.eql(u8, op.name.name, dialect_mod.AccyDialect.GatherOp.operation_name)) return .gather;
625     if (std.mem.eql(u8, op.name.name, dialect_mod.AccyDialect.ScatterOp.operation_name)) return .scatter;
626     return null;
627 }