Skip to documentation
SLOP

tiny.choir.passes.memref_views

Reference tiny.choir passes memref_views

Defined in passes.

API (2)

Actions

Public operations.

Types and contracts

Public types and contracts.

No direct callersNo direct callspassesmemref views
Static calls · unresolved targets: unknown · external targets: unknown.

Source

Called byCallsNo direct callersprivate sourcelib.choir.src.passes.views.MemrefViewLegalizerdeinitprivate sourcelib.choir.src.passes.views.MemrefViewLegalizerinitprivate sourcelib.choir.src.passes.views.MemrefViewLegalizerrunpasses.memref_viewslegalizeMemrefViews
Static calls · unresolved targets: 0 · external targets: 0.

Source: lib/choir/src/passes/root.zig:181

zig
pub const memref_views = @import("views.zig");

Source: lib/choir/src/passes/views.zig

zig
const std = @import("std");const ir = @import("../core/root.zig");const dialects = @import("../dialects/root.zig");const ArithDialect = dialects.arith.ArithDialect;const MemrefDialect = dialects.memref.MemrefDialect;const ScfDialect = dialects.scf.ScfDialect;pub const LegalizeError = anyerror;const ViewInfo = struct {    source: *ir.Value,    offset: u64,    shape: []u64,    stride: []u64,};const Insertion = struct {    block: *ir.Block,    before: ?*ir.Operation,    fn insert(self: Insertion, op: *ir.Operation) !void {        if (self.before) |before| {            try self.block.insertBefore(op, before);        } else {            try self.block.addOperation(op);        }    }};const ResolvedIndex = struct {    base: *ir.Value,    index: *ir.Value,};pub fn legalizeMemrefViews(module: *ir.Operation, allocator: std.mem.Allocator) LegalizeError!void {    const ctx = module.context;    var legalizer = MemrefViewLegalizer.init(allocator, ctx);    defer legalizer.deinit();    try legalizer.run(module);}const MemrefViewLegalizer = struct {    allocator: std.mem.Allocator,    ctx: *ir.Context,    view_map: std.AutoHashMap(*ir.Value, ViewInfo),    const HandleResult = enum { kept, removed };    fn init(allocator: std.mem.Allocator, ctx: *ir.Context) MemrefViewLegalizer {        return .{            .allocator = allocator,            .ctx = ctx,            .view_map = std.AutoHashMap(*ir.Value, ViewInfo).init(allocator),        };    }    fn deinit(self: *MemrefViewLegalizer) void {        var iter = self.view_map.valueIterator();        while (iter.next()) |info| {            self.allocator.free(info.shape);            self.allocator.free(info.stride);        }        self.view_map.deinit();    }    fn run(self: *MemrefViewLegalizer, module: *ir.Operation) LegalizeError!void {        try self.walkOp(module);        try self.cleanupViews();    }    fn walkOp(self: *MemrefViewLegalizer, op: *ir.Operation) LegalizeError!void {        const handled = try self.handleOp(op);        if (handled == .removed) return;        for (op.regions.items) |*region| {            var block_iter = region.getBlocks();            while (block_iter.next()) |block| {                var current: ?*ir.Operation = @ptrCast(@alignCast(block.operations.head));                while (current) |op_ptr| {                    const next = op_ptr.next_op;                    try self.walkOp(op_ptr);                    current = next;                }            }        }    }    fn handleOp(self: *MemrefViewLegalizer, op: *ir.Operation) LegalizeError!HandleResult {        const name = op.name.name;        if (std.mem.eql(u8, name, MemrefDialect.LoadOp.operation_name)) {            try self.rewriteLoad(op);            return .kept;        }        if (std.mem.eql(u8, name, MemrefDialect.StoreOp.operation_name)) {            try self.rewriteStore(op);            return .kept;        }        if (std.mem.eql(u8, name, MemrefDialect.CopyOp.operation_name)) {            if (try self.rewriteCopy(op)) {                return .removed;            }            return .kept;        }        if (std.mem.eql(u8, name, MemrefDialect.DeallocOp.operation_name)) {            try self.rewriteDealloc(op);            return .kept;        }        if (std.mem.eql(u8, name, MemrefDialect.SubviewOp.operation_name) or            std.mem.eql(u8, name, MemrefDialect.TransposeOp.operation_name))        {            _ = try self.getViewInfo(op.getResult(0).?);            return .kept;        }        if (try self.opUsesView(op)) {            return error.UnsupportedViewUse;        }        return .kept;    }    fn opUsesView(self: *MemrefViewLegalizer, op: *ir.Operation) LegalizeError!bool {        for (op.operands.items) |operand| {            if (try self.getViewInfo(operand.value) != null) {                return true;            }        }        return false;    }    fn rewriteLoad(self: *MemrefViewLegalizer, op: *ir.Operation) LegalizeError!void {        const load = MemrefDialect.LoadOp{ .op = op };        const memref_val = load.getMemref();        const index_val = load.getIndex();        if (try self.getViewInfo(memref_val) == null) return;        const block = op.getBlock() orelse return error.MissingBlock;        const insertion = Insertion{ .block = block, .before = op };        const resolved = try self.resolveViewIndex(insertion, op.getLoc(), memref_val, index_val);        op.setOperandValue(0, resolved.base);        op.setOperandValue(1, resolved.index);    }    fn rewriteStore(self: *MemrefViewLegalizer, op: *ir.Operation) LegalizeError!void {        const store = MemrefDialect.StoreOp{ .op = op };        const memref_val = store.getMemref();        const index_val = store.getIndex();        if (try self.getViewInfo(memref_val) == null) return;        const block = op.getBlock() orelse return error.MissingBlock;        const insertion = Insertion{ .block = block, .before = op };        const resolved = try self.resolveViewIndex(insertion, op.getLoc(), memref_val, index_val);        op.setOperandValue(1, resolved.base);        op.setOperandValue(2, resolved.index);    }    fn rewriteDealloc(self: *MemrefViewLegalizer, op: *ir.Operation) LegalizeError!void {        const dealloc = MemrefDialect.DeallocOp{ .op = op };        const memref_val = dealloc.getMemref();        if (try self.getViewInfo(memref_val) == null) return;        const base = try self.resolveBaseMemref(memref_val);        op.setOperandValue(0, base);    }    fn rewriteCopy(self: *MemrefViewLegalizer, op: *ir.Operation) LegalizeError!bool {        const copy = MemrefDialect.CopyOp{ .op = op };        const src = copy.getSrc();        const dst = copy.getDst();        const src_view = try self.getViewInfo(src);        const dst_view = try self.getViewInfo(dst);        if (src_view == null and dst_view == null) return false;        const block = op.getBlock() orelse return error.MissingBlock;        const loc = op.getLoc();        const len = try self.getCopyLength(op, src, dst);        const index_type = ArithDialect.getIndexType(self.ctx) catch return error.InvalidLayout;        const insertion = Insertion{ .block = block, .before = op };        const zero = try self.emitConstIndex(insertion, loc, index_type, 0);        const len_val = try self.emitConstIndex(insertion, loc, index_type, len);        const one = try self.emitConstIndex(insertion, loc, index_type, 1);        var for_op = try ScfDialect.ForOp.create(self.ctx, loc, zero, len_val, one, &.{}, &.{});        try block.insertBefore(for_op.op, op);        const body_block = for_op.getBodyBlock();        const iv = for_op.getInductionVar();        var body_insert = Insertion{ .block = body_block, .before = null };        const src_resolved = try self.resolveViewIndex(body_insert, loc, src, iv);        const elem_type = try self.memrefElementType(src_resolved.base);        const load_op = try MemrefDialect.LoadOp.create(self.ctx, loc, src_resolved.base, src_resolved.index, elem_type);        try body_insert.insert(load_op.op);        const dst_resolved = try self.resolveViewIndex(body_insert, loc, dst, iv);        const store_op = try MemrefDialect.StoreOp.create(self.ctx, loc, load_op.getResult(), dst_resolved.base, dst_resolved.index);        try body_insert.insert(store_op.op);        const yield_op = try ScfDialect.YieldOp.create(self.ctx, loc, &.{});        try body_insert.insert(yield_op.op);        block.removeOperation(op);        return true;    }    fn resolveBaseMemref(self: *MemrefViewLegalizer, memref: *ir.Value) LegalizeError!*ir.Value {        var current = memref;        while (true) {            const info = try self.getViewInfo(current) orelse return current;            current = info.source;        }    }    fn resolveViewIndex(        self: *MemrefViewLegalizer,        insertion: Insertion,        loc: ir.Location,        memref: *ir.Value,        index: *ir.Value,    ) LegalizeError!ResolvedIndex {        const info = try self.getViewInfo(memref) orelse return .{ .base = memref, .index = index };        const parent_index = try self.emitLinearToOffset(insertion, loc, index, info.shape, info.stride, info.offset);        return self.resolveViewIndex(insertion, loc, info.source, parent_index);    }    fn getViewInfo(self: *MemrefViewLegalizer, value: *ir.Value) LegalizeError!?*ViewInfo {        if (self.view_map.getPtr(value)) |info| return info;        const def_op_opaque = value.getDefiningOp() orelse return null;        const def_op: *ir.Operation = @ptrCast(@alignCast(def_op_opaque));        const is_subview = std.mem.eql(u8, def_op.name.name, MemrefDialect.SubviewOp.operation_name);        const is_transpose = std.mem.eql(u8, def_op.name.name, MemrefDialect.TransposeOp.operation_name);        if (!is_subview and !is_transpose) return null;        if (is_subview) {            const subview = MemrefDialect.SubviewOp{ .op = def_op };            const shape_payload = subview.getShapePayload() orelse return error.MissingLayout;            const stride_payload = subview.getStridePayload() orelse return error.MissingLayout;            const offset_payload = subview.getOffsetPayload() orelse return error.MissingLayout;            const shape = try parseDims(self.allocator, shape_payload);            errdefer self.allocator.free(shape);            const stride = try parseDims(self.allocator, stride_payload);            errdefer self.allocator.free(stride);            if (shape.len != stride.len) return error.InvalidLayout;            const offset = try parseOffset(offset_payload);            const info = ViewInfo{                .source = subview.getSource(),                .offset = offset,                .shape = shape,                .stride = stride,            };            try self.view_map.put(value, info);            return self.view_map.getPtr(value).?;        }        const transpose = MemrefDialect.TransposeOp{ .op = def_op };        const shape_payload = transpose.getShapePayload() orelse return error.MissingLayout;        const stride_payload = transpose.getStridePayload() orelse return error.MissingLayout;        const shape = try parseDims(self.allocator, shape_payload);        errdefer self.allocator.free(shape);        const stride = try parseDims(self.allocator, stride_payload);        errdefer self.allocator.free(stride);        if (shape.len != stride.len) return error.InvalidLayout;        const info = ViewInfo{            .source = transpose.getSource(),            .offset = 0,            .shape = shape,            .stride = stride,        };        try self.view_map.put(value, info);        return self.view_map.getPtr(value).?;    }    fn parseOffset(payload: []const u8) LegalizeError!u64 {        if (payload.len == 0) return error.InvalidLayout;        return std.fmt.parseInt(u64, payload, 10) catch return error.InvalidLayout;    }    fn parseDims(allocator: std.mem.Allocator, payload: []const u8) LegalizeError![]u64 {        if (payload.len == 0) return allocator.alloc(u64, 0);        var dims: std.ArrayListUnmanaged(u64) = .empty;        errdefer dims.deinit(allocator);        var it = std.mem.splitScalar(u8, payload, ',');        while (it.next()) |part| {            if (part.len == 0) return error.InvalidLayout;            const dim = std.fmt.parseInt(u64, part, 10) catch return error.InvalidLayout;            try dims.append(allocator, dim);        }        return dims.toOwnedSlice(allocator);    }    fn emitConstIndex(        self: *MemrefViewLegalizer,        insertion: Insertion,        loc: ir.Location,        index_type: ir.Type,        value: u64,    ) LegalizeError!*ir.Value {        const int_value = std.math.cast(i64, value) orelse return error.LayoutOverflow;        var const_op = ArithDialect.ConstantOp.createInt(self.ctx, loc, index_type, int_value) catch return error.InvalidLayout;        try insertion.insert(const_op.op);        return const_op.getResult();    }    fn emitLinearToOffset(        self: *MemrefViewLegalizer,        insertion: Insertion,        loc: ir.Location,        index: *ir.Value,        shape: []const u64,        stride: []const u64,        offset: u64,    ) LegalizeError!*ir.Value {        if (shape.len != stride.len) return error.InvalidLayout;        const index_type = index.type;        if (shape.len == 0) {            if (offset == 0) return index;            const base = try self.emitConstIndex(insertion, loc, index_type, offset);            var add_op = ArithDialect.AddOp.create(self.ctx, loc, base, index) catch return error.InvalidLayout;            try insertion.insert(add_op.op);            return add_op.getResult();        }        var accum = try self.emitConstIndex(insertion, loc, index_type, offset);        var current = index;        var i: usize = 0;        while (i < shape.len) : (i += 1) {            const stride_row = try product(shape[i + 1 ..]);            if (stride_row == 0) return error.InvalidLayout;            var idx_val: *ir.Value = undefined;            if (stride_row == 1) {                idx_val = current;            } else {                const stride_const = try self.emitConstIndex(insertion, loc, index_type, stride_row);                var div_op = ArithDialect.DivOp.create(self.ctx, loc, current, stride_const) catch return error.InvalidLayout;                try insertion.insert(div_op.op);                idx_val = div_op.getResult();                var rem_op = ArithDialect.RemOp.create(self.ctx, loc, current, stride_const) catch return error.InvalidLayout;                try insertion.insert(rem_op.op);                current = rem_op.getResult();            }            const stride_val = stride[i];            if (stride_val != 0) {                var term = idx_val;                if (stride_val != 1) {                    const stride_const2 = try self.emitConstIndex(insertion, loc, index_type, stride_val);                    var mul_op = ArithDialect.MulOp.create(self.ctx, loc, idx_val, stride_const2) catch return error.InvalidLayout;                    try insertion.insert(mul_op.op);                    term = mul_op.getResult();                }                var add_op = ArithDialect.AddOp.create(self.ctx, loc, accum, term) catch return error.InvalidLayout;                try insertion.insert(add_op.op);                accum = add_op.getResult();            }        }        return accum;    }    fn product(dims: []const u64) LegalizeError!u64 {        var total: u64 = 1;        for (dims) |dim| {            total = std.math.mul(u64, total, dim) catch return error.LayoutOverflow;        }        return total;    }    fn getCopyLength(self: *MemrefViewLegalizer, op: *ir.Operation, src: *ir.Value, dst: *ir.Value) LegalizeError!u64 {        if (op.getAttrAs(ir.Attribute.IntegerAttr, "len")) |len_attr| {            const len_int = len_attr.getValue();            if (len_int < 0) return error.InvalidLayout;            return std.math.cast(u64, len_int) orelse return error.LayoutOverflow;        }        const src_size = try self.memrefSize(src);        const dst_size = try self.memrefSize(dst);        if (src_size) |s| {            if (dst_size) |d| {                if (s != d) return error.InvalidLayout;            }            return s;        }        if (dst_size) |d| return d;        return error.UnsupportedCopyLength;    }    fn memrefSize(self: *MemrefViewLegalizer, memref: *ir.Value) LegalizeError!?u64 {        _ = self;        const params = memrefParams(memref.type) orelse return error.InvalidLayout;        return params.size;    }    fn memrefElementType(self: *MemrefViewLegalizer, memref: *ir.Value) LegalizeError!ir.Type {        const params = memrefParams(memref.type) orelse return error.InvalidLayout;        return self.ctx.getDialectTypeFromName(params.element_type_name) catch return error.InvalidLayout;    }    fn memrefParams(memref_type: ir.Type) ?MemrefDialect.MemrefParams {        const param_key = memref_type.getDialectParamKey() orelse return null;        return MemrefDialect.parseMemrefParams(param_key);    }    fn cleanupViews(self: *MemrefViewLegalizer) LegalizeError!void {        var iter = self.view_map.iterator();        while (iter.next()) |entry| {            const value = entry.key_ptr.*;            if (!value.hasNoUses()) continue;            const def_op_opaque = value.getDefiningOp() orelse continue;            const def_op: *ir.Operation = @ptrCast(@alignCast(def_op_opaque));            if (!std.mem.eql(u8, def_op.name.name, MemrefDialect.SubviewOp.operation_name) and                !std.mem.eql(u8, def_op.name.name, MemrefDialect.TransposeOp.operation_name))            {                continue;            }            if (def_op.getBlock()) |block| {                block.removeOperation(def_op);            }        }    }};

Audit

Definitions3
Public names3
Members0
Version26.7.0
Revisiondaab053ee433