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 }