lib/choir/src/backends/gpu/spirv/emitter/memory.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const choir = @import("../../../../root.zig");
3
4 const ir = choir.ir;
5 const dialects = choir.dialects;
6 const scalar = @import("scalar.zig");
7 const spec = @import("spec.zig");
8 const spirv_ops = @import("ops.zig");
9
10 const MemrefDialect = dialects.memref.MemrefDialect;
11 const MemrefAddressSpace = dialects.AddressSpace;
12 const AtomicRmwKind = dialects.memref.AtomicRmwKind;
13 const SpirvOp = spirv_ops.SpirvOp;
14
15 pub const Layout = enum { buffer, array };
16
17 pub const Binding = struct {
18 storage_class: u32,
19 elem_kind: scalar.Kind,
20 elem_type: u32,
21 storage_elem_kind: scalar.Kind,
22 storage_elem_type: u32,
23 ptr_elem_type: u32,
24 layout: Layout,
25 };
26
27 pub fn emitLoad(self: anytype, op: *ir.Operation) !void {
28 const load = MemrefDialect.LoadOp{ .op = op };
29 const result = load.getResult();
30
31 const memref = load.getMemref();
32 const index = load.getIndex();
33
34 const binding = self.memref_bindings.get(memref) orelse return error.InvalidMemrefType;
35 const result_type_id = try self.getTypeForValue(result);
36 if (result_type_id != binding.elem_type) return error.UnsupportedType;
37 const access_id = try emitElementAccess(self, binding, memref, index);
38
39 const loaded_id = self.builder.newId();
40 try self.builder.emit(&self.builder.functions, SpirvOp.Load, &.{
41 binding.storage_elem_type,
42 loaded_id,
43 access_id,
44 });
45
46 const result_id = try emitLoadResult(self, binding, loaded_id);
47 try self.bindValue(result, result_id);
48 }
49
50 pub fn emitStore(self: anytype, op: *ir.Operation) !void {
51 const store = MemrefDialect.StoreOp{ .op = op };
52 const value = store.getValue();
53 const memref = store.getMemref();
54 const index = store.getIndex();
55
56 const binding = self.memref_bindings.get(memref) orelse return error.InvalidMemrefType;
57 const value_type_id = try self.getTypeForValue(value);
58 if (value_type_id != binding.elem_type) return error.UnsupportedType;
59 const value_id = try self.getValue(value);
60 const access_id = try emitElementAccess(self, binding, memref, index);
61 const storage_value_id = try emitStoreValue(self, binding, value_id);
62
63 try self.builder.emit(&self.builder.functions, SpirvOp.Store, &.{
64 access_id,
65 storage_value_id,
66 });
67 }
68
69 pub fn emitAtomicRmw(self: anytype, op: *ir.Operation) !void {
70 const atomic = MemrefDialect.AtomicRmwOp{ .op = op };
71 const result = atomic.getResult();
72 const value = atomic.getValue();
73 const memref = atomic.getMemref();
74 const index = atomic.getIndex();
75
76 const binding = self.memref_bindings.get(memref) orelse return error.InvalidMemrefType;
77 const result_type_id = try self.getTypeForValue(result);
78 const value_type_id = try self.getTypeForValue(value);
79 if (result_type_id != binding.elem_type or value_type_id != binding.elem_type) {
80 return error.UnsupportedType;
81 }
82 if (binding.storage_elem_type != binding.elem_type) return error.UnsupportedType;
83
84 const kind = scalar.kindFromType(result.type) orelse return error.UnsupportedType;
85 const opcode = atomicRmwOpcode(atomic.getKind() orelse return error.MissingAttribute, kind) orelse return error.UnsupportedType;
86 const memory_operands = atomicMemoryOperands(binding.storage_class) orelse return error.UnsupportedAddressSpace;
87 const u32_type = try self.getScalarType(.u32);
88 const scope_id = try self.getIntConstant(u32_type, .u32, memory_operands.scope);
89 const semantics_id = try self.getIntConstant(
90 u32_type,
91 .u32,
92 spec.MemorySemanticsMask.AcquireRelease | memory_operands.semantics,
93 );
94 const access_id = try emitElementAccess(self, binding, memref, index);
95 const value_id = try self.getValue(value);
96
97 const result_id = self.builder.newId();
98 try self.builder.emit(&self.builder.functions, opcode, &.{
99 result_type_id,
100 result_id,
101 access_id,
102 scope_id,
103 semantics_id,
104 value_id,
105 });
106 try self.bindValue(result, result_id);
107 }
108
109 pub fn emitAtomicCas(self: anytype, op: *ir.Operation) !void {
110 const atomic = MemrefDialect.AtomicCasOp{ .op = op };
111 const result = atomic.getResult();
112 const expected = atomic.getExpected();
113 const desired = atomic.getDesired();
114 const memref = atomic.getMemref();
115 const index = atomic.getIndex();
116
117 const binding = self.memref_bindings.get(memref) orelse return error.InvalidMemrefType;
118 const result_type_id = try self.getTypeForValue(result);
119 const expected_type_id = try self.getTypeForValue(expected);
120 const desired_type_id = try self.getTypeForValue(desired);
121 if (result_type_id != binding.elem_type or
122 expected_type_id != binding.elem_type or
123 desired_type_id != binding.elem_type)
124 {
125 return error.UnsupportedType;
126 }
127 if (binding.storage_elem_type != binding.elem_type) return error.UnsupportedType;
128
129 const kind = scalar.kindFromType(result.type) orelse return error.UnsupportedType;
130 if (!atomicScalarSupported(kind)) return error.UnsupportedType;
131
132 const memory_operands = atomicMemoryOperands(binding.storage_class) orelse return error.UnsupportedAddressSpace;
133 const u32_type = try self.getScalarType(.u32);
134 const scope_id = try self.getIntConstant(u32_type, .u32, memory_operands.scope);
135 const equal_semantics_id = try self.getIntConstant(
136 u32_type,
137 .u32,
138 spec.MemorySemanticsMask.AcquireRelease | memory_operands.semantics,
139 );
140 const unequal_semantics_id = try self.getIntConstant(
141 u32_type,
142 .u32,
143 spec.MemorySemanticsMask.Acquire | memory_operands.semantics,
144 );
145 const access_id = try emitElementAccess(self, binding, memref, index);
146 const expected_id = try self.getValue(expected);
147 const desired_id = try self.getValue(desired);
148
149 const result_id = self.builder.newId();
150 try self.builder.emit(&self.builder.functions, SpirvOp.AtomicCompareExchange, &.{
151 result_type_id,
152 result_id,
153 access_id,
154 scope_id,
155 equal_semantics_id,
156 unequal_semantics_id,
157 desired_id,
158 expected_id,
159 });
160 try self.bindValue(result, result_id);
161 }
162
163 pub fn emitAlloc(self: anytype, op: *ir.Operation) !void {
164 const alloc = MemrefDialect.AllocOp{ .op = op };
165 const result = alloc.getResult();
166 const memref = parseType(result.type) orelse return error.InvalidMemrefType;
167
168 if (memref.addr_space != .shared) return error.UnsupportedAddressSpace;
169 const size = memref.size orelse return error.InvalidMemrefType;
170
171 const elem_kind = scalar.kindFromName(memref.element_type_name) orelse return error.UnsupportedType;
172 const elem_type_id = try self.getScalarType(elem_kind);
173 const storage_elem_kind = storageElementKind(elem_kind);
174 const storage_elem_type_id = try self.getScalarType(storage_elem_kind);
175 const array_type = try self.getArrayType(storage_elem_type_id, @intCast(size));
176 const ptr_array = try self.getPointerType(spec.StorageClass.Workgroup, array_type);
177 const ptr_elem = try self.getPointerType(spec.StorageClass.Workgroup, storage_elem_type_id);
178
179 const var_id = self.builder.newId();
180 try self.builder.emit(&self.builder.globals, SpirvOp.Variable, &.{
181 ptr_array,
182 var_id,
183 spec.StorageClass.Workgroup,
184 });
185
186 try self.bindValue(result, var_id);
187 try self.memref_bindings.put(self.allocator, result, .{
188 .storage_class = spec.StorageClass.Workgroup,
189 .elem_kind = elem_kind,
190 .elem_type = elem_type_id,
191 .storage_elem_kind = storage_elem_kind,
192 .storage_elem_type = storage_elem_type_id,
193 .ptr_elem_type = ptr_elem,
194 .layout = .array,
195 });
196 }
197
198 /// Function variables must precede every ordinary instruction in the entry
199 /// block, including instructions that appear before an alloca in the IR.
200 pub fn declareAlloca(self: anytype, op: *ir.Operation) !void {
201 const alloca = MemrefDialect.AllocaOp{ .op = op };
202 const result = alloca.getResult();
203 const memref = parseType(result.type) orelse return error.InvalidMemrefType;
204 if (memref.addr_space != .local) return error.UnsupportedAddressSpace;
205 if (alloca.getDynamicSize() != null) return error.UnsupportedOperation;
206 const size = memref.size orelse return error.UnsupportedOperation;
207 if (size == 0 or size > std.math.maxInt(u32)) return error.UnsupportedOperation;
208
209 const elem_kind = scalar.kindFromName(memref.element_type_name) orelse return error.UnsupportedType;
210 const elem_type = try self.getScalarType(elem_kind);
211 const storage_kind = storageElementKind(elem_kind);
212 const storage_type = try self.getScalarType(storage_kind);
213 const array_type = try self.getArrayType(storage_type, @intCast(size));
214 const ptr_array = try self.getPointerType(spec.StorageClass.Function, array_type);
215 const ptr_elem = try self.getPointerType(spec.StorageClass.Function, storage_type);
216 const var_id = self.builder.newId();
217 try self.builder.emit(&self.builder.functions, SpirvOp.Variable, &.{
218 ptr_array,
219 var_id,
220 spec.StorageClass.Function,
221 });
222 try self.bindValue(result, var_id);
223 try self.memref_bindings.put(self.allocator, result, .{
224 .storage_class = spec.StorageClass.Function,
225 .elem_kind = elem_kind,
226 .elem_type = elem_type,
227 .storage_elem_kind = storage_kind,
228 .storage_elem_type = storage_type,
229 .ptr_elem_type = ptr_elem,
230 .layout = .array,
231 });
232 }
233
234 pub fn emitDeclaredAlloca(self: anytype, op: *ir.Operation) !void {
235 const result = (MemrefDialect.AllocaOp{ .op = op }).getResult();
236 if (self.memref_bindings.get(result) == null) return error.InvalidMemrefType;
237 }
238
239 fn emitLoadResult(self: anytype, binding: Binding, loaded_id: u32) !u32 {
240 if (binding.storage_elem_type == binding.elem_type) return loaded_id;
241 if (binding.elem_kind != .bool or binding.storage_elem_kind != .u8) return error.UnsupportedType;
242
243 const zero_id = try self.getIntConstant(binding.storage_elem_type, binding.storage_elem_kind, 0);
244 const result_id = self.builder.newId();
245 try self.builder.emit(&self.builder.functions, SpirvOp.INotEqual, &.{
246 binding.elem_type,
247 result_id,
248 loaded_id,
249 zero_id,
250 });
251 return result_id;
252 }
253
254 fn emitStoreValue(self: anytype, binding: Binding, value_id: u32) !u32 {
255 if (binding.storage_elem_type == binding.elem_type) return value_id;
256 if (binding.elem_kind != .bool or binding.storage_elem_kind != .u8) return error.UnsupportedType;
257
258 const zero_id = try self.getIntConstant(binding.storage_elem_type, binding.storage_elem_kind, 0);
259 const one_id = try self.getIntConstant(binding.storage_elem_type, binding.storage_elem_kind, 1);
260 const result_id = self.builder.newId();
261 try self.builder.emit(&self.builder.functions, SpirvOp.Select, &.{
262 binding.storage_elem_type,
263 result_id,
264 value_id,
265 one_id,
266 zero_id,
267 });
268 return result_id;
269 }
270
271 fn emitElementAccess(
272 self: anytype,
273 binding: Binding,
274 memref: *ir.Value,
275 index: *ir.Value,
276 ) !u32 {
277 const memref_id = try self.getValue(memref);
278 const index_id = try self.getValue(index);
279 const access_id = self.builder.newId();
280 switch (binding.layout) {
281 .buffer => {
282 const const_zero = try self.getIntConstant(
283 try self.getScalarType(.u32),
284 .u32,
285 0,
286 );
287 try self.builder.emit(&self.builder.functions, SpirvOp.AccessChain, &.{
288 binding.ptr_elem_type,
289 access_id,
290 memref_id,
291 const_zero,
292 index_id,
293 });
294 },
295 .array => {
296 try self.builder.emit(&self.builder.functions, SpirvOp.AccessChain, &.{
297 binding.ptr_elem_type,
298 access_id,
299 memref_id,
300 index_id,
301 });
302 },
303 }
304 return access_id;
305 }
306
307 pub fn decorateBufferVariable(
308 self: anytype,
309 var_id: u32,
310 binding: u32,
311 storage_class: u32,
312 addr_space: MemrefAddressSpace,
313 ) !void {
314 _ = storage_class;
315 try self.builder.emit(&self.builder.annotations, SpirvOp.Decorate, &.{
316 var_id,
317 spec.Decoration.DescriptorSet,
318 0,
319 });
320 try self.builder.emit(&self.builder.annotations, SpirvOp.Decorate, &.{
321 var_id,
322 spec.Decoration.Binding,
323 binding,
324 });
325
326 if (addr_space == .constant) {
327 try self.builder.emit(&self.builder.annotations, SpirvOp.Decorate, &.{
328 var_id,
329 spec.Decoration.NonWritable,
330 });
331 }
332 }
333
334 pub fn getBufferLayout(
335 self: anytype,
336 storage_elem_type: u32,
337 storage_elem_kind: scalar.Kind,
338 storage_class: u32,
339 ) !struct {
340 ptr_struct_type: u32,
341 ptr_elem_type: u32,
342 } {
343 const stride = scalar.elementByteSize(storage_elem_kind) orelse return error.UnsupportedType;
344 const runtime_array = try self.getRuntimeArrayType(storage_elem_type, @intCast(stride));
345 const struct_type = try self.getStructType(runtime_array);
346
347 const ptr_struct = try self.getPointerType(storage_class, struct_type);
348 const ptr_elem = try self.getPointerType(storage_class, storage_elem_type);
349 return .{ .ptr_struct_type = ptr_struct, .ptr_elem_type = ptr_elem };
350 }
351
352 pub fn storageElementKind(elem_kind: scalar.Kind) scalar.Kind {
353 return switch (elem_kind) {
354 .bool => .u8,
355 else => elem_kind,
356 };
357 }
358
359 pub fn parseType(ty: ir.Type) ?MemrefDialect.MemrefParams {
360 const name = ty.getDialectTypeName() orelse return null;
361 if (!std.mem.eql(u8, name, MemrefDialect.name)) return null;
362 const key = ty.getDialectParamKey() orelse return null;
363 return MemrefDialect.parseMemrefParams(key);
364 }
365
366 pub fn storageClassForAddressSpace(addr_space: MemrefAddressSpace) ?u32 {
367 return switch (addr_space) {
368 .device, .host, .unified => spec.StorageClass.StorageBuffer,
369 .constant => spec.StorageClass.Uniform,
370 .shared => spec.StorageClass.Workgroup,
371 .local => null,
372 };
373 }
374
375 fn atomicRmwOpcode(kind: AtomicRmwKind, elem_kind: scalar.Kind) ?u16 {
376 if (!atomicScalarSupported(elem_kind)) return null;
377 return switch (kind) {
378 .add => SpirvOp.AtomicIAdd,
379 .min => if (scalar.isSignedInt(elem_kind)) SpirvOp.AtomicSMin else SpirvOp.AtomicUMin,
380 .max => if (scalar.isSignedInt(elem_kind)) SpirvOp.AtomicSMax else SpirvOp.AtomicUMax,
381 .bit_and => SpirvOp.AtomicAnd,
382 .bit_or => SpirvOp.AtomicOr,
383 .bit_xor => SpirvOp.AtomicXor,
384 .exchange => SpirvOp.AtomicExchange,
385 };
386 }
387
388 fn atomicScalarSupported(kind: scalar.Kind) bool {
389 return switch (kind) {
390 .i32, .u32 => true,
391 else => false,
392 };
393 }
394
395 fn atomicMemoryOperands(storage_class: u32) ?struct {
396 scope: u32,
397 semantics: u32,
398 } {
399 if (storage_class == spec.StorageClass.Workgroup) {
400 return .{
401 .scope = spec.Scope.Workgroup,
402 .semantics = spec.MemorySemanticsMask.WorkgroupMemory,
403 };
404 }
405 if (storage_class == spec.StorageClass.StorageBuffer) {
406 return .{
407 .scope = spec.Scope.Device,
408 .semantics = spec.MemorySemanticsMask.UniformMemory,
409 };
410 }
411 return null;
412 }
413
414 test "spirv memory owner maps address spaces" {
415 try std.testing.expectEqual(@as(u32, spec.StorageClass.Workgroup), storageClassForAddressSpace(.shared).?);
416 try std.testing.expectEqual(@as(u32, spec.StorageClass.StorageBuffer), storageClassForAddressSpace(.device).?);
417 try std.testing.expect(storageClassForAddressSpace(.constant).? == spec.StorageClass.Uniform);
418 }
419
420 test "spirv memory owner maps atomic opcodes and memory operands" {
421 try std.testing.expectEqual(@as(u16, SpirvOp.AtomicIAdd), atomicRmwOpcode(.add, .i32).?);
422 try std.testing.expectEqual(@as(u16, SpirvOp.AtomicUMin), atomicRmwOpcode(.min, .u32).?);
423 try std.testing.expectEqual(@as(u16, SpirvOp.AtomicSMax), atomicRmwOpcode(.max, .i32).?);
424 try std.testing.expectEqual(@as(u16, SpirvOp.AtomicAnd), atomicRmwOpcode(.bit_and, .u32).?);
425 try std.testing.expectEqual(@as(?u16, null), atomicRmwOpcode(.add, .f32));
426
427 const workgroup = atomicMemoryOperands(spec.StorageClass.Workgroup).?;
428 try std.testing.expectEqual(spec.Scope.Workgroup, workgroup.scope);
429 try std.testing.expectEqual(spec.MemorySemanticsMask.WorkgroupMemory, workgroup.semantics);
430
431 const storage_buffer = atomicMemoryOperands(spec.StorageClass.StorageBuffer).?;
432 try std.testing.expectEqual(spec.Scope.Device, storage_buffer.scope);
433 try std.testing.expectEqual(spec.MemorySemanticsMask.UniformMemory, storage_buffer.semantics);
434 try std.testing.expect(atomicMemoryOperands(spec.StorageClass.Uniform) == null);
435 }