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 }