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 };