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 }