lib/choir/src/backends/gpu/nvptx/conversion.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const choir = @import("../../../root.zig");
3
4 const ir = choir.ir;
5 const rewrite = ir.rewrite;
6 const passes = choir.passes;
7 const dialects = choir.dialects;
8 const backends = choir.backends;
9 const gpu = @import("../../../dialects/gpu/root.zig");
10 const nvptx = @import("dialect.zig");
11
12 const Pass = passes.Pass;
13 const PassContext = passes.PassContext;
14 const PassResult = passes.PassResult;
15
16 const ConversionTarget = passes.ConversionTarget;
17 const RewritePatternSet = rewrite.RewritePatternSet;
18 const RewritePattern = rewrite.RewritePattern;
19 const PatternRewriter = rewrite.PatternRewriter;
20
21 const ArithDialect = dialects.ArithDialect;
22 const GpuDialect = gpu.GpuDialect;
23 const MemrefDialect = dialects.MemrefDialect;
24 const MemrefAddressSpace = dialects.AddressSpace;
25 const NvptxDialect = nvptx.NvptxDialect;
26 const Dimension = gpu.Dimension;
27 const Scope = gpu.Scope;
28
29 pub const gpu_to_nvptx_pass_name = "gpu-to-nvptx";
30 pub const gpu_to_nvptx_pass_description = "Lower GPU + memref ops to NVPTX dialect";
31
32 fn conversionPatternSpec(root_op_name: []const u8) rewrite.RewritePatternSpec {
33 return .{
34 .name = root_op_name,
35 .root_op_name = root_op_name,
36 };
37 }
38
39 pub const illegal_ops = [_][]const u8{
40 GpuDialect.ThreadIdxOp.operation_name,
41 GpuDialect.BlockIdxOp.operation_name,
42 GpuDialect.BlockDimOp.operation_name,
43 GpuDialect.GridDimOp.operation_name,
44 GpuDialect.GlobalIdxOp.operation_name,
45 GpuDialect.LaneIdOp.operation_name,
46 GpuDialect.WarpIdOp.operation_name,
47 GpuDialect.BarrierOp.operation_name,
48 GpuDialect.SyncWarpOp.operation_name,
49 GpuDialect.ActiveMaskOp.operation_name,
50 GpuDialect.AllSyncOp.operation_name,
51 GpuDialect.AnySyncOp.operation_name,
52 GpuDialect.BallotSyncOp.operation_name,
53 GpuDialect.ShflSyncOp.operation_name,
54 GpuDialect.WarpReduceOp.operation_name,
55 GpuDialect.WarpScanOp.operation_name,
56 GpuDialect.MmaSyncOp.operation_name,
57 GpuDialect.FenceOp.operation_name,
58 GpuDialect.CpAsyncSharedOp.operation_name,
59 GpuDialect.CpAsyncCommitOp.operation_name,
60 GpuDialect.CpAsyncWaitOp.operation_name,
61 MemrefDialect.LoadOp.operation_name,
62 MemrefDialect.StoreOp.operation_name,
63 MemrefDialect.AtomicRmwOp.operation_name,
64 MemrefDialect.AtomicCasOp.operation_name,
65 };
66
67 pub const legal_dialects = [_][]const u8{
68 dialects.BuiltinDialect.name,
69 dialects.FuncDialect.name,
70 dialects.ScfDialect.name,
71 ArithDialect.name,
72 GpuDialect.name,
73 MemrefDialect.name,
74 NvptxDialect.name,
75 };
76
77 pub const conversion_pattern_entries = [_]backends.ConversionPatternEntry{
78 .{ .spec = conversionPatternSpec(GpuDialect.ThreadIdxOp.operation_name), .rewrite = rewriteThreadIdx },
79 .{ .spec = conversionPatternSpec(GpuDialect.BlockIdxOp.operation_name), .rewrite = rewriteBlockIdx },
80 .{ .spec = conversionPatternSpec(GpuDialect.BlockDimOp.operation_name), .rewrite = rewriteBlockDim },
81 .{ .spec = conversionPatternSpec(GpuDialect.GridDimOp.operation_name), .rewrite = rewriteGridDim },
82 .{ .spec = conversionPatternSpec(GpuDialect.GlobalIdxOp.operation_name), .rewrite = rewriteGlobalIdx },
83 .{ .spec = conversionPatternSpec(GpuDialect.LaneIdOp.operation_name), .rewrite = rewriteLaneId },
84 .{ .spec = conversionPatternSpec(GpuDialect.WarpIdOp.operation_name), .rewrite = rewriteWarpId },
85 .{ .spec = conversionPatternSpec(GpuDialect.BarrierOp.operation_name), .rewrite = rewriteBarrier },
86 .{ .spec = conversionPatternSpec(GpuDialect.SyncWarpOp.operation_name), .rewrite = rewriteSyncWarp },
87 .{ .spec = conversionPatternSpec(GpuDialect.ActiveMaskOp.operation_name), .rewrite = rewriteActiveMask },
88 .{ .spec = conversionPatternSpec(GpuDialect.AllSyncOp.operation_name), .rewrite = rewriteAllSync },
89 .{ .spec = conversionPatternSpec(GpuDialect.AnySyncOp.operation_name), .rewrite = rewriteAnySync },
90 .{ .spec = conversionPatternSpec(GpuDialect.BallotSyncOp.operation_name), .rewrite = rewriteBallotSync },
91 .{ .spec = conversionPatternSpec(GpuDialect.ShflSyncOp.operation_name), .rewrite = rewriteShflSync },
92 .{ .spec = conversionPatternSpec(GpuDialect.WarpReduceOp.operation_name), .rewrite = rewriteWarpReduce },
93 .{ .spec = conversionPatternSpec(GpuDialect.WarpScanOp.operation_name), .rewrite = rewriteWarpScan },
94 .{ .spec = conversionPatternSpec(GpuDialect.MmaSyncOp.operation_name), .rewrite = rewriteMmaSync },
95 .{ .spec = conversionPatternSpec(GpuDialect.FenceOp.operation_name), .rewrite = rewriteFence },
96 .{ .spec = conversionPatternSpec(GpuDialect.CpAsyncSharedOp.operation_name), .rewrite = rewriteCpAsyncShared },
97 .{ .spec = conversionPatternSpec(GpuDialect.CpAsyncCommitOp.operation_name), .rewrite = rewriteCpAsyncCommit },
98 .{ .spec = conversionPatternSpec(GpuDialect.CpAsyncWaitOp.operation_name), .rewrite = rewriteCpAsyncWait },
99 .{ .spec = conversionPatternSpec(MemrefDialect.LoadOp.operation_name), .rewrite = rewriteMemrefLoad },
100 .{ .spec = conversionPatternSpec(MemrefDialect.StoreOp.operation_name), .rewrite = rewriteMemrefStore },
101 .{ .spec = conversionPatternSpec(MemrefDialect.AtomicRmwOp.operation_name), .rewrite = rewriteMemrefAtomicRmw },
102 .{ .spec = conversionPatternSpec(MemrefDialect.AtomicCasOp.operation_name), .rewrite = rewriteMemrefAtomicCas },
103 };
104
105 fn conversionPatternSpecs() [conversion_pattern_entries.len]rewrite.RewritePatternSpec {
106 comptime {
107 var specs: [conversion_pattern_entries.len]rewrite.RewritePatternSpec = undefined;
108 for (conversion_pattern_entries, 0..) |entry, index| {
109 specs[index] = entry.spec;
110 }
111 return specs;
112 }
113 }
114
115 pub const conversion_patterns = conversionPatternSpecs();
116
117 pub const target_spec = backends.TargetSpec{
118 .name = "nvptx",
119 .description = gpu_to_nvptx_pass_description,
120 .target_dialect_name = NvptxDialect.name,
121 .legality = .{
122 .legal_dialects = legal_dialects[0..],
123 .illegal_ops = illegal_ops[0..],
124 },
125 .conversion_patterns = conversion_patterns[0..],
126 .pass_name = gpu_to_nvptx_pass_name,
127 .pass_description = gpu_to_nvptx_pass_description,
128 };
129
130 pub fn createGpuToNvptxPass() Pass {
131 return .{
132 .name = gpu_to_nvptx_pass_name,
133 .description = gpu_to_nvptx_pass_description,
134 .run_fn = runGpuToNvptx,
135 .mutation_scope = .isolated,
136 .dependent_dialects = passes.dialectDependencies(&.{"nvptx"}),
137 };
138 }
139
140 pub const pass_registration = passes.PassRegistration{
141 .name = gpu_to_nvptx_pass_name,
142 .description = gpu_to_nvptx_pass_description,
143 .pass = createGpuToNvptxPass(),
144 };
145
146 fn runGpuToNvptx(ctx: *PassContext) PassResult {
147 ir.dialects.loadDialectSpec(ctx.ir_ctx, NvptxDialect.spec) catch return .failure;
148
149 var target = ConversionTarget.init(ctx.allocator);
150 defer target.deinit();
151 target_spec.applyLegality(&target) catch return .failure;
152
153 const had_illegal_ops = containsIllegalOps(ctx.op, &target);
154
155 var patterns = RewritePatternSet.init(ctx.allocator);
156 defer patterns.deinit();
157
158 for (conversion_pattern_entries) |entry| {
159 patterns.add(RewritePattern.init(entry.spec, entry.rewrite)) catch return .failure;
160 }
161
162 const result = passes.conversion.applyFullConversion(ctx.allocator, ctx.ir_ctx, ctx.op, &target, &patterns);
163 if (result == .failure) return .failure;
164
165 if (had_illegal_ops) ctx.markModified();
166 return .success;
167 }
168
169 fn containsIllegalOps(op: *ir.Operation, target: *const ConversionTarget) bool {
170 if (target.isIllegal(op)) return true;
171 for (op.regions.items) |*region| {
172 var block_iter = region.getBlocks();
173 while (block_iter.next()) |block| {
174 var current: ?*ir.Operation = @ptrCast(@alignCast(block.operations.head));
175 while (current) |child| {
176 if (containsIllegalOps(child, target)) return true;
177 current = child.next_op;
178 }
179 }
180 }
181 return false;
182 }
183
184 fn setDialectAttr(rewriter: *PatternRewriter, op: *ir.Operation, name: []const u8, dialect: []const u8, payload: []const u8) !void {
185 const attr = try rewriter.ir_ctx.getDialectAttr(dialect, payload);
186 try rewriter.setAttr(op, name, attr);
187 }
188
189 fn getDimensionAttr(op: *const ir.Operation) ?Dimension {
190 const dialect_attr = op.getAttrAs(ir.Attribute.DialectAttr, "dim") orelse return null;
191 return Dimension.fromString(dialect_attr.payload);
192 }
193
194 fn getScopeAttr(op: *const ir.Operation) ?Scope {
195 const dialect_attr = op.getAttrAs(ir.Attribute.DialectAttr, "scope") orelse return null;
196 return Scope.fromString(dialect_attr.payload);
197 }
198
199 fn rewriteThreadIdx(op: *ir.Operation, rewriter: *PatternRewriter) rewrite.PatternResult {
200 return rewriteGpuIndex(op, rewriter, NvptxDialect.ThreadIdxOp.operation_name);
201 }
202
203 fn rewriteBlockIdx(op: *ir.Operation, rewriter: *PatternRewriter) rewrite.PatternResult {
204 return rewriteGpuIndex(op, rewriter, NvptxDialect.BlockIdxOp.operation_name);
205 }
206
207 fn rewriteBlockDim(op: *ir.Operation, rewriter: *PatternRewriter) rewrite.PatternResult {
208 return rewriteGpuIndex(op, rewriter, NvptxDialect.BlockDimOp.operation_name);
209 }
210
211 fn rewriteGridDim(op: *ir.Operation, rewriter: *PatternRewriter) rewrite.PatternResult {
212 return rewriteGpuIndex(op, rewriter, NvptxDialect.GridDimOp.operation_name);
213 }
214
215 fn rewriteGpuIndex(
216 op: *ir.Operation,
217 rewriter: *PatternRewriter,
218 target_name: []const u8,
219 ) rewrite.PatternResult {
220 const dim = getDimensionAttr(op) orelse return .failure;
221 const result = op.getResult(0) orelse return .failure;
222
223 var state = ir.Operation.State.init(target_name, op.location);
224 state.addTypes(&.{result.type});
225 const dim_attr = rewriter.ir_ctx.getDialectAttr("nvptx.dim", dim.toString()) catch return .failure;
226 state.addAttributes(&.{.{ .name = "dim", .value = dim_attr }});
227 _ = rewriter.replaceOpWithNewOp(op, state) catch return .failure;
228 return .success;
229 }
230
231 fn rewriteGlobalIdx(op: *ir.Operation, rewriter: *PatternRewriter) rewrite.PatternResult {
232 const dim = getDimensionAttr(op) orelse return .failure;
233 const result = op.getResult(0) orelse return .failure;
234
235 var cta_state = ir.Operation.State.init(NvptxDialect.BlockIdxOp.operation_name, op.location);
236 cta_state.addTypes(&.{result.type});
237 const cta = rewriter.create(cta_state) catch return .failure;
238 setDialectAttr(rewriter, cta, "dim", "nvptx.dim", dim.toString()) catch return .failure;
239
240 var ntid_state = ir.Operation.State.init(NvptxDialect.BlockDimOp.operation_name, op.location);
241 ntid_state.addTypes(&.{result.type});
242 const ntid = rewriter.create(ntid_state) catch return .failure;
243 setDialectAttr(rewriter, ntid, "dim", "nvptx.dim", dim.toString()) catch return .failure;
244
245 var tid_state = ir.Operation.State.init(NvptxDialect.ThreadIdxOp.operation_name, op.location);
246 tid_state.addTypes(&.{result.type});
247 const tid = rewriter.create(tid_state) catch return .failure;
248 setDialectAttr(rewriter, tid, "dim", "nvptx.dim", dim.toString()) catch return .failure;
249
250 var mul_state = ir.Operation.State.init(ArithDialect.MulOp.operation_name, op.location);
251 mul_state.addOperands(&.{ cta.getResult(0).?, ntid.getResult(0).? });
252 mul_state.addTypes(&.{result.type});
253 const mul_op = rewriter.create(mul_state) catch return .failure;
254
255 var add_state = ir.Operation.State.init(ArithDialect.AddOp.operation_name, op.location);
256 add_state.addOperands(&.{ mul_op.getResult(0).?, tid.getResult(0).? });
257 add_state.addTypes(&.{result.type});
258 _ = rewriter.replaceOpWithNewOp(op, add_state) catch return .failure;
259 return .success;
260 }
261
262 fn rewriteLaneId(op: *ir.Operation, rewriter: *PatternRewriter) rewrite.PatternResult {
263 return rewriteGpuRegister(op, rewriter, NvptxDialect.LaneIdOp.operation_name);
264 }
265
266 fn rewriteWarpId(op: *ir.Operation, rewriter: *PatternRewriter) rewrite.PatternResult {
267 return rewriteGpuRegister(op, rewriter, NvptxDialect.WarpIdOp.operation_name);
268 }
269
270 fn rewriteGpuRegister(
271 op: *ir.Operation,
272 rewriter: *PatternRewriter,
273 target_name: []const u8,
274 ) rewrite.PatternResult {
275 const result = op.getResult(0) orelse return .failure;
276 var state = ir.Operation.State.init(target_name, op.location);
277 state.addTypes(&.{result.type});
278 _ = rewriter.replaceOpWithNewOp(op, state) catch return .failure;
279 return .success;
280 }
281
282 fn rewriteBarrier(op: *ir.Operation, rewriter: *PatternRewriter) rewrite.PatternResult {
283 const scope = getScopeAttr(op) orelse return .failure;
284 const target_name = switch (scope) {
285 .thread => {
286 rewriter.eraseOp(op) catch return .failure;
287 return .success;
288 },
289 .warp => NvptxDialect.WarpBarrierAllOp.operation_name,
290 .block => NvptxDialect.Barrier0Op.operation_name,
291 .cluster, .device, .system => return .failure,
292 };
293 const state = ir.Operation.State.init(target_name, op.location);
294 _ = rewriter.replaceOpWithNewOp(op, state) catch return .failure;
295 return .success;
296 }
297
298 fn rewriteSyncWarp(op: *ir.Operation, rewriter: *PatternRewriter) rewrite.PatternResult {
299 const sync = GpuDialect.SyncWarpOp{ .op = op };
300 var state = ir.Operation.State.init(NvptxDialect.SyncWarpOp.operation_name, op.location);
301 state.addOperands(&.{sync.getMask()});
302 _ = rewriter.replaceOpWithNewOp(op, state) catch return .failure;
303 return .success;
304 }
305
306 fn rewriteActiveMask(op: *ir.Operation, rewriter: *PatternRewriter) rewrite.PatternResult {
307 const result = op.getResult(0) orelse return .failure;
308 var state = ir.Operation.State.init(NvptxDialect.ActiveMaskOp.operation_name, op.location);
309 state.addTypes(&.{result.type});
310 _ = rewriter.replaceOpWithNewOp(op, state) catch return .failure;
311 return .success;
312 }
313
314 fn rewriteAllSync(op: *ir.Operation, rewriter: *PatternRewriter) rewrite.PatternResult {
315 const all = GpuDialect.AllSyncOp{ .op = op };
316 const result = op.getResult(0) orelse return .failure;
317 var state = ir.Operation.State.init(NvptxDialect.AllSyncOp.operation_name, op.location);
318 state.addOperands(&.{ all.getMask(), all.getPredicate() });
319 state.addTypes(&.{result.type});
320 _ = rewriter.replaceOpWithNewOp(op, state) catch return .failure;
321 return .success;
322 }
323
324 fn rewriteAnySync(op: *ir.Operation, rewriter: *PatternRewriter) rewrite.PatternResult {
325 const any = GpuDialect.AnySyncOp{ .op = op };
326 const result = op.getResult(0) orelse return .failure;
327 var state = ir.Operation.State.init(NvptxDialect.AnySyncOp.operation_name, op.location);
328 state.addOperands(&.{ any.getMask(), any.getPredicate() });
329 state.addTypes(&.{result.type});
330 _ = rewriter.replaceOpWithNewOp(op, state) catch return .failure;
331 return .success;
332 }
333
334 fn rewriteBallotSync(op: *ir.Operation, rewriter: *PatternRewriter) rewrite.PatternResult {
335 const ballot = GpuDialect.BallotSyncOp{ .op = op };
336 const result = op.getResult(0) orelse return .failure;
337 var state = ir.Operation.State.init(NvptxDialect.BallotSyncOp.operation_name, op.location);
338 state.addOperands(&.{ ballot.getMask(), ballot.getPredicate() });
339 state.addTypes(&.{result.type});
340 _ = rewriter.replaceOpWithNewOp(op, state) catch return .failure;
341 return .success;
342 }
343
344 fn rewriteShflSync(op: *ir.Operation, rewriter: *PatternRewriter) rewrite.PatternResult {
345 const shfl = GpuDialect.ShflSyncOp{ .op = op };
346 const mode = shfl.getMode() orelse return .failure;
347 const result = op.getResult(0) orelse return .failure;
348 var state = ir.Operation.State.init(NvptxDialect.ShflSyncOp.operation_name, op.location);
349 state.addOperands(&.{ shfl.getMask(), shfl.getSrc(), shfl.getLaneOrDelta() });
350 state.addTypes(&.{result.type});
351 const mode_attr = rewriter.ir_ctx.getDialectAttr("nvptx.shuffle_mode", mode.toString()) catch return .failure;
352 state.addAttributes(&.{.{ .name = "mode", .value = mode_attr }});
353 _ = rewriter.replaceOpWithNewOp(op, state) catch return .failure;
354 return .success;
355 }
356
357 fn rewriteWarpReduce(op: *ir.Operation, rewriter: *PatternRewriter) rewrite.PatternResult {
358 const reduce = GpuDialect.WarpReduceOp{ .op = op };
359 const kind = reduce.getOpKind() orelse return .failure;
360 const result = op.getResult(0) orelse return .failure;
361 var state = ir.Operation.State.init(NvptxDialect.WarpReduceOp.operation_name, op.location);
362 state.addOperands(&.{ reduce.getMask(), reduce.getValue() });
363 state.addTypes(&.{result.type});
364 const op_attr = rewriter.ir_ctx.getDialectAttr("nvptx.warp_op", kind.toString()) catch return .failure;
365 state.addAttributes(&.{.{ .name = "op", .value = op_attr }});
366 _ = rewriter.replaceOpWithNewOp(op, state) catch return .failure;
367 return .success;
368 }
369
370 fn rewriteWarpScan(op: *ir.Operation, rewriter: *PatternRewriter) rewrite.PatternResult {
371 const scan = GpuDialect.WarpScanOp{ .op = op };
372 const kind = scan.getOpKind() orelse return .failure;
373 const inclusive = scan.isInclusive();
374 const result = op.getResult(0) orelse return .failure;
375 var state = ir.Operation.State.init(NvptxDialect.WarpScanOp.operation_name, op.location);
376 state.addOperands(&.{ scan.getMask(), scan.getValue() });
377 state.addTypes(&.{result.type});
378 const op_attr = rewriter.ir_ctx.getDialectAttr("nvptx.warp_op", kind.toString()) catch return .failure;
379 const inclusive_attr = rewriter.ir_ctx.getBoolAttr(inclusive) catch return .failure;
380 state.addAttributes(&.{
381 .{ .name = "op", .value = op_attr },
382 .{ .name = "inclusive", .value = inclusive_attr },
383 });
384 _ = rewriter.replaceOpWithNewOp(op, state) catch return .failure;
385 return .success;
386 }
387
388 fn rewriteMmaSync(op: *ir.Operation, rewriter: *PatternRewriter) rewrite.PatternResult {
389 const mma = GpuDialect.MmaSyncOp{ .op = op };
390 const shape = mma.getShape() orelse return .failure;
391 if (op.operands.items.len != 10 or op.getNumResults() != 4) return .failure;
392 var state = ir.Operation.State.init(NvptxDialect.MmaSyncOp.operation_name, op.location);
393 state.addOperands(&.{
394 mma.getA(0), mma.getA(1), mma.getA(2), mma.getA(3),
395 mma.getB(0), mma.getB(1), mma.getC(0), mma.getC(1),
396 mma.getC(2), mma.getC(3),
397 });
398 state.addTypes(&.{
399 op.getResult(0).?.type,
400 op.getResult(1).?.type,
401 op.getResult(2).?.type,
402 op.getResult(3).?.type,
403 });
404 var shape_buf: [32]u8 = undefined;
405 const shape_str = shape.toString(shape_buf[0..]) catch return .failure;
406 const shape_attr = rewriter.ir_ctx.getDialectAttr("nvptx.mma_shape", shape_str) catch return .failure;
407 state.addAttributes(&.{.{ .name = "shape", .value = shape_attr }});
408 _ = rewriter.replaceOpWithNewOp(op, state) catch return .failure;
409 return .success;
410 }
411
412 fn rewriteFence(op: *ir.Operation, rewriter: *PatternRewriter) rewrite.PatternResult {
413 const state = ir.Operation.State.init(NvptxDialect.FenceDeviceOp.operation_name, op.location);
414 _ = rewriter.replaceOpWithNewOp(op, state) catch return .failure;
415 return .success;
416 }
417
418 fn rewriteCpAsyncShared(op: *ir.Operation, rewriter: *PatternRewriter) rewrite.PatternResult {
419 const copy = GpuDialect.CpAsyncSharedOp{ .op = op };
420 const bytes = copy.getBytes() orelse return .failure;
421 var state = ir.Operation.State.init(NvptxDialect.CpAsyncSharedOp.operation_name, op.location);
422 state.addOperands(&.{ copy.getDst(), copy.getDstIndex(), copy.getSrc(), copy.getSrcIndex() });
423 const bytes_attr = rewriter.ir_ctx.getI64Attr(@intCast(bytes)) catch return .failure;
424 state.addAttributes(&.{.{ .name = "bytes", .value = bytes_attr }});
425 _ = rewriter.replaceOpWithNewOp(op, state) catch return .failure;
426 return .success;
427 }
428
429 fn rewriteCpAsyncCommit(op: *ir.Operation, rewriter: *PatternRewriter) rewrite.PatternResult {
430 const state = ir.Operation.State.init(NvptxDialect.CpAsyncCommitOp.operation_name, op.location);
431 _ = rewriter.replaceOpWithNewOp(op, state) catch return .failure;
432 return .success;
433 }
434
435 fn rewriteCpAsyncWait(op: *ir.Operation, rewriter: *PatternRewriter) rewrite.PatternResult {
436 const wait = GpuDialect.CpAsyncWaitOp{ .op = op };
437 const groups = wait.getGroups() orelse return .failure;
438 var state = ir.Operation.State.init(NvptxDialect.CpAsyncWaitOp.operation_name, op.location);
439 const groups_attr = rewriter.ir_ctx.getI64Attr(@intCast(groups)) catch return .failure;
440 state.addAttributes(&.{.{ .name = "groups", .value = groups_attr }});
441 _ = rewriter.replaceOpWithNewOp(op, state) catch return .failure;
442 return .success;
443 }
444
445 fn rewriteMemrefLoad(op: *ir.Operation, rewriter: *PatternRewriter) rewrite.PatternResult {
446 const load = MemrefDialect.LoadOp{ .op = op };
447 const result = op.getResult(0) orelse return .failure;
448 const memref = load.getMemref();
449 const addr_space = memrefAddressSpace(memref.type) orelse return .failure;
450 const op_name = switch (addr_space) {
451 .local => NvptxDialect.LoadLocalOp.operation_name,
452 .shared => NvptxDialect.LoadSharedOp.operation_name,
453 else => NvptxDialect.LoadGlobalOp.operation_name,
454 };
455
456 var state = ir.Operation.State.init(op_name, op.location);
457 state.addOperands(&.{ memref, load.getIndex() });
458 state.addTypes(&.{result.type});
459 _ = rewriter.replaceOpWithNewOp(op, state) catch return .failure;
460 return .success;
461 }
462
463 fn rewriteMemrefStore(op: *ir.Operation, rewriter: *PatternRewriter) rewrite.PatternResult {
464 const store = MemrefDialect.StoreOp{ .op = op };
465 const memref = store.getMemref();
466 const addr_space = memrefAddressSpace(memref.type) orelse return .failure;
467 const op_name = switch (addr_space) {
468 .local => NvptxDialect.StoreLocalOp.operation_name,
469 .shared => NvptxDialect.StoreSharedOp.operation_name,
470 else => NvptxDialect.StoreGlobalOp.operation_name,
471 };
472
473 var state = ir.Operation.State.init(op_name, op.location);
474 state.addOperands(&.{ store.getValue(), memref, store.getIndex() });
475 _ = rewriter.replaceOpWithNewOp(op, state) catch return .failure;
476 return .success;
477 }
478
479 fn rewriteMemrefAtomicRmw(op: *ir.Operation, rewriter: *PatternRewriter) rewrite.PatternResult {
480 const atomic = MemrefDialect.AtomicRmwOp{ .op = op };
481 const result = op.getResult(0) orelse return .failure;
482 const kind = atomic.getKind() orelse return .failure;
483 const memref = atomic.getMemref();
484 const addr_space = memrefAddressSpace(memref.type) orelse return .failure;
485 const op_name = switch (addr_space) {
486 .shared => NvptxDialect.AtomicSharedOp.operation_name,
487 .local => return .failure,
488 else => NvptxDialect.AtomicGlobalOp.operation_name,
489 };
490
491 var state = ir.Operation.State.init(op_name, op.location);
492 state.addOperands(&.{ atomic.getValue(), memref, atomic.getIndex() });
493 state.addTypes(&.{result.type});
494 const kind_attr = rewriter.ir_ctx.getDialectAttr("nvptx.atomic_kind", kind.toString()) catch return .failure;
495 state.addAttributes(&.{.{ .name = "kind", .value = kind_attr }});
496 _ = rewriter.replaceOpWithNewOp(op, state) catch return .failure;
497 return .success;
498 }
499
500 fn rewriteMemrefAtomicCas(op: *ir.Operation, rewriter: *PatternRewriter) rewrite.PatternResult {
501 const atomic = MemrefDialect.AtomicCasOp{ .op = op };
502 const result = op.getResult(0) orelse return .failure;
503 const memref = atomic.getMemref();
504 const addr_space = memrefAddressSpace(memref.type) orelse return .failure;
505 const op_name = switch (addr_space) {
506 .shared => NvptxDialect.AtomicCasSharedOp.operation_name,
507 .local => return .failure,
508 else => NvptxDialect.AtomicCasGlobalOp.operation_name,
509 };
510
511 var state = ir.Operation.State.init(op_name, op.location);
512 state.addOperands(&.{ atomic.getExpected(), atomic.getDesired(), memref, atomic.getIndex() });
513 state.addTypes(&.{result.type});
514 _ = rewriter.replaceOpWithNewOp(op, state) catch return .failure;
515 return .success;
516 }
517
518 fn memrefAddressSpace(typ: ir.Type) ?MemrefAddressSpace {
519 const name = typ.getDialectTypeName() orelse return null;
520 if (!std.mem.eql(u8, name, MemrefDialect.name)) return null;
521 const param_key = typ.getDialectParamKey() orelse return null;
522 const params = MemrefDialect.parseMemrefParams(param_key) orelse return null;
523 return params.addr_space;
524 }
525
526 test "nvptx target spec publishes legality and conversion roots" {
527 const testing = std.testing;
528
529 try testing.expectEqualStrings("nvptx", target_spec.name);
530 try testing.expectEqualStrings(NvptxDialect.name, target_spec.target_dialect_name);
531 try testing.expectEqualStrings(gpu_to_nvptx_pass_name, target_spec.pass_name);
532 try testing.expectEqual(illegal_ops.len, target_spec.conversionPatternRootCount());
533 try testing.expect(target_spec.legalizesDialect(NvptxDialect.name));
534 try testing.expect(target_spec.marksIllegalOp(GpuDialect.ThreadIdxOp.operation_name));
535 try testing.expect(target_spec.marksIllegalOp(MemrefDialect.LoadOp.operation_name));
536 try testing.expect(target_spec.hasConversionPatternRoot(GpuDialect.GlobalIdxOp.operation_name));
537 try testing.expect(target_spec.hasConversionPatternRoot(MemrefDialect.AtomicCasOp.operation_name));
538
539 var target = ConversionTarget.init(testing.allocator);
540 defer target.deinit();
541 try target_spec.applyLegality(&target);
542 try testing.expect(target.legal_dialects.contains(NvptxDialect.name));
543 try testing.expect(target.illegal_ops.contains(MemrefDialect.LoadOp.operation_name));
544 }
545
546 test "nvptx conversion lowers gpu idx and memref ops" {
547 const testing = std.testing;
548 var arena = std.heap.ArenaAllocator.init(std.testing.allocator);
549 defer arena.deinit();
550 const allocator = arena.allocator();
551
552 var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
553 defer ctx.deinit(allocator);
554
555 const loc = ir.Location.getUnknown();
556 const module = try dialects.BuiltinDialect.ModuleOp.create(&ctx, loc);
557 const module_block = module.getBodyBlock();
558
559 const i32_type = try ArithDialect.getI32Type(&ctx);
560 const memref_type = try MemrefDialect.getMemrefType1D(&ctx, 4, i32_type, .device);
561
562 var func = try dialects.FuncDialect.FuncOp.createKernel(&ctx, loc, "kernel", &.{memref_type});
563 try module_block.addOperation(func.op);
564
565 const entry = func.getEntryBlock();
566 const arg0 = entry.arguments.items[0];
567
568 const gid = try GpuDialect.GlobalIdxOp.create(&ctx, loc, .x);
569 try entry.addOperation(gid.op);
570
571 const load = try MemrefDialect.LoadOp.create(&ctx, loc, arg0, gid.getResult(), i32_type);
572 try entry.addOperation(load.op);
573
574 const one = try ArithDialect.ConstantOp.createInt(&ctx, loc, i32_type, 1);
575 try entry.addOperation(one.op);
576 const add = try ArithDialect.AddOp.create(&ctx, loc, load.getResult(), one.getResult());
577 try entry.addOperation(add.op);
578
579 const store = try MemrefDialect.StoreOp.create(&ctx, loc, add.getResult(), arg0, gid.getResult());
580 try entry.addOperation(store.op);
581
582 const lane = try GpuDialect.LaneIdOp.create(&ctx, loc);
583 try entry.addOperation(lane.op);
584
585 const warp = try GpuDialect.WarpIdOp.create(&ctx, loc);
586 try entry.addOperation(warp.op);
587
588 const active = try GpuDialect.ActiveMaskOp.create(&ctx, loc);
589 try entry.addOperation(active.op);
590
591 const mask = try ArithDialect.ConstantOp.createInt(&ctx, loc, i32_type, -1);
592 try entry.addOperation(mask.op);
593
594 const predicate = try ArithDialect.ConstantOp.createBool(&ctx, loc, true);
595 try entry.addOperation(predicate.op);
596
597 const sync = try GpuDialect.SyncWarpOp.create(&ctx, loc, mask.getResult());
598 try entry.addOperation(sync.op);
599
600 const all = try GpuDialect.AllSyncOp.create(&ctx, loc, mask.getResult(), predicate.getResult());
601 try entry.addOperation(all.op);
602
603 const any = try GpuDialect.AnySyncOp.create(&ctx, loc, mask.getResult(), predicate.getResult());
604 try entry.addOperation(any.op);
605
606 const ballot = try GpuDialect.BallotSyncOp.create(&ctx, loc, mask.getResult(), predicate.getResult());
607 try entry.addOperation(ballot.op);
608
609 const shuffle = try GpuDialect.ShflSyncOp.create(&ctx, loc, .xor, mask.getResult(), lane.getResult(), one.getResult());
610 try entry.addOperation(shuffle.op);
611
612 const reduce = try GpuDialect.WarpReduceOp.create(&ctx, loc, .add, mask.getResult(), add.getResult());
613 try entry.addOperation(reduce.op);
614
615 const scan = try GpuDialect.WarpScanOp.create(&ctx, loc, .xor, true, mask.getResult(), add.getResult());
616 try entry.addOperation(scan.op);
617
618 const warp_barrier = try GpuDialect.BarrierOp.create(&ctx, loc, .warp);
619 try entry.addOperation(warp_barrier.op);
620
621 const block_barrier = try GpuDialect.BarrierOp.create(&ctx, loc, .block);
622 try entry.addOperation(block_barrier.op);
623
624 const ret = try dialects.FuncDialect.ReturnOp.create(&ctx, loc, &.{});
625 try entry.addOperation(ret.op);
626
627 var analysis_cache = passes.AnalysisCache.init(allocator, null);
628 defer analysis_cache.deinit();
629 var pass_ctx = PassContext.init(module.op, &ctx, allocator, &analysis_cache);
630 defer pass_ctx.deinit();
631
632 const pass = createGpuToNvptxPass();
633 try testing.expectEqual(PassResult.success, pass.run(&pass_ctx));
634
635 var saw_nvptx = false;
636 var saw_barrier0 = false;
637 var saw_warp_barrier = false;
638 var saw_warp_control = false;
639 var saw_vote = false;
640 var saw_shuffle = false;
641 var saw_collective = false;
642 var iter = entry.operations.head;
643 while (iter) |op_ptr| {
644 const op: *ir.Operation = @ptrCast(@alignCast(op_ptr));
645 if (std.mem.eql(u8, op.name.name, NvptxDialect.ThreadIdxOp.operation_name) or
646 std.mem.eql(u8, op.name.name, NvptxDialect.BlockIdxOp.operation_name) or
647 std.mem.eql(u8, op.name.name, NvptxDialect.LoadGlobalOp.operation_name))
648 {
649 saw_nvptx = true;
650 }
651 if (std.mem.eql(u8, op.name.name, NvptxDialect.Barrier0Op.operation_name)) {
652 saw_barrier0 = true;
653 }
654 if (std.mem.eql(u8, op.name.name, NvptxDialect.WarpBarrierAllOp.operation_name)) {
655 saw_warp_barrier = true;
656 }
657 if (std.mem.eql(u8, op.name.name, NvptxDialect.LaneIdOp.operation_name) or
658 std.mem.eql(u8, op.name.name, NvptxDialect.WarpIdOp.operation_name) or
659 std.mem.eql(u8, op.name.name, NvptxDialect.ActiveMaskOp.operation_name) or
660 std.mem.eql(u8, op.name.name, NvptxDialect.SyncWarpOp.operation_name))
661 {
662 saw_warp_control = true;
663 }
664 if (std.mem.eql(u8, op.name.name, NvptxDialect.AllSyncOp.operation_name) or
665 std.mem.eql(u8, op.name.name, NvptxDialect.AnySyncOp.operation_name) or
666 std.mem.eql(u8, op.name.name, NvptxDialect.BallotSyncOp.operation_name))
667 {
668 saw_vote = true;
669 }
670 if (std.mem.eql(u8, op.name.name, NvptxDialect.ShflSyncOp.operation_name)) {
671 saw_shuffle = true;
672 }
673 if (std.mem.eql(u8, op.name.name, NvptxDialect.WarpReduceOp.operation_name) or
674 std.mem.eql(u8, op.name.name, NvptxDialect.WarpScanOp.operation_name))
675 {
676 saw_collective = true;
677 }
678 iter = op.next_op;
679 }
680
681 try testing.expect(saw_nvptx);
682 try testing.expect(saw_barrier0);
683 try testing.expect(saw_warp_barrier);
684 try testing.expect(saw_warp_control);
685 try testing.expect(saw_vote);
686 try testing.expect(saw_shuffle);
687 try testing.expect(saw_collective);
688 }
689
690 test "nvptx conversion rejects unknown target ops" {
691 const testing = std.testing;
692 var arena = std.heap.ArenaAllocator.init(std.testing.allocator);
693 defer arena.deinit();
694 const allocator = arena.allocator();
695
696 var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
697 defer ctx.deinit(allocator);
698 try ctx.allowUnregistered();
699
700 const loc = ir.Location.getUnknown();
701 const module = try dialects.BuiltinDialect.ModuleOp.create(&ctx, loc);
702 const unknown = try ctx.createOperation(ir.Operation.State.init("external.unknown", loc));
703 try module.getBodyBlock().addOperation(unknown);
704
705 var analysis_cache = passes.AnalysisCache.init(allocator, null);
706 defer analysis_cache.deinit();
707 var pass_ctx = PassContext.init(module.op, &ctx, allocator, &analysis_cache);
708 defer pass_ctx.deinit();
709
710 const pass = createGpuToNvptxPass();
711 try testing.expectEqual(PassResult.failure, pass.run(&pass_ctx));
712 }