lib/choir/src/backends/wasm/emission/function/instruction/memory.zig

daab053ee43316e1809a84551d573ddd1e5bf3d2

  1 const std = @import("std");
  2 const wasm = @import("../../../root.zig");
  3 const ir = @import("../../../../../core/root.zig");
  4 const dialects = @import("../../../../../dialects/root.zig");
  5 const emission = @import("../../root.zig");
  6 const instruction = @import("root.zig");
  7 
  8 const MemrefDialect = dialects.MemrefDialect;
  9 
 10 pub fn writeLoad(
 11     comptime Out: type,
 12     writer: *instruction.Writer(Out),
 13     operation: *ir.Operation,
 14 ) emission.Error!void {
 15     const load = MemrefDialect.LoadOp{ .op = operation };
 16     const layout = try memrefParams(load.getMemref().type);
 17     try writeAddress(Out, writer, load.getMemref(), load.getIndex(), layout);
 18     try writer.out.writeByte(try loadOpcode(layout.element_type_name));
 19     try wasm.binary.writeUleb(
 20         writer.out,
 21         alignmentPower(try elementByteSize(layout.element_type_name)),
 22     );
 23     try wasm.binary.writeUleb(writer.out, @as(u32, 0));
 24     try writer.writeLocalSet(try writer.localFor(load.getResult()));
 25 }
 26 
 27 pub fn writeStore(
 28     comptime Out: type,
 29     writer: *instruction.Writer(Out),
 30     operation: *ir.Operation,
 31 ) emission.Error!void {
 32     const store = MemrefDialect.StoreOp{ .op = operation };
 33     const layout = try memrefParams(store.getMemref().type);
 34     try writeAddress(Out, writer, store.getMemref(), store.getIndex(), layout);
 35     try writer.writeValue(store.getValue());
 36     try writer.out.writeByte(try storeOpcode(layout.element_type_name));
 37     try wasm.binary.writeUleb(
 38         writer.out,
 39         alignmentPower(try elementByteSize(layout.element_type_name)),
 40     );
 41     try wasm.binary.writeUleb(writer.out, @as(u32, 0));
 42 }
 43 
 44 fn writeAddress(
 45     comptime Out: type,
 46     writer: *instruction.Writer(Out),
 47     memref: *ir.Value,
 48     index: *ir.Value,
 49     layout: MemrefDialect.MemrefParams,
 50 ) emission.Error!void {
 51     try requireMemoryAddressSpace(layout);
 52     const element_size = try elementByteSize(layout.element_type_name);
 53     try writer.writeValue(memref);
 54     try writer.writeValue(index);
 55     if (element_size != 1) {
 56         try writer.out.writeByte(0x41);
 57         try wasm.binary.writeSleb(writer.out, @intCast(element_size));
 58         try writer.out.writeByte(0x6c);
 59     }
 60     try writer.out.writeByte(0x6a);
 61 }
 62 
 63 fn memrefParams(memref_type: ir.Type) error{CodeGenFailed}!MemrefDialect.MemrefParams {
 64     const key = memref_type.getDialectParamKey() orelse return error.CodeGenFailed;
 65     return MemrefDialect.parseMemrefParams(key) orelse return error.CodeGenFailed;
 66 }
 67 
 68 fn requireMemoryAddressSpace(layout: MemrefDialect.MemrefParams) error{CodeGenFailed}!void {
 69     if (layout.addr_space != .host and layout.addr_space != .unified) {
 70         return error.CodeGenFailed;
 71     }
 72 }
 73 
 74 fn elementByteSize(element_type_name: []const u8) error{CodeGenFailed}!u32 {
 75     const kind = dialects.arith.scalarKindFromTypeName(element_type_name) orelse
 76         return error.CodeGenFailed;
 77     return switch (kind) {
 78         .i8, .u8 => 1,
 79         .i16, .u16, .f16, .bf16 => 2,
 80         .i32, .u32, .f32, .index, .bool => 4,
 81         .i64, .u64, .f64 => 8,
 82     };
 83 }
 84 
 85 fn alignmentPower(size: u32) u32 {
 86     return switch (size) {
 87         1 => 0,
 88         2 => 1,
 89         4 => 2,
 90         8 => 3,
 91         else => 0,
 92     };
 93 }
 94 
 95 fn loadOpcode(element_type_name: []const u8) error{CodeGenFailed}!u8 {
 96     const kind = dialects.arith.scalarKindFromTypeName(element_type_name) orelse
 97         return error.CodeGenFailed;
 98     return switch (kind) {
 99         .i8 => 0x2c,
100         .u8 => 0x2d,
101         .i16 => 0x2e,
102         .u16 => 0x2f,
103         .i32, .u32, .index, .bool => 0x28,
104         .i64, .u64 => 0x29,
105         .f32 => 0x2a,
106         .f64 => 0x2b,
107         .f16, .bf16 => error.CodeGenFailed,
108     };
109 }
110 
111 fn storeOpcode(element_type_name: []const u8) error{CodeGenFailed}!u8 {
112     const kind = dialects.arith.scalarKindFromTypeName(element_type_name) orelse
113         return error.CodeGenFailed;
114     return switch (kind) {
115         .i8, .u8 => 0x3a,
116         .i16, .u16 => 0x3b,
117         .i32, .u32, .index, .bool => 0x36,
118         .i64, .u64 => 0x37,
119         .f32 => 0x38,
120         .f64 => 0x39,
121         .f16, .bf16 => error.CodeGenFailed,
122     };
123 }
124 
125 test "wasm memory opcodes preserve element width and signedness" {
126     try std.testing.expectEqual(@as(u32, 1), try elementByteSize("arith.i8"));
127     try std.testing.expectEqual(@as(u32, 2), try elementByteSize("arith.f16"));
128     try std.testing.expectEqual(@as(u32, 4), try elementByteSize("arith.index"));
129     try std.testing.expectEqual(@as(u32, 8), try elementByteSize("arith.f64"));
130     try std.testing.expectEqual(@as(u8, 0x2c), try loadOpcode("arith.i8"));
131     try std.testing.expectEqual(@as(u8, 0x2d), try loadOpcode("arith.u8"));
132     try std.testing.expectEqual(@as(u8, 0x3a), try storeOpcode("arith.i8"));
133     try std.testing.expectEqual(@as(u8, 0x3a), try storeOpcode("arith.u8"));
134     try std.testing.expectEqual(@as(u32, 0), alignmentPower(1));
135     try std.testing.expectEqual(@as(u32, 3), alignmentPower(8));
136 }
137 
138 test "wasm memory accepts only host-visible address spaces" {
139     const base = MemrefDialect.MemrefParams{
140         .size = null,
141         .element_type_name = "arith.i32",
142         .addr_space = .host,
143     };
144     try requireMemoryAddressSpace(base);
145 
146     var unified = base;
147     unified.addr_space = .unified;
148     try requireMemoryAddressSpace(unified);
149 
150     var device = base;
151     device.addr_space = .device;
152     try std.testing.expectError(error.CodeGenFailed, requireMemoryAddressSpace(device));
153 }