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 }