lib/choir/src/passes/views.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const ir = @import("../core/root.zig");
3 const dialects = @import("../dialects/root.zig");
4
5 const ArithDialect = dialects.arith.ArithDialect;
6 const MemrefDialect = dialects.memref.MemrefDialect;
7 const ScfDialect = dialects.scf.ScfDialect;
8
9 pub const LegalizeError = anyerror;
10
11 const ViewInfo = struct {
12 source: *ir.Value,
13 offset: u64,
14 shape: []u64,
15 stride: []u64,
16 };
17
18 const Insertion = struct {
19 block: *ir.Block,
20 before: ?*ir.Operation,
21
22 fn insert(self: Insertion, op: *ir.Operation) !void {
23 if (self.before) |before| {
24 try self.block.insertBefore(op, before);
25 } else {
26 try self.block.addOperation(op);
27 }
28 }
29 };
30
31 const ResolvedIndex = struct {
32 base: *ir.Value,
33 index: *ir.Value,
34 };
35
36 pub fn legalizeMemrefViews(module: *ir.Operation, allocator: std.mem.Allocator) LegalizeError!void {
37 const ctx = module.context;
38 var legalizer = MemrefViewLegalizer.init(allocator, ctx);
39 defer legalizer.deinit();
40 try legalizer.run(module);
41 }
42
43 const MemrefViewLegalizer = struct {
44 allocator: std.mem.Allocator,
45 ctx: *ir.Context,
46 view_map: std.AutoHashMap(*ir.Value, ViewInfo),
47
48 const HandleResult = enum { kept, removed };
49
50 fn init(allocator: std.mem.Allocator, ctx: *ir.Context) MemrefViewLegalizer {
51 return .{
52 .allocator = allocator,
53 .ctx = ctx,
54 .view_map = std.AutoHashMap(*ir.Value, ViewInfo).init(allocator),
55 };
56 }
57
58 fn deinit(self: *MemrefViewLegalizer) void {
59 var iter = self.view_map.valueIterator();
60 while (iter.next()) |info| {
61 self.allocator.free(info.shape);
62 self.allocator.free(info.stride);
63 }
64 self.view_map.deinit();
65 }
66
67 fn run(self: *MemrefViewLegalizer, module: *ir.Operation) LegalizeError!void {
68 try self.walkOp(module);
69 try self.cleanupViews();
70 }
71
72 fn walkOp(self: *MemrefViewLegalizer, op: *ir.Operation) LegalizeError!void {
73 const handled = try self.handleOp(op);
74 if (handled == .removed) return;
75
76 for (op.regions.items) |*region| {
77 var block_iter = region.getBlocks();
78 while (block_iter.next()) |block| {
79 var current: ?*ir.Operation = @ptrCast(@alignCast(block.operations.head));
80 while (current) |op_ptr| {
81 const next = op_ptr.next_op;
82 try self.walkOp(op_ptr);
83 current = next;
84 }
85 }
86 }
87 }
88
89 fn handleOp(self: *MemrefViewLegalizer, op: *ir.Operation) LegalizeError!HandleResult {
90 const name = op.name.name;
91 if (std.mem.eql(u8, name, MemrefDialect.LoadOp.operation_name)) {
92 try self.rewriteLoad(op);
93 return .kept;
94 }
95 if (std.mem.eql(u8, name, MemrefDialect.StoreOp.operation_name)) {
96 try self.rewriteStore(op);
97 return .kept;
98 }
99 if (std.mem.eql(u8, name, MemrefDialect.CopyOp.operation_name)) {
100 if (try self.rewriteCopy(op)) {
101 return .removed;
102 }
103 return .kept;
104 }
105 if (std.mem.eql(u8, name, MemrefDialect.DeallocOp.operation_name)) {
106 try self.rewriteDealloc(op);
107 return .kept;
108 }
109 if (std.mem.eql(u8, name, MemrefDialect.SubviewOp.operation_name) or
110 std.mem.eql(u8, name, MemrefDialect.TransposeOp.operation_name))
111 {
112 _ = try self.getViewInfo(op.getResult(0).?);
113 return .kept;
114 }
115
116 if (try self.opUsesView(op)) {
117 return error.UnsupportedViewUse;
118 }
119
120 return .kept;
121 }
122
123 fn opUsesView(self: *MemrefViewLegalizer, op: *ir.Operation) LegalizeError!bool {
124 for (op.operands.items) |operand| {
125 if (try self.getViewInfo(operand.value) != null) {
126 return true;
127 }
128 }
129 return false;
130 }
131
132 fn rewriteLoad(self: *MemrefViewLegalizer, op: *ir.Operation) LegalizeError!void {
133 const load = MemrefDialect.LoadOp{ .op = op };
134 const memref_val = load.getMemref();
135 const index_val = load.getIndex();
136 if (try self.getViewInfo(memref_val) == null) return;
137
138 const block = op.getBlock() orelse return error.MissingBlock;
139 const insertion = Insertion{ .block = block, .before = op };
140 const resolved = try self.resolveViewIndex(insertion, op.getLoc(), memref_val, index_val);
141
142 op.setOperandValue(0, resolved.base);
143 op.setOperandValue(1, resolved.index);
144 }
145
146 fn rewriteStore(self: *MemrefViewLegalizer, op: *ir.Operation) LegalizeError!void {
147 const store = MemrefDialect.StoreOp{ .op = op };
148 const memref_val = store.getMemref();
149 const index_val = store.getIndex();
150 if (try self.getViewInfo(memref_val) == null) return;
151
152 const block = op.getBlock() orelse return error.MissingBlock;
153 const insertion = Insertion{ .block = block, .before = op };
154 const resolved = try self.resolveViewIndex(insertion, op.getLoc(), memref_val, index_val);
155
156 op.setOperandValue(1, resolved.base);
157 op.setOperandValue(2, resolved.index);
158 }
159
160 fn rewriteDealloc(self: *MemrefViewLegalizer, op: *ir.Operation) LegalizeError!void {
161 const dealloc = MemrefDialect.DeallocOp{ .op = op };
162 const memref_val = dealloc.getMemref();
163 if (try self.getViewInfo(memref_val) == null) return;
164 const base = try self.resolveBaseMemref(memref_val);
165 op.setOperandValue(0, base);
166 }
167
168 fn rewriteCopy(self: *MemrefViewLegalizer, op: *ir.Operation) LegalizeError!bool {
169 const copy = MemrefDialect.CopyOp{ .op = op };
170 const src = copy.getSrc();
171 const dst = copy.getDst();
172 const src_view = try self.getViewInfo(src);
173 const dst_view = try self.getViewInfo(dst);
174 if (src_view == null and dst_view == null) return false;
175
176 const block = op.getBlock() orelse return error.MissingBlock;
177 const loc = op.getLoc();
178 const len = try self.getCopyLength(op, src, dst);
179
180 const index_type = ArithDialect.getIndexType(self.ctx) catch return error.InvalidLayout;
181
182 const insertion = Insertion{ .block = block, .before = op };
183 const zero = try self.emitConstIndex(insertion, loc, index_type, 0);
184 const len_val = try self.emitConstIndex(insertion, loc, index_type, len);
185 const one = try self.emitConstIndex(insertion, loc, index_type, 1);
186
187 var for_op = try ScfDialect.ForOp.create(self.ctx, loc, zero, len_val, one, &.{}, &.{});
188 try block.insertBefore(for_op.op, op);
189
190 const body_block = for_op.getBodyBlock();
191 const iv = for_op.getInductionVar();
192 var body_insert = Insertion{ .block = body_block, .before = null };
193
194 const src_resolved = try self.resolveViewIndex(body_insert, loc, src, iv);
195 const elem_type = try self.memrefElementType(src_resolved.base);
196 const load_op = try MemrefDialect.LoadOp.create(self.ctx, loc, src_resolved.base, src_resolved.index, elem_type);
197 try body_insert.insert(load_op.op);
198
199 const dst_resolved = try self.resolveViewIndex(body_insert, loc, dst, iv);
200 const store_op = try MemrefDialect.StoreOp.create(self.ctx, loc, load_op.getResult(), dst_resolved.base, dst_resolved.index);
201 try body_insert.insert(store_op.op);
202
203 const yield_op = try ScfDialect.YieldOp.create(self.ctx, loc, &.{});
204 try body_insert.insert(yield_op.op);
205
206 block.removeOperation(op);
207 return true;
208 }
209
210 fn resolveBaseMemref(self: *MemrefViewLegalizer, memref: *ir.Value) LegalizeError!*ir.Value {
211 var current = memref;
212 while (true) {
213 const info = try self.getViewInfo(current) orelse return current;
214 current = info.source;
215 }
216 }
217
218 fn resolveViewIndex(
219 self: *MemrefViewLegalizer,
220 insertion: Insertion,
221 loc: ir.Location,
222 memref: *ir.Value,
223 index: *ir.Value,
224 ) LegalizeError!ResolvedIndex {
225 const info = try self.getViewInfo(memref) orelse return .{ .base = memref, .index = index };
226 const parent_index = try self.emitLinearToOffset(insertion, loc, index, info.shape, info.stride, info.offset);
227 return self.resolveViewIndex(insertion, loc, info.source, parent_index);
228 }
229
230 fn getViewInfo(self: *MemrefViewLegalizer, value: *ir.Value) LegalizeError!?*ViewInfo {
231 if (self.view_map.getPtr(value)) |info| return info;
232
233 const def_op_opaque = value.getDefiningOp() orelse return null;
234 const def_op: *ir.Operation = @ptrCast(@alignCast(def_op_opaque));
235
236 const is_subview = std.mem.eql(u8, def_op.name.name, MemrefDialect.SubviewOp.operation_name);
237 const is_transpose = std.mem.eql(u8, def_op.name.name, MemrefDialect.TransposeOp.operation_name);
238 if (!is_subview and !is_transpose) return null;
239
240 if (is_subview) {
241 const subview = MemrefDialect.SubviewOp{ .op = def_op };
242 const shape_payload = subview.getShapePayload() orelse return error.MissingLayout;
243 const stride_payload = subview.getStridePayload() orelse return error.MissingLayout;
244 const offset_payload = subview.getOffsetPayload() orelse return error.MissingLayout;
245
246 const shape = try parseDims(self.allocator, shape_payload);
247 errdefer self.allocator.free(shape);
248 const stride = try parseDims(self.allocator, stride_payload);
249 errdefer self.allocator.free(stride);
250 if (shape.len != stride.len) return error.InvalidLayout;
251
252 const offset = try parseOffset(offset_payload);
253
254 const info = ViewInfo{
255 .source = subview.getSource(),
256 .offset = offset,
257 .shape = shape,
258 .stride = stride,
259 };
260 try self.view_map.put(value, info);
261 return self.view_map.getPtr(value).?;
262 }
263
264 const transpose = MemrefDialect.TransposeOp{ .op = def_op };
265 const shape_payload = transpose.getShapePayload() orelse return error.MissingLayout;
266 const stride_payload = transpose.getStridePayload() orelse return error.MissingLayout;
267
268 const shape = try parseDims(self.allocator, shape_payload);
269 errdefer self.allocator.free(shape);
270 const stride = try parseDims(self.allocator, stride_payload);
271 errdefer self.allocator.free(stride);
272 if (shape.len != stride.len) return error.InvalidLayout;
273
274 const info = ViewInfo{
275 .source = transpose.getSource(),
276 .offset = 0,
277 .shape = shape,
278 .stride = stride,
279 };
280 try self.view_map.put(value, info);
281 return self.view_map.getPtr(value).?;
282 }
283
284 fn parseOffset(payload: []const u8) LegalizeError!u64 {
285 if (payload.len == 0) return error.InvalidLayout;
286 return std.fmt.parseInt(u64, payload, 10) catch return error.InvalidLayout;
287 }
288
289 fn parseDims(allocator: std.mem.Allocator, payload: []const u8) LegalizeError![]u64 {
290 if (payload.len == 0) return allocator.alloc(u64, 0);
291 var dims: std.ArrayListUnmanaged(u64) = .empty;
292 errdefer dims.deinit(allocator);
293
294 var it = std.mem.splitScalar(u8, payload, ',');
295 while (it.next()) |part| {
296 if (part.len == 0) return error.InvalidLayout;
297 const dim = std.fmt.parseInt(u64, part, 10) catch return error.InvalidLayout;
298 try dims.append(allocator, dim);
299 }
300
301 return dims.toOwnedSlice(allocator);
302 }
303
304 fn emitConstIndex(
305 self: *MemrefViewLegalizer,
306 insertion: Insertion,
307 loc: ir.Location,
308 index_type: ir.Type,
309 value: u64,
310 ) LegalizeError!*ir.Value {
311 const int_value = std.math.cast(i64, value) orelse return error.LayoutOverflow;
312 var const_op = ArithDialect.ConstantOp.createInt(self.ctx, loc, index_type, int_value) catch return error.InvalidLayout;
313 try insertion.insert(const_op.op);
314 return const_op.getResult();
315 }
316
317 fn emitLinearToOffset(
318 self: *MemrefViewLegalizer,
319 insertion: Insertion,
320 loc: ir.Location,
321 index: *ir.Value,
322 shape: []const u64,
323 stride: []const u64,
324 offset: u64,
325 ) LegalizeError!*ir.Value {
326 if (shape.len != stride.len) return error.InvalidLayout;
327 const index_type = index.type;
328
329 if (shape.len == 0) {
330 if (offset == 0) return index;
331 const base = try self.emitConstIndex(insertion, loc, index_type, offset);
332 var add_op = ArithDialect.AddOp.create(self.ctx, loc, base, index) catch return error.InvalidLayout;
333 try insertion.insert(add_op.op);
334 return add_op.getResult();
335 }
336
337 var accum = try self.emitConstIndex(insertion, loc, index_type, offset);
338 var current = index;
339
340 var i: usize = 0;
341 while (i < shape.len) : (i += 1) {
342 const stride_row = try product(shape[i + 1 ..]);
343 if (stride_row == 0) return error.InvalidLayout;
344
345 var idx_val: *ir.Value = undefined;
346 if (stride_row == 1) {
347 idx_val = current;
348 } else {
349 const stride_const = try self.emitConstIndex(insertion, loc, index_type, stride_row);
350 var div_op = ArithDialect.DivOp.create(self.ctx, loc, current, stride_const) catch return error.InvalidLayout;
351 try insertion.insert(div_op.op);
352 idx_val = div_op.getResult();
353
354 var rem_op = ArithDialect.RemOp.create(self.ctx, loc, current, stride_const) catch return error.InvalidLayout;
355 try insertion.insert(rem_op.op);
356 current = rem_op.getResult();
357 }
358
359 const stride_val = stride[i];
360 if (stride_val != 0) {
361 var term = idx_val;
362 if (stride_val != 1) {
363 const stride_const2 = try self.emitConstIndex(insertion, loc, index_type, stride_val);
364 var mul_op = ArithDialect.MulOp.create(self.ctx, loc, idx_val, stride_const2) catch return error.InvalidLayout;
365 try insertion.insert(mul_op.op);
366 term = mul_op.getResult();
367 }
368
369 var add_op = ArithDialect.AddOp.create(self.ctx, loc, accum, term) catch return error.InvalidLayout;
370 try insertion.insert(add_op.op);
371 accum = add_op.getResult();
372 }
373 }
374
375 return accum;
376 }
377
378 fn product(dims: []const u64) LegalizeError!u64 {
379 var total: u64 = 1;
380 for (dims) |dim| {
381 total = std.math.mul(u64, total, dim) catch return error.LayoutOverflow;
382 }
383 return total;
384 }
385
386 fn getCopyLength(self: *MemrefViewLegalizer, op: *ir.Operation, src: *ir.Value, dst: *ir.Value) LegalizeError!u64 {
387 if (op.getAttrAs(ir.Attribute.IntegerAttr, "len")) |len_attr| {
388 const len_int = len_attr.getValue();
389 if (len_int < 0) return error.InvalidLayout;
390 return std.math.cast(u64, len_int) orelse return error.LayoutOverflow;
391 }
392
393 const src_size = try self.memrefSize(src);
394 const dst_size = try self.memrefSize(dst);
395
396 if (src_size) |s| {
397 if (dst_size) |d| {
398 if (s != d) return error.InvalidLayout;
399 }
400 return s;
401 }
402 if (dst_size) |d| return d;
403
404 return error.UnsupportedCopyLength;
405 }
406
407 fn memrefSize(self: *MemrefViewLegalizer, memref: *ir.Value) LegalizeError!?u64 {
408 _ = self;
409 const params = memrefParams(memref.type) orelse return error.InvalidLayout;
410 return params.size;
411 }
412
413 fn memrefElementType(self: *MemrefViewLegalizer, memref: *ir.Value) LegalizeError!ir.Type {
414 const params = memrefParams(memref.type) orelse return error.InvalidLayout;
415 return self.ctx.getDialectTypeFromName(params.element_type_name) catch return error.InvalidLayout;
416 }
417
418 fn memrefParams(memref_type: ir.Type) ?MemrefDialect.MemrefParams {
419 const param_key = memref_type.getDialectParamKey() orelse return null;
420 return MemrefDialect.parseMemrefParams(param_key);
421 }
422
423 fn cleanupViews(self: *MemrefViewLegalizer) LegalizeError!void {
424 var iter = self.view_map.iterator();
425 while (iter.next()) |entry| {
426 const value = entry.key_ptr.*;
427 if (!value.hasNoUses()) continue;
428 const def_op_opaque = value.getDefiningOp() orelse continue;
429 const def_op: *ir.Operation = @ptrCast(@alignCast(def_op_opaque));
430 if (!std.mem.eql(u8, def_op.name.name, MemrefDialect.SubviewOp.operation_name) and
431 !std.mem.eql(u8, def_op.name.name, MemrefDialect.TransposeOp.operation_name))
432 {
433 continue;
434 }
435 if (def_op.getBlock()) |block| {
436 block.removeOperation(def_op);
437 }
438 }
439 }
440 };