lib/choir/src/backends/gpu/spirv/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 spirv = @import("dialect.zig");
 11 const spirv_emit = @import("emitter/root.zig").emit;
 12 
 13 const Pass = passes.Pass;
 14 const PassContext = passes.PassContext;
 15 const PassResult = passes.PassResult;
 16 
 17 const ConversionTarget = passes.ConversionTarget;
 18 const RewritePatternSet = rewrite.RewritePatternSet;
 19 const RewritePattern = rewrite.RewritePattern;
 20 const PatternRewriter = rewrite.PatternRewriter;
 21 
 22 const ArithDialect = dialects.ArithDialect;
 23 const GpuDialect = gpu.GpuDialect;
 24 const MemrefDialect = dialects.MemrefDialect;
 25 const SpirvDialect = spirv.SpirvDialect;
 26 const Dimension = gpu.Dimension;
 27 const Scope = gpu.Scope;
 28 
 29 pub const gpu_to_spirv_pass_name = "gpu-to-spirv";
 30 pub const gpu_to_spirv_pass_description = "Lower GPU + arith ops to SPIR-V 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.BarrierOp.operation_name,
 46     GpuDialect.SyncWarpOp.operation_name,
 47     GpuDialect.ActiveMaskOp.operation_name,
 48     GpuDialect.AllSyncOp.operation_name,
 49     GpuDialect.AnySyncOp.operation_name,
 50     GpuDialect.BallotSyncOp.operation_name,
 51     GpuDialect.ShflSyncOp.operation_name,
 52     GpuDialect.WarpReduceOp.operation_name,
 53     GpuDialect.WarpScanOp.operation_name,
 54     ArithDialect.ConstantOp.operation_name,
 55     ArithDialect.AddOp.operation_name,
 56     ArithDialect.SubOp.operation_name,
 57     ArithDialect.MulOp.operation_name,
 58     ArithDialect.DivOp.operation_name,
 59 };
 60 
 61 pub const legal_dialects = [_][]const u8{
 62     dialects.BuiltinDialect.name,
 63     dialects.FuncDialect.name,
 64     dialects.ScfDialect.name,
 65     ArithDialect.name,
 66     GpuDialect.name,
 67     MemrefDialect.name,
 68     SpirvDialect.name,
 69 };
 70 
 71 pub const conversion_pattern_entries = [_]backends.ConversionPatternEntry{
 72     .{ .spec = conversionPatternSpec(GpuDialect.ThreadIdxOp.operation_name), .rewrite = rewriteThreadIdx },
 73     .{ .spec = conversionPatternSpec(GpuDialect.BlockIdxOp.operation_name), .rewrite = rewriteBlockIdx },
 74     .{ .spec = conversionPatternSpec(GpuDialect.BlockDimOp.operation_name), .rewrite = rewriteBlockDim },
 75     .{ .spec = conversionPatternSpec(GpuDialect.GridDimOp.operation_name), .rewrite = rewriteGridDim },
 76     .{ .spec = conversionPatternSpec(GpuDialect.GlobalIdxOp.operation_name), .rewrite = rewriteGlobalIdx },
 77     .{ .spec = conversionPatternSpec(GpuDialect.BarrierOp.operation_name), .rewrite = rewriteBarrier },
 78     .{ .spec = conversionPatternSpec(GpuDialect.SyncWarpOp.operation_name), .rewrite = rewriteSyncWarp },
 79     .{ .spec = conversionPatternSpec(GpuDialect.ActiveMaskOp.operation_name), .rewrite = rewriteActiveMask },
 80     .{ .spec = conversionPatternSpec(GpuDialect.AllSyncOp.operation_name), .rewrite = rewriteAllSync },
 81     .{ .spec = conversionPatternSpec(GpuDialect.AnySyncOp.operation_name), .rewrite = rewriteAnySync },
 82     .{ .spec = conversionPatternSpec(GpuDialect.BallotSyncOp.operation_name), .rewrite = rewriteBallotSync },
 83     .{ .spec = conversionPatternSpec(GpuDialect.ShflSyncOp.operation_name), .rewrite = rewriteShflSync },
 84     .{ .spec = conversionPatternSpec(GpuDialect.WarpReduceOp.operation_name), .rewrite = rewriteWarpReduce },
 85     .{ .spec = conversionPatternSpec(GpuDialect.WarpScanOp.operation_name), .rewrite = rewriteWarpScan },
 86     .{ .spec = conversionPatternSpec(ArithDialect.ConstantOp.operation_name), .rewrite = rewriteArithConstant },
 87     .{ .spec = conversionPatternSpec(ArithDialect.AddOp.operation_name), .rewrite = rewriteArithAdd },
 88     .{ .spec = conversionPatternSpec(ArithDialect.SubOp.operation_name), .rewrite = rewriteArithSub },
 89     .{ .spec = conversionPatternSpec(ArithDialect.MulOp.operation_name), .rewrite = rewriteArithMul },
 90     .{ .spec = conversionPatternSpec(ArithDialect.DivOp.operation_name), .rewrite = rewriteArithDiv },
 91 };
 92 
 93 fn conversionPatternSpecs() [conversion_pattern_entries.len]rewrite.RewritePatternSpec {
 94     comptime {
 95         var specs: [conversion_pattern_entries.len]rewrite.RewritePatternSpec = undefined;
 96         for (conversion_pattern_entries, 0..) |entry, index| {
 97             specs[index] = entry.spec;
 98         }
 99         return specs;
100     }
101 }
102 
103 pub const conversion_patterns = conversionPatternSpecs();
104 
105 pub const target_spec = backends.TargetSpec{
106     .name = "spirv",
107     .description = gpu_to_spirv_pass_description,
108     .target_dialect_name = SpirvDialect.name,
109     .legality = .{
110         .legal_dialects = legal_dialects[0..],
111         .illegal_ops = illegal_ops[0..],
112     },
113     .conversion_patterns = conversion_patterns[0..],
114     .pass_name = gpu_to_spirv_pass_name,
115     .pass_description = gpu_to_spirv_pass_description,
116 };
117 
118 pub fn createGpuToSpirvPass() Pass {
119     return .{
120         .name = gpu_to_spirv_pass_name,
121         .description = gpu_to_spirv_pass_description,
122         .run_fn = runGpuToSpirv,
123         .mutation_scope = .isolated,
124         .dependent_dialects = passes.dialectDependencies(&.{"spirv"}),
125     };
126 }
127 
128 pub const pass_registration = passes.PassRegistration{
129     .name = gpu_to_spirv_pass_name,
130     .description = gpu_to_spirv_pass_description,
131     .pass = createGpuToSpirvPass(),
132 };
133 
134 fn runGpuToSpirv(ctx: *PassContext) PassResult {
135     ir.dialects.loadDialectSpec(ctx.ir_ctx, SpirvDialect.spec) catch return .failure;
136 
137     var target = ConversionTarget.init(ctx.allocator);
138     defer target.deinit();
139     target_spec.applyLegality(&target) catch return .failure;
140 
141     const had_illegal_ops = containsIllegalOps(ctx.op, &target);
142 
143     var patterns = RewritePatternSet.init(ctx.allocator);
144     defer patterns.deinit();
145 
146     for (conversion_pattern_entries) |entry| {
147         patterns.add(RewritePattern.init(entry.spec, entry.rewrite)) catch return .failure;
148     }
149 
150     const result = passes.conversion.applyFullConversion(ctx.allocator, ctx.ir_ctx, ctx.op, &target, &patterns);
151     if (result == .failure) return .failure;
152 
153     if (had_illegal_ops) ctx.markModified();
154     return .success;
155 }
156 
157 fn containsIllegalOps(op: *ir.Operation, target: *const ConversionTarget) bool {
158     if (target.isIllegal(op)) return true;
159     for (op.regions.items) |*region| {
160         var block_iter = region.getBlocks();
161         while (block_iter.next()) |block| {
162             var current: ?*ir.Operation = @ptrCast(@alignCast(block.operations.head));
163             while (current) |child| {
164                 if (containsIllegalOps(child, target)) return true;
165                 current = child.next_op;
166             }
167         }
168     }
169     return false;
170 }
171 
172 fn getDimensionAttr(op: *const ir.Operation) ?Dimension {
173     const dialect_attr = op.getAttrAs(ir.Attribute.DialectAttr, "dim") orelse return null;
174     return Dimension.fromString(dialect_attr.payload);
175 }
176 
177 fn getScopeAttr(op: *const ir.Operation) ?Scope {
178     const dialect_attr = op.getAttrAs(ir.Attribute.DialectAttr, "scope") orelse return null;
179     return Scope.fromString(dialect_attr.payload);
180 }
181 
182 fn rewriteThreadIdx(op: *ir.Operation, rewriter: *PatternRewriter) rewrite.PatternResult {
183     return rewriteGpuIndex(op, rewriter, SpirvDialect.LocalInvocationIdOp.operation_name);
184 }
185 
186 fn rewriteBlockIdx(op: *ir.Operation, rewriter: *PatternRewriter) rewrite.PatternResult {
187     return rewriteGpuIndex(op, rewriter, SpirvDialect.WorkgroupIdOp.operation_name);
188 }
189 
190 fn rewriteBlockDim(op: *ir.Operation, rewriter: *PatternRewriter) rewrite.PatternResult {
191     return rewriteGpuIndex(op, rewriter, SpirvDialect.WorkgroupSizeOp.operation_name);
192 }
193 
194 fn rewriteGridDim(op: *ir.Operation, rewriter: *PatternRewriter) rewrite.PatternResult {
195     return rewriteGpuIndex(op, rewriter, SpirvDialect.NumWorkgroupsOp.operation_name);
196 }
197 
198 fn rewriteGlobalIdx(op: *ir.Operation, rewriter: *PatternRewriter) rewrite.PatternResult {
199     return rewriteGpuIndex(op, rewriter, SpirvDialect.GlobalInvocationIdOp.operation_name);
200 }
201 
202 fn rewriteGpuIndex(
203     op: *ir.Operation,
204     rewriter: *PatternRewriter,
205     target_name: []const u8,
206 ) rewrite.PatternResult {
207     const dim = getDimensionAttr(op) orelse return .failure;
208     const result = op.getResult(0) orelse return .failure;
209 
210     var state = ir.Operation.State.init(target_name, op.location);
211     state.addTypes(&.{result.type});
212     const dim_attr = rewriter.ir_ctx.getDialectAttr("spirv.dim", dim.toString()) catch return .failure;
213     state.addAttributes(&.{.{ .name = "dim", .value = dim_attr }});
214     _ = rewriter.replaceOpWithNewOp(op, state) catch return .failure;
215     return .success;
216 }
217 
218 fn rewriteBarrier(op: *ir.Operation, rewriter: *PatternRewriter) rewrite.PatternResult {
219     const scope = getScopeAttr(op) orelse return .failure;
220     var state = ir.Operation.State.init(SpirvDialect.BarrierOp.operation_name, op.location);
221     const scope_attr = rewriter.ir_ctx.getDialectAttr("spirv.scope", scope.toString()) catch return .failure;
222     state.addAttributes(&.{.{ .name = "scope", .value = scope_attr }});
223     _ = rewriter.replaceOpWithNewOp(op, state) catch return .failure;
224     return .success;
225 }
226 
227 fn rewriteSyncWarp(op: *ir.Operation, rewriter: *PatternRewriter) rewrite.PatternResult {
228     const sync = GpuDialect.SyncWarpOp{ .op = op };
229     var state = ir.Operation.State.init(SpirvDialect.SyncWarpOp.operation_name, op.location);
230     state.addOperands(&.{sync.getMask()});
231     _ = rewriter.replaceOpWithNewOp(op, state) catch return .failure;
232     return .success;
233 }
234 
235 fn rewriteActiveMask(op: *ir.Operation, rewriter: *PatternRewriter) rewrite.PatternResult {
236     const result = op.getResult(0) orelse return .failure;
237     var state = ir.Operation.State.init(SpirvDialect.ActiveMaskOp.operation_name, op.location);
238     state.addTypes(&.{result.type});
239     _ = rewriter.replaceOpWithNewOp(op, state) catch return .failure;
240     return .success;
241 }
242 
243 fn rewriteAllSync(op: *ir.Operation, rewriter: *PatternRewriter) rewrite.PatternResult {
244     const all = GpuDialect.AllSyncOp{ .op = op };
245     const result = op.getResult(0) orelse return .failure;
246     var state = ir.Operation.State.init(SpirvDialect.AllSyncOp.operation_name, op.location);
247     state.addOperands(&.{ all.getMask(), all.getPredicate() });
248     state.addTypes(&.{result.type});
249     _ = rewriter.replaceOpWithNewOp(op, state) catch return .failure;
250     return .success;
251 }
252 
253 fn rewriteAnySync(op: *ir.Operation, rewriter: *PatternRewriter) rewrite.PatternResult {
254     const any = GpuDialect.AnySyncOp{ .op = op };
255     const result = op.getResult(0) orelse return .failure;
256     var state = ir.Operation.State.init(SpirvDialect.AnySyncOp.operation_name, op.location);
257     state.addOperands(&.{ any.getMask(), any.getPredicate() });
258     state.addTypes(&.{result.type});
259     _ = rewriter.replaceOpWithNewOp(op, state) catch return .failure;
260     return .success;
261 }
262 
263 fn rewriteBallotSync(op: *ir.Operation, rewriter: *PatternRewriter) rewrite.PatternResult {
264     const ballot = GpuDialect.BallotSyncOp{ .op = op };
265     const result = op.getResult(0) orelse return .failure;
266     var state = ir.Operation.State.init(SpirvDialect.BallotSyncOp.operation_name, op.location);
267     state.addOperands(&.{ ballot.getMask(), ballot.getPredicate() });
268     state.addTypes(&.{result.type});
269     _ = rewriter.replaceOpWithNewOp(op, state) catch return .failure;
270     return .success;
271 }
272 
273 fn rewriteShflSync(op: *ir.Operation, rewriter: *PatternRewriter) rewrite.PatternResult {
274     const shfl = GpuDialect.ShflSyncOp{ .op = op };
275     const mode = shfl.getMode() orelse return .failure;
276     const result = op.getResult(0) orelse return .failure;
277     var state = ir.Operation.State.init(SpirvDialect.ShflSyncOp.operation_name, op.location);
278     state.addOperands(&.{ shfl.getMask(), shfl.getSrc(), shfl.getLaneOrDelta() });
279     state.addTypes(&.{result.type});
280     const mode_attr = rewriter.ir_ctx.getDialectAttr("spirv.shuffle_mode", mode.toString()) catch return .failure;
281     state.addAttributes(&.{.{ .name = "mode", .value = mode_attr }});
282     _ = rewriter.replaceOpWithNewOp(op, state) catch return .failure;
283     return .success;
284 }
285 
286 fn rewriteWarpReduce(op: *ir.Operation, rewriter: *PatternRewriter) rewrite.PatternResult {
287     const reduce = GpuDialect.WarpReduceOp{ .op = op };
288     const kind = reduce.getOpKind() orelse return .failure;
289     const result = op.getResult(0) orelse return .failure;
290     var state = ir.Operation.State.init(SpirvDialect.WarpReduceOp.operation_name, op.location);
291     state.addOperands(&.{ reduce.getMask(), reduce.getValue() });
292     state.addTypes(&.{result.type});
293     const op_attr = rewriter.ir_ctx.getDialectAttr("spirv.warp_op", kind.toString()) catch return .failure;
294     state.addAttributes(&.{.{ .name = "op", .value = op_attr }});
295     _ = rewriter.replaceOpWithNewOp(op, state) catch return .failure;
296     return .success;
297 }
298 
299 fn rewriteWarpScan(op: *ir.Operation, rewriter: *PatternRewriter) rewrite.PatternResult {
300     const scan = GpuDialect.WarpScanOp{ .op = op };
301     const kind = scan.getOpKind() orelse return .failure;
302     const inclusive = scan.isInclusive();
303     const result = op.getResult(0) orelse return .failure;
304     var state = ir.Operation.State.init(SpirvDialect.WarpScanOp.operation_name, op.location);
305     state.addOperands(&.{ scan.getMask(), scan.getValue() });
306     state.addTypes(&.{result.type});
307     const op_attr = rewriter.ir_ctx.getDialectAttr("spirv.warp_op", kind.toString()) catch return .failure;
308     const inclusive_attr = rewriter.ir_ctx.getBoolAttr(inclusive) catch return .failure;
309     state.addAttributes(&.{
310         .{ .name = "op", .value = op_attr },
311         .{ .name = "inclusive", .value = inclusive_attr },
312     });
313     _ = rewriter.replaceOpWithNewOp(op, state) catch return .failure;
314     return .success;
315 }
316 
317 const ScalarFlavor = enum { float, signed, unsigned };
318 
319 fn classifyScalar(ty: ir.Type) ?ScalarFlavor {
320     const name = ty.getDialectTypeName() orelse return null;
321     const kind = dialects.arith.scalarKindFromTypeName(name) orelse return null;
322     return switch (kind) {
323         .f16, .f32, .f64 => .float,
324         .u8, .u16, .u32, .u64, .index => .unsigned,
325         .i8, .i16, .i32, .i64 => .signed,
326         .bool, .bf16 => null,
327     };
328 }
329 
330 const BinaryKind = enum { add, sub, mul, div };
331 
332 fn rewriteArithConstant(op: *ir.Operation, rewriter: *PatternRewriter) rewrite.PatternResult {
333     const result = op.getResult(0) orelse return .failure;
334     const constant = ArithDialect.ConstantOp{ .op = op };
335 
336     if (constant.getIntValue()) |int_value| {
337         var state = ir.Operation.State.init(SpirvDialect.ConstantOp.operation_name, op.location);
338         state.addTypes(&.{result.type});
339         const value_attr = rewriter.ir_ctx.getI64Attr(int_value) catch return .failure;
340         state.addAttributes(&.{.{ .name = "value", .value = value_attr }});
341         _ = rewriter.replaceOpWithNewOp(op, state) catch return .failure;
342         return .success;
343     }
344 
345     if (constant.getFloatValue()) |float_value| {
346         var state = ir.Operation.State.init(SpirvDialect.ConstantOp.operation_name, op.location);
347         state.addTypes(&.{result.type});
348         const value_attr = rewriter.ir_ctx.getF64Attr(float_value) catch return .failure;
349         state.addAttributes(&.{.{ .name = "value", .value = value_attr }});
350         _ = rewriter.replaceOpWithNewOp(op, state) catch return .failure;
351         return .success;
352     }
353 
354     if (op.getAttrAs(ir.Attribute.BoolAttr, "value")) |bool_attr| {
355         var state = ir.Operation.State.init(SpirvDialect.ConstantOp.operation_name, op.location);
356         state.addTypes(&.{result.type});
357         const value_attr = rewriter.ir_ctx.getBoolAttr(bool_attr.getValue()) catch return .failure;
358         state.addAttributes(&.{.{ .name = "value", .value = value_attr }});
359         _ = rewriter.replaceOpWithNewOp(op, state) catch return .failure;
360         return .success;
361     }
362 
363     return .failure;
364 }
365 
366 fn rewriteArithAdd(op: *ir.Operation, rewriter: *PatternRewriter) rewrite.PatternResult {
367     return rewriteArithBinary(op, rewriter, .add);
368 }
369 
370 fn rewriteArithSub(op: *ir.Operation, rewriter: *PatternRewriter) rewrite.PatternResult {
371     return rewriteArithBinary(op, rewriter, .sub);
372 }
373 
374 fn rewriteArithMul(op: *ir.Operation, rewriter: *PatternRewriter) rewrite.PatternResult {
375     return rewriteArithBinary(op, rewriter, .mul);
376 }
377 
378 fn rewriteArithDiv(op: *ir.Operation, rewriter: *PatternRewriter) rewrite.PatternResult {
379     return rewriteArithBinary(op, rewriter, .div);
380 }
381 
382 fn rewriteArithBinary(
383     op: *ir.Operation,
384     rewriter: *PatternRewriter,
385     kind: BinaryKind,
386 ) rewrite.PatternResult {
387     if (op.operands.items.len != 2) return .failure;
388     const lhs = op.operands.items[0].value;
389     const rhs = op.operands.items[1].value;
390     const result = op.getResult(0) orelse return .failure;
391     const flavor = classifyScalar(result.type) orelse return .failure;
392 
393     const target_name = switch (kind) {
394         .add => if (flavor == .float) SpirvDialect.FAddOp.operation_name else SpirvDialect.IAddOp.operation_name,
395         .sub => if (flavor == .float) SpirvDialect.FSubOp.operation_name else SpirvDialect.ISubOp.operation_name,
396         .mul => if (flavor == .float) SpirvDialect.FMulOp.operation_name else SpirvDialect.IMulOp.operation_name,
397         .div => switch (flavor) {
398             .float => SpirvDialect.FDivOp.operation_name,
399             .signed => SpirvDialect.SDivOp.operation_name,
400             .unsigned => SpirvDialect.UDivOp.operation_name,
401         },
402     };
403 
404     var state = ir.Operation.State.init(target_name, op.location);
405     state.addOperands(&.{ lhs, rhs });
406     state.addTypes(&.{result.type});
407     _ = rewriter.replaceOpWithNewOp(op, state) catch return .failure;
408     return .success;
409 }
410 
411 test "spirv target spec publishes legality and conversion roots" {
412     const testing = std.testing;
413 
414     try testing.expectEqualStrings("spirv", target_spec.name);
415     try testing.expectEqualStrings(SpirvDialect.name, target_spec.target_dialect_name);
416     try testing.expectEqualStrings(gpu_to_spirv_pass_name, target_spec.pass_name);
417     try testing.expectEqual(illegal_ops.len, target_spec.conversionPatternRootCount());
418     try testing.expect(target_spec.legalizesDialect(SpirvDialect.name));
419     try testing.expect(target_spec.marksIllegalOp(GpuDialect.ThreadIdxOp.operation_name));
420     try testing.expect(target_spec.marksIllegalOp(ArithDialect.AddOp.operation_name));
421     try testing.expect(target_spec.hasConversionPatternRoot(GpuDialect.ThreadIdxOp.operation_name));
422     try testing.expect(target_spec.hasConversionPatternRoot(ArithDialect.DivOp.operation_name));
423 
424     var target = ConversionTarget.init(testing.allocator);
425     defer target.deinit();
426     try target_spec.applyLegality(&target);
427     try testing.expect(target.legal_dialects.contains(SpirvDialect.name));
428     try testing.expect(target.illegal_ops.contains(GpuDialect.ThreadIdxOp.operation_name));
429 }
430 
431 test "gpu to spirv conversion rewrites gpu + arith ops" {
432     const testing = std.testing;
433     var arena = std.heap.ArenaAllocator.init(std.testing.allocator);
434     defer arena.deinit();
435     const allocator = arena.allocator();
436 
437     var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
438     defer ctx.deinit(allocator);
439 
440     const loc = ir.Location.getUnknown();
441     const module = try dialects.BuiltinDialect.ModuleOp.create(&ctx, loc);
442     const module_block = module.getBodyBlock();
443 
444     var func = try dialects.FuncDialect.FuncOp.createKernel(&ctx, loc, "kernel", &.{});
445     try module_block.addOperation(func.op);
446 
447     const entry = func.getEntryBlock();
448     const tid = try GpuDialect.ThreadIdxOp.create(&ctx, loc, .x);
449     try entry.addOperation(tid.op);
450 
451     const index_type = try ArithDialect.getIndexType(&ctx);
452     const one = try ArithDialect.ConstantOp.createInt(&ctx, loc, index_type, 1);
453     try entry.addOperation(one.op);
454 
455     const add = try ArithDialect.AddOp.create(&ctx, loc, tid.getResult(), one.getResult());
456     try entry.addOperation(add.op);
457 
458     const barrier = try GpuDialect.BarrierOp.create(&ctx, loc, .block);
459     try entry.addOperation(barrier.op);
460 
461     const ret = try dialects.FuncDialect.ReturnOp.create(&ctx, loc, &.{});
462     try entry.addOperation(ret.op);
463 
464     var analysis_cache = passes.AnalysisCache.init(allocator, null);
465     defer analysis_cache.deinit();
466     var pass_ctx = PassContext.init(module.op, &ctx, allocator, &analysis_cache);
467     defer pass_ctx.deinit();
468 
469     const pass = createGpuToSpirvPass();
470     try testing.expectEqual(PassResult.success, pass.run(&pass_ctx));
471 
472     var saw_spirv_index = false;
473     var saw_spirv_add = false;
474     var saw_spirv_const = false;
475     var saw_spirv_barrier = false;
476 
477     var op_iter = entry.operations.head;
478     while (op_iter) |op_ptr| {
479         const op: *ir.Operation = @ptrCast(@alignCast(op_ptr));
480         if (std.mem.eql(u8, op.name.name, SpirvDialect.LocalInvocationIdOp.operation_name)) {
481             saw_spirv_index = true;
482         }
483         if (std.mem.eql(u8, op.name.name, SpirvDialect.IAddOp.operation_name)) {
484             saw_spirv_add = true;
485         }
486         if (std.mem.eql(u8, op.name.name, SpirvDialect.ConstantOp.operation_name)) {
487             saw_spirv_const = true;
488         }
489         if (std.mem.eql(u8, op.name.name, SpirvDialect.BarrierOp.operation_name)) {
490             saw_spirv_barrier = true;
491         }
492         op_iter = op.next_op;
493     }
494 
495     try testing.expect(saw_spirv_index);
496     try testing.expect(saw_spirv_add);
497     try testing.expect(saw_spirv_const);
498     try testing.expect(saw_spirv_barrier);
499 }
500 
501 test "spirv backend emits after gpu-to-spirv conversion" {
502     const testing = std.testing;
503     var arena = std.heap.ArenaAllocator.init(std.testing.allocator);
504     defer arena.deinit();
505     const allocator = arena.allocator();
506 
507     var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
508     defer ctx.deinit(allocator);
509 
510     const loc = ir.Location.getUnknown();
511     const module = try dialects.BuiltinDialect.ModuleOp.create(&ctx, loc);
512     const module_block = module.getBodyBlock();
513 
514     var func = try dialects.FuncDialect.FuncOp.createKernel(&ctx, loc, "kernel", &.{});
515     try module_block.addOperation(func.op);
516 
517     const entry = func.getEntryBlock();
518     const tid = try GpuDialect.ThreadIdxOp.create(&ctx, loc, .x);
519     try entry.addOperation(tid.op);
520 
521     const index_type = try ArithDialect.getIndexType(&ctx);
522     const one = try ArithDialect.ConstantOp.createInt(&ctx, loc, index_type, 1);
523     try entry.addOperation(one.op);
524 
525     const add = try ArithDialect.AddOp.create(&ctx, loc, tid.getResult(), one.getResult());
526     try entry.addOperation(add.op);
527 
528     const ret = try dialects.FuncDialect.ReturnOp.create(&ctx, loc, &.{});
529     try entry.addOperation(ret.op);
530 
531     var analysis_cache = passes.AnalysisCache.init(allocator, null);
532     defer analysis_cache.deinit();
533     var pass_ctx = PassContext.init(module.op, &ctx, allocator, &analysis_cache);
534     defer pass_ctx.deinit();
535 
536     const pass = createGpuToSpirvPass();
537     try testing.expectEqual(PassResult.success, pass.run(&pass_ctx));
538 
539     var emitter = spirv_emit.Emitter.init(allocator);
540     defer emitter.deinit();
541 
542     const bytes = try emitter.emitModuleBytes(module.op);
543     defer allocator.free(bytes);
544     try testing.expect(bytes.len > 0);
545 }
546 
547 test "spirv conversion rejects unknown target ops" {
548     const testing = std.testing;
549     var arena = std.heap.ArenaAllocator.init(std.testing.allocator);
550     defer arena.deinit();
551     const allocator = arena.allocator();
552 
553     var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
554     defer ctx.deinit(allocator);
555     try ctx.allowUnregistered();
556 
557     const loc = ir.Location.getUnknown();
558     const module = try dialects.BuiltinDialect.ModuleOp.create(&ctx, loc);
559     const unknown = try ctx.createOperation(ir.Operation.State.init("external.unknown", loc));
560     try module.getBodyBlock().addOperation(unknown);
561 
562     var analysis_cache = passes.AnalysisCache.init(allocator, null);
563     defer analysis_cache.deinit();
564     var pass_ctx = PassContext.init(module.op, &ctx, allocator, &analysis_cache);
565     defer pass_ctx.deinit();
566 
567     const pass = createGpuToSpirvPass();
568     try testing.expectEqual(PassResult.failure, pass.run(&pass_ctx));
569 }