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 }