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 }