lib/choir/src/backends/gpu/features.zig

daab053ee43316e1809a84551d573ddd1e5bf3d2

  1 const std = @import("std");
  2 const abi = @import("choir_abi");
  3 const choir = @import("../../root.zig");
  4 const gpu = @import("../../dialects/gpu/root.zig");
  5 
  6 const ir = choir.ir;
  7 const dialects = choir.dialects;
  8 const arith = dialects.arith;
  9 const MemrefDialect = dialects.memref.MemrefDialect;
 10 const GpuDialect = gpu.GpuDialect;
 11 
 12 pub fn requirementsForModule(module: *ir.Operation) abi.Features {
 13     var requirements: abi.Features = .{};
 14     appendRequirements(&requirements, module);
 15     return requirements;
 16 }
 17 
 18 pub fn requirementsForOperation(op: *ir.Operation) abi.Features {
 19     const name = op.name.name;
 20     if (isName(name, MemrefDialect.AllocOp.operation_name)) return memrefAllocRequirements(MemrefDialect.AllocOp{ .op = op });
 21     if (isName(name, MemrefDialect.AtomicRmwOp.operation_name)) return memrefAtomicRmwRequirements(MemrefDialect.AtomicRmwOp{ .op = op });
 22     if (isName(name, MemrefDialect.AtomicCasOp.operation_name)) return memrefAtomicCasRequirements(MemrefDialect.AtomicCasOp{ .op = op });
 23     if (isName(name, GpuDialect.AtomicLoadOp.operation_name)) return .{ .unsupported_atomic = true };
 24     if (isName(name, GpuDialect.AtomicStoreOp.operation_name)) return .{ .unsupported_atomic = true };
 25     if (isName(name, GpuDialect.AtomicAddOp.operation_name)) return .{ .unsupported_atomic = true };
 26     if (isName(name, GpuDialect.AtomicMaxOp.operation_name)) return .{ .unsupported_atomic = true };
 27     if (isName(name, GpuDialect.AtomicCasOp.operation_name)) return .{ .unsupported_atomic = true };
 28     if (isName(name, GpuDialect.MemcpyAsyncOp.operation_name)) return .{ .async_copy = true };
 29     if (isName(name, GpuDialect.MmaSyncOp.operation_name)) return .{ .tensor_cores = true };
 30     if (isName(name, GpuDialect.CpAsyncSharedOp.operation_name)) return .{ .tensor_cores = true };
 31     if (isName(name, GpuDialect.CpAsyncCommitOp.operation_name)) return .{ .tensor_cores = true };
 32     if (isName(name, GpuDialect.CpAsyncWaitOp.operation_name)) return .{ .tensor_cores = true };
 33     if (isName(name, arith.ArithDialect.Tf32RoundOp.operation_name)) return .{ .tensor_cores = true };
 34     return .{};
 35 }
 36 
 37 fn appendRequirements(requirements: *abi.Features, op: *ir.Operation) void {
 38     mergeRequirements(requirements, requirementsForOperation(op));
 39     for (op.regions.items) |*region| {
 40         var block_iter = region.getBlocks();
 41         while (block_iter.next()) |block| {
 42             var current: ?*ir.Operation = @ptrCast(@alignCast(block.operations.head));
 43             while (current) |current_op| {
 44                 appendRequirements(requirements, current_op);
 45                 current = current_op.next_op;
 46             }
 47         }
 48     }
 49 }
 50 
 51 fn mergeRequirements(requirements: *abi.Features, required: abi.Features) void {
 52     requirements.atomic_i32 = requirements.atomic_i32 or required.atomic_i32;
 53     requirements.atomic_u32 = requirements.atomic_u32 or required.atomic_u32;
 54     requirements.atomic_index = requirements.atomic_index or required.atomic_index;
 55     requirements.atomic_f32_add_device = requirements.atomic_f32_add_device or required.atomic_f32_add_device;
 56     requirements.atomic_f32_add_shared = requirements.atomic_f32_add_shared or required.atomic_f32_add_shared;
 57     requirements.unsupported_atomic = requirements.unsupported_atomic or required.unsupported_atomic;
 58     requirements.async_copy = requirements.async_copy or required.async_copy;
 59     requirements.tensor_cores = requirements.tensor_cores or required.tensor_cores;
 60     requirements.cooperative_matrix = requirements.cooperative_matrix or required.cooperative_matrix;
 61     requirements.dynamic_shared_memory = requirements.dynamic_shared_memory or required.dynamic_shared_memory;
 62     requirements.indirect_launch = requirements.indirect_launch or required.indirect_launch;
 63 }
 64 
 65 fn isName(actual: []const u8, expected: []const u8) bool {
 66     return std.mem.eql(u8, actual, expected);
 67 }
 68 
 69 fn memrefAllocRequirements(op: MemrefDialect.AllocOp) abi.Features {
 70     if (op.getDynamicSize() == null) return .{};
 71     const params = memrefParams(op.getResult().type) orelse return .{};
 72     return switch (params.addr_space) {
 73         .shared => .{ .dynamic_shared_memory = true },
 74         .host, .device, .constant, .unified, .local => .{},
 75     };
 76 }
 77 
 78 fn memrefAtomicRmwRequirements(op: MemrefDialect.AtomicRmwOp) abi.Features {
 79     const params = memrefParams(op.getMemref().type) orelse return .{ .unsupported_atomic = true };
 80     const kind = op.getKind() orelse return .{ .unsupported_atomic = true };
 81     const scalar = arith.scalarKindFromTypeName(params.element_type_name) orelse return .{ .unsupported_atomic = true };
 82     return atomicRmwRequirements(kind, scalar, params.addr_space);
 83 }
 84 
 85 fn memrefAtomicCasRequirements(op: MemrefDialect.AtomicCasOp) abi.Features {
 86     const params = memrefParams(op.getMemref().type) orelse return .{ .unsupported_atomic = true };
 87     if (params.addr_space == .local) return .{ .unsupported_atomic = true };
 88     const scalar = arith.scalarKindFromTypeName(params.element_type_name) orelse return .{ .unsupported_atomic = true };
 89     return atomicIntegerRequirements(scalar);
 90 }
 91 
 92 fn memrefParams(memref_type: ir.Type) ?MemrefDialect.MemrefParams {
 93     if (!isName(memref_type.getDialectTypeName() orelse return null, MemrefDialect.name)) return null;
 94     return MemrefDialect.parseMemrefParams(memref_type.getDialectParamKey() orelse return null);
 95 }
 96 
 97 fn atomicRmwRequirements(
 98     kind: dialects.AtomicRmwKind,
 99     scalar: arith.ScalarKind,
100     addr_space: dialects.AddressSpace,
101 ) abi.Features {
102     if (addr_space == .local) return .{ .unsupported_atomic = true };
103     switch (scalar) {
104         .i32, .u32, .index => return atomicIntegerRequirements(scalar),
105         .f32 => {
106             if (kind != .add) return .{ .unsupported_atomic = true };
107             return switch (addr_space) {
108                 .host, .device, .unified => .{ .atomic_f32_add_device = true },
109                 .shared => .{ .atomic_f32_add_shared = true },
110                 .constant, .local => .{ .unsupported_atomic = true },
111             };
112         },
113         .i8, .i16, .i64, .u8, .u16, .u64, .f16, .bf16, .f64, .bool => return .{ .unsupported_atomic = true },
114     }
115 }
116 
117 fn atomicIntegerRequirements(scalar: arith.ScalarKind) abi.Features {
118     return switch (scalar) {
119         .i32 => .{ .atomic_i32 = true },
120         .u32 => .{ .atomic_u32 = true },
121         .index => .{ .atomic_index = true },
122         else => .{ .unsupported_atomic = true },
123     };
124 }
125 
126 test "target features infer atomics from memref atomic operations" {
127     var ctx = try ir.Context.init(std.testing.allocator, ir.Context.Limits.testing);
128     defer ctx.deinit(std.testing.allocator);
129 
130     const loc = ir.Location.getUnknown();
131     const i32_type = try dialects.arith.ArithDialect.getI32Type(&ctx);
132     const index_type = try dialects.arith.ArithDialect.getIndexType(&ctx);
133     const memref_type = try MemrefDialect.getMemrefTypeDynamic(&ctx, i32_type, .device);
134     const alloc = try MemrefDialect.AllocOp.createStatic(&ctx, loc, memref_type);
135     const idx = try dialects.arith.ArithDialect.ConstantOp.createInt(&ctx, loc, index_type, 0);
136     const value = try dialects.arith.ArithDialect.ConstantOp.createInt(&ctx, loc, i32_type, 1);
137     const atomic = try MemrefDialect.AtomicRmwOp.create(
138         &ctx,
139         loc,
140         .add,
141         value.getResult(),
142         alloc.getResult(),
143         idx.getResult(),
144         i32_type,
145     );
146 
147     try std.testing.expect(requirementsForOperation(atomic.op).atomic_i32);
148 }
149 
150 test "target features infer dynamic shared memory from shared dynamic allocation" {
151     var ctx = try ir.Context.init(std.testing.allocator, ir.Context.Limits.testing);
152     defer ctx.deinit(std.testing.allocator);
153 
154     const loc = ir.Location.getUnknown();
155     const f32_type = try dialects.arith.ArithDialect.getScalarType(&ctx, .f32);
156     const index_type = try dialects.arith.ArithDialect.getIndexType(&ctx);
157     const size = try dialects.arith.ArithDialect.ConstantOp.createInt(&ctx, loc, index_type, 16);
158 
159     const shared_type = try MemrefDialect.getMemrefType1D(&ctx, 16, f32_type, .shared);
160     const shared = try MemrefDialect.AllocOp.createDynamic(&ctx, loc, size.getResult(), shared_type);
161     const shared_requirements = requirementsForOperation(shared.op);
162     try std.testing.expect(shared_requirements.dynamic_shared_memory);
163 
164     const static_shared = try MemrefDialect.AllocOp.createStatic(&ctx, loc, shared_type);
165     try std.testing.expect(!requirementsForOperation(static_shared.op).dynamic_shared_memory);
166 
167     const device_type = try MemrefDialect.getMemrefType1D(&ctx, 16, f32_type, .device);
168     const device = try MemrefDialect.AllocOp.createDynamic(&ctx, loc, size.getResult(), device_type);
169     try std.testing.expect(!requirementsForOperation(device.op).dynamic_shared_memory);
170 }
171 
172 test "target features classify f32 atomic add by address space" {
173     var ctx = try ir.Context.init(std.testing.allocator, ir.Context.Limits.testing);
174     defer ctx.deinit(std.testing.allocator);
175 
176     const loc = ir.Location.getUnknown();
177     const f32_type = try dialects.arith.ArithDialect.getScalarType(&ctx, .f32);
178     const index_type = try dialects.arith.ArithDialect.getIndexType(&ctx);
179     const idx = try dialects.arith.ArithDialect.ConstantOp.createInt(&ctx, loc, index_type, 0);
180     const value = try dialects.arith.ArithDialect.ConstantOp.createFloat(&ctx, loc, f32_type, 1.0);
181 
182     const device_type = try MemrefDialect.getMemrefTypeDynamic(&ctx, f32_type, .device);
183     const device = try MemrefDialect.AllocOp.createStatic(&ctx, loc, device_type);
184     const device_atomic = try MemrefDialect.AtomicRmwOp.create(
185         &ctx,
186         loc,
187         .add,
188         value.getResult(),
189         device.getResult(),
190         idx.getResult(),
191         f32_type,
192     );
193     const device_requirements = requirementsForOperation(device_atomic.op);
194     try std.testing.expect(device_requirements.atomic_f32_add_device);
195     try std.testing.expect(!device_requirements.atomic_f32_add_shared);
196     try std.testing.expect(!device_requirements.unsupported_atomic);
197 
198     const shared_type = try MemrefDialect.getMemrefTypeDynamic(&ctx, f32_type, .shared);
199     const shared = try MemrefDialect.AllocOp.createStatic(&ctx, loc, shared_type);
200     const shared_atomic = try MemrefDialect.AtomicRmwOp.create(
201         &ctx,
202         loc,
203         .add,
204         value.getResult(),
205         shared.getResult(),
206         idx.getResult(),
207         f32_type,
208     );
209     const shared_requirements = requirementsForOperation(shared_atomic.op);
210     try std.testing.expect(!shared_requirements.atomic_f32_add_device);
211     try std.testing.expect(shared_requirements.atomic_f32_add_shared);
212     try std.testing.expect(!shared_requirements.unsupported_atomic);
213 }
214 
215 test "target features mark unsupported atomic dtype and float kind" {
216     var ctx = try ir.Context.init(std.testing.allocator, ir.Context.Limits.testing);
217     defer ctx.deinit(std.testing.allocator);
218 
219     const loc = ir.Location.getUnknown();
220     const f32_type = try dialects.arith.ArithDialect.getScalarType(&ctx, .f32);
221     const f64_type = try dialects.arith.ArithDialect.getF64Type(&ctx);
222     const index_type = try dialects.arith.ArithDialect.getIndexType(&ctx);
223     const idx = try dialects.arith.ArithDialect.ConstantOp.createInt(&ctx, loc, index_type, 0);
224     const f32_value = try dialects.arith.ArithDialect.ConstantOp.createFloat(&ctx, loc, f32_type, 1.0);
225     const f64_value = try dialects.arith.ArithDialect.ConstantOp.createFloat(&ctx, loc, f64_type, 1.0);
226 
227     const f32_memref_type = try MemrefDialect.getMemrefTypeDynamic(&ctx, f32_type, .device);
228     const f32_memref = try MemrefDialect.AllocOp.createStatic(&ctx, loc, f32_memref_type);
229     const exchange = try MemrefDialect.AtomicRmwOp.create(
230         &ctx,
231         loc,
232         .exchange,
233         f32_value.getResult(),
234         f32_memref.getResult(),
235         idx.getResult(),
236         f32_type,
237     );
238     try std.testing.expect(requirementsForOperation(exchange.op).unsupported_atomic);
239 
240     const f64_memref_type = try MemrefDialect.getMemrefTypeDynamic(&ctx, f64_type, .device);
241     const f64_memref = try MemrefDialect.AllocOp.createStatic(&ctx, loc, f64_memref_type);
242     const add = try MemrefDialect.AtomicRmwOp.create(
243         &ctx,
244         loc,
245         .add,
246         f64_value.getResult(),
247         f64_memref.getResult(),
248         idx.getResult(),
249         f64_type,
250     );
251     try std.testing.expect(requirementsForOperation(add.op).unsupported_atomic);
252 }
253 
254 test "target features infer async copy from gpu memcpy async operations" {
255     var ctx = try ir.Context.init(std.testing.allocator, ir.Context.Limits.testing);
256     defer ctx.deinit(std.testing.allocator);
257 
258     const loc = ir.Location.getUnknown();
259     const i32_type = try dialects.arith.ArithDialect.getI32Type(&ctx);
260     const index_type = try dialects.arith.ArithDialect.getIndexType(&ctx);
261     const memref_type = try MemrefDialect.getMemrefType1D(&ctx, 16, i32_type, .device);
262     const src = try MemrefDialect.AllocOp.createStatic(&ctx, loc, memref_type);
263     const dst = try MemrefDialect.AllocOp.createStatic(&ctx, loc, memref_type);
264     const size = try dialects.arith.ArithDialect.ConstantOp.createInt(&ctx, loc, index_type, 16);
265     const copy = try GpuDialect.MemcpyAsyncOp.create(
266         &ctx,
267         loc,
268         src.getResult(),
269         dst.getResult(),
270         size.getResult(),
271         null,
272         null,
273     );
274 
275     const requirements = requirementsForOperation(copy.op);
276     try std.testing.expect(requirements.async_copy);
277     try std.testing.expect(!requirements.atomic_i32);
278     try std.testing.expect(!requirements.atomic_f32_add_device);
279     try std.testing.expect(!requirements.unsupported_atomic);
280 }
281 
282 test "target features infer tensor cores from gpu mma operations" {
283     var ctx = try ir.Context.init(std.testing.allocator, ir.Context.Limits.testing);
284     defer ctx.deinit(std.testing.allocator);
285 
286     const loc = ir.Location.getUnknown();
287     const f32_type = try dialects.arith.ArithDialect.getScalarType(&ctx, .f32);
288     const shape = gpu.MmaShape{ .m = 16, .n = 8, .k = 8 };
289 
290     try ctx.allowUnregistered();
291     var builder = ir.OperationBuilder.init(&ctx);
292     var source_state = ir.Operation.State.init("test.frag_source", loc);
293     source_state.addTypes(&.{ f32_type, f32_type, f32_type, f32_type, f32_type, f32_type, f32_type, f32_type, f32_type, f32_type });
294     const source = try builder.create(source_state);
295 
296     const mma = try GpuDialect.MmaSyncOp.create(&ctx, loc, .{
297         source.getResult(0).?,
298         source.getResult(1).?,
299         source.getResult(2).?,
300         source.getResult(3).?,
301     }, .{
302         source.getResult(4).?,
303         source.getResult(5).?,
304     }, .{
305         source.getResult(6).?,
306         source.getResult(7).?,
307         source.getResult(8).?,
308         source.getResult(9).?,
309     }, shape);
310 
311     const round = try dialects.arith.ArithDialect.Tf32RoundOp.create(&ctx, loc, source.getResult(0).?);
312 
313     try std.testing.expect(requirementsForOperation(mma.op).tensor_cores);
314     try std.testing.expect(requirementsForOperation(round.op).tensor_cores);
315     try std.testing.expect(!requirementsForOperation(source).tensor_cores);
316 }