lib/choir/src/backends/gpu/spirv/emitter/catalog.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const choir = @import("../../../../root.zig");
3
4 const ir = choir.ir;
5 const dialects = choir.dialects;
6 const spirv_target = @import("../root.zig");
7 const gpu_target = @import("../../../../dialects/gpu/root.zig");
8 const dialect_writer = @import("dialect.zig");
9 const gpu = @import("gpu.zig");
10 const memory = @import("memory.zig");
11 const scalar = @import("scalar.zig");
12 const stage = @import("stage.zig");
13 const SpirvOp = @import("ops.zig").SpirvOp;
14
15 const GpuDialect = gpu_target.GpuDialect;
16 const ArithDialect = dialects.arith.ArithDialect;
17 const MemrefDialect = dialects.memref.MemrefDialect;
18 const FuncDialect = dialects.func.FuncDialect;
19 const ScfDialect = dialects.scf.ScfDialect;
20 const SpirvDialect = spirv_target.SpirvDialect;
21 const BuiltinKind = gpu.BuiltinKind;
22 const builtinLoad = stage.emitBuiltinLoad;
23
24 fn direct(comptime function: anytype) type {
25 return struct {
26 fn emit(
27 comptime ErrorSet: type,
28 comptime _: anytype,
29 self: anytype,
30 op: *ir.Operation,
31 ) ErrorSet!void {
32 return function(self, op);
33 }
34 };
35 }
36
37 fn withArg(comptime function: anytype, comptime argument: anytype) type {
38 return struct {
39 fn emit(
40 comptime ErrorSet: type,
41 comptime _: anytype,
42 self: anytype,
43 op: *ir.Operation,
44 ) ErrorSet!void {
45 return function(self, op, argument);
46 }
47 };
48 }
49
50 fn withTwoArgs(
51 comptime function: anytype,
52 comptime first: anytype,
53 comptime second: anytype,
54 ) type {
55 return struct {
56 fn emit(
57 comptime ErrorSet: type,
58 comptime _: anytype,
59 self: anytype,
60 op: *ir.Operation,
61 ) ErrorSet!void {
62 return function(self, op, first, second);
63 }
64 };
65 }
66
67 fn spirvBinary(
68 comptime opcode: u32,
69 comptime expectation: dialect_writer.SpirvBinaryExpectation,
70 ) type {
71 return withTwoArgs(dialect_writer.emitBinary, opcode, expectation);
72 }
73
74 fn arithBinary(comptime kind: scalar.BinaryKind) type {
75 return withArg(scalar.emitBinaryArith, kind);
76 }
77
78 fn arithMinMax(comptime kind: scalar.MinMaxKind) type {
79 return withArg(scalar.emitMinMaxArith, kind);
80 }
81
82 fn glslUnary(comptime opcode: u32) type {
83 return withTwoArgs(scalar.emitGlslExtUnary, opcode, scalar.FloatExtConstraint.float16_or_32);
84 }
85
86 fn arithBitwise(comptime kind: scalar.BitwiseBinaryKind) type {
87 return withArg(scalar.emitBitwiseBinary, kind);
88 }
89
90 fn arithShift(comptime kind: scalar.ShiftKind) type {
91 return withArg(scalar.emitShiftArith, kind);
92 }
93
94 fn gpuIndex(comptime builtin: gpu.BuiltinKind) type {
95 return withArg(gpu.emitIndex, builtin);
96 }
97
98 fn gpuSubgroupIndex(comptime builtin: gpu.BuiltinKind) type {
99 return withArg(gpu.emitSubgroupIndex, builtin);
100 }
101
102 fn method(comptime name: []const u8) type {
103 return struct {
104 fn emit(
105 comptime ErrorSet: type,
106 comptime Control: anytype,
107 self: anytype,
108 op: *ir.Operation,
109 ) ErrorSet!void {
110 return @call(.auto, @field(Control, name), .{ self, op });
111 }
112 };
113 }
114
115 fn unsupportedControlFlow() type {
116 return struct {
117 fn emit(
118 comptime ErrorSet: type,
119 comptime _: anytype,
120 _: anytype,
121 _: *ir.Operation,
122 ) ErrorSet!void {
123 return error.UnsupportedControlFlow;
124 }
125 };
126 }
127
128 fn operation(comptime Operation: type, comptime Emitter: type) type {
129 if (!@hasDecl(Operation, "operation_name")) @compileError("SPIR-V catalog name is missing");
130 if (!@hasDecl(Emitter, "emit")) @compileError("SPIR-V catalog emitter is missing");
131 return struct {
132 pub const name = Operation.operation_name;
133
134 fn emit(
135 comptime ErrorSet: type,
136 comptime Control: anytype,
137 self: anytype,
138 op: *ir.Operation,
139 ) ErrorSet!void {
140 return Emitter.emit(ErrorSet, Control, self, op);
141 }
142 };
143 }
144
145 pub const entries = .{
146 operation(SpirvDialect.ConstantOp, direct(dialect_writer.emitConstant)),
147 operation(SpirvDialect.IAddOp, spirvBinary(SpirvOp.IAdd, .int_any)),
148 operation(SpirvDialect.FAddOp, spirvBinary(SpirvOp.FAdd, .float)),
149 operation(SpirvDialect.ISubOp, spirvBinary(SpirvOp.ISub, .int_any)),
150 operation(SpirvDialect.FSubOp, spirvBinary(SpirvOp.FSub, .float)),
151 operation(SpirvDialect.IMulOp, spirvBinary(SpirvOp.IMul, .int_any)),
152 operation(SpirvDialect.FMulOp, spirvBinary(SpirvOp.FMul, .float)),
153 operation(SpirvDialect.UDivOp, spirvBinary(SpirvOp.UDiv, .int_unsigned)),
154 operation(SpirvDialect.SDivOp, spirvBinary(SpirvOp.SDiv, .int_signed)),
155 operation(SpirvDialect.FDivOp, spirvBinary(SpirvOp.FDiv, .float)),
156 operation(ArithDialect.ConstantOp, direct(scalar.emitConstant)),
157 operation(ArithDialect.AddOp, arithBinary(.add)),
158 operation(ArithDialect.SubOp, arithBinary(.sub)),
159 operation(ArithDialect.MulOp, arithBinary(.mul)),
160 operation(ArithDialect.DivOp, arithBinary(.div)),
161 operation(ArithDialect.MaxOp, arithMinMax(.max)),
162 operation(ArithDialect.MinOp, arithMinMax(.min)),
163 operation(ArithDialect.NegOp, direct(scalar.emitNegArith)),
164 operation(ArithDialect.ExpOp, glslUnary(scalar.GLSLstd450.Exp)),
165 operation(ArithDialect.LogOp, glslUnary(scalar.GLSLstd450.Log)),
166 operation(ArithDialect.TanhOp, glslUnary(scalar.GLSLstd450.Tanh)),
167 operation(ArithDialect.SqrtOp, direct(scalar.emitSqrtArith)),
168 operation(ArithDialect.AbsOp, direct(scalar.emitAbsArith)),
169 operation(ArithDialect.SinOp, glslUnary(scalar.GLSLstd450.Sin)),
170 operation(ArithDialect.CosOp, glslUnary(scalar.GLSLstd450.Cos)),
171 operation(ArithDialect.TanOp, glslUnary(scalar.GLSLstd450.Tan)),
172 operation(ArithDialect.FloorOp, glslUnary(scalar.GLSLstd450.Floor)),
173 operation(ArithDialect.RoundOp, direct(scalar.emitRoundArith)),
174 operation(ArithDialect.TruncOp, glslUnary(scalar.GLSLstd450.Trunc)),
175 operation(ArithDialect.PowOp, direct(scalar.emitPowArith)),
176 operation(ArithDialect.Atan2Op, direct(scalar.emitAtan2Arith)),
177 operation(ArithDialect.FmaOp, direct(scalar.emitFmaArith)),
178 operation(ArithDialect.AndOp, arithBitwise(.band)),
179 operation(ArithDialect.OrOp, arithBitwise(.bor)),
180 operation(ArithDialect.XorOp, arithBitwise(.bxor)),
181 operation(ArithDialect.NotOp, direct(scalar.emitNotArith)),
182 operation(ArithDialect.ShlOp, arithShift(.shl)),
183 operation(ArithDialect.ShrOp, arithShift(.shr)),
184 operation(ArithDialect.UshrOp, arithShift(.ushr)),
185 operation(ArithDialect.UmulhiOp, direct(scalar.emitUmulhi)),
186 operation(ArithDialect.CastOp, direct(scalar.emitCast)),
187 operation(ArithDialect.BitcastOp, direct(scalar.emitBitcast)),
188 operation(ArithDialect.CmpOp, direct(scalar.emitCmp)),
189 operation(ArithDialect.SelectOp, direct(scalar.emitSelect)),
190 operation(MemrefDialect.LoadOp, direct(memory.emitLoad)),
191 operation(MemrefDialect.StoreOp, direct(memory.emitStore)),
192 operation(MemrefDialect.AtomicRmwOp, direct(memory.emitAtomicRmw)),
193 operation(MemrefDialect.AtomicCasOp, direct(memory.emitAtomicCas)),
194 operation(MemrefDialect.AllocOp, direct(memory.emitAlloc)),
195 operation(MemrefDialect.AllocaOp, direct(memory.emitDeclaredAlloca)),
196 operation(FuncDialect.CallOp, method("emitCall")),
197 operation(SpirvDialect.LocalInvocationIdOp, gpuIndex(.local_invocation_id)),
198 operation(SpirvDialect.WorkgroupIdOp, gpuIndex(.workgroup_id)),
199 operation(SpirvDialect.WorkgroupSizeOp, gpuIndex(.workgroup_size)),
200 operation(SpirvDialect.NumWorkgroupsOp, gpuIndex(.num_workgroups)),
201 operation(SpirvDialect.GlobalInvocationIdOp, gpuIndex(.global_invocation_id)),
202 operation(SpirvDialect.BarrierOp, direct(gpu.emitBarrier)),
203 operation(SpirvDialect.SyncWarpOp, direct(gpu.emitSyncWarp)),
204 operation(SpirvDialect.ActiveMaskOp, direct(gpu.emitActiveMask)),
205 operation(SpirvDialect.AllSyncOp, withArg(gpu.emitAllAny, gpu.AllAnyKind.all)),
206 operation(SpirvDialect.AnySyncOp, withArg(gpu.emitAllAny, gpu.AllAnyKind.any)),
207 operation(SpirvDialect.BallotSyncOp, direct(gpu.emitBallot)),
208 operation(SpirvDialect.ShflSyncOp, direct(gpu.emitShuffle)),
209 operation(SpirvDialect.WarpReduceOp, direct(gpu.emitWarpReduce)),
210 operation(SpirvDialect.WarpScanOp, direct(gpu.emitWarpScan)),
211 operation(GpuDialect.ThreadIdxOp, gpuIndex(.local_invocation_id)),
212 operation(GpuDialect.BlockIdxOp, gpuIndex(.workgroup_id)),
213 operation(GpuDialect.BlockDimOp, gpuIndex(.workgroup_size)),
214 operation(GpuDialect.GridDimOp, gpuIndex(.num_workgroups)),
215 operation(GpuDialect.GlobalIdxOp, gpuIndex(.global_invocation_id)),
216 operation(GpuDialect.LaneIdOp, gpuSubgroupIndex(.subgroup_local_invocation_id)),
217 operation(GpuDialect.WarpIdOp, gpuSubgroupIndex(.subgroup_id)),
218 operation(GpuDialect.BarrierOp, direct(gpu.emitBarrier)),
219 operation(GpuDialect.SyncWarpOp, direct(gpu.emitSyncWarp)),
220 operation(GpuDialect.ActiveMaskOp, direct(gpu.emitActiveMask)),
221 operation(GpuDialect.AllSyncOp, withArg(gpu.emitAllAny, gpu.AllAnyKind.all)),
222 operation(GpuDialect.AnySyncOp, withArg(gpu.emitAllAny, gpu.AllAnyKind.any)),
223 operation(GpuDialect.BallotSyncOp, direct(gpu.emitBallot)),
224 operation(GpuDialect.ShflSyncOp, direct(gpu.emitShuffle)),
225 operation(GpuDialect.WarpReduceOp, direct(gpu.emitWarpReduce)),
226 operation(GpuDialect.WarpScanOp, direct(gpu.emitWarpScan)),
227 operation(GpuDialect.StageInputOp, direct(stage.emitStageInput)),
228 operation(GpuDialect.StageOutputOp, direct(stage.emitStageOutput)),
229 operation(GpuDialect.PositionOp, direct(stage.emitPosition)),
230 operation(GpuDialect.FragCoordOp, withArg(builtinLoad, BuiltinKind.frag_coord)),
231 operation(GpuDialect.VertexIndexOp, withArg(builtinLoad, BuiltinKind.vertex_index)),
232 operation(GpuDialect.InstanceIndexOp, withArg(builtinLoad, BuiltinKind.instance_index)),
233 operation(GpuDialect.FrontFacingOp, withArg(builtinLoad, BuiltinKind.front_facing)),
234 operation(GpuDialect.SampledTextureOp, direct(stage.emitSampledTexture)),
235 operation(GpuDialect.SampleOp, withArg(stage.emitSample, stage.Lod.implicit)),
236 operation(GpuDialect.SampleLodOp, withArg(stage.emitSample, stage.Lod.explicit)),
237 operation(GpuDialect.DpdxOp, withArg(stage.emitDerivative, SpirvOp.DPdxFine)),
238 operation(GpuDialect.DpdyOp, withArg(stage.emitDerivative, SpirvOp.DPdyFine)),
239 operation(GpuDialect.FwidthOp, withArg(stage.emitDerivative, SpirvOp.FwidthFine)),
240 operation(GpuDialect.PushConstantOp, direct(stage.emitPushConstant)),
241 operation(GpuDialect.UniformOp, direct(stage.emitUniform)),
242 operation(ScfDialect.IfOp, method("emitScfIf")),
243 operation(ScfDialect.ForOp, method("emitScfFor")),
244 operation(ScfDialect.WhileOp, method("emitScfWhile")),
245 operation(ScfDialect.YieldOp, unsupportedControlFlow()),
246 operation(ScfDialect.ConditionOp, unsupportedControlFlow()),
247 operation(FuncDialect.ReturnOp, method("emitReturn")),
248 };
249
250 pub fn supports(name: []const u8) bool {
251 inline for (entries) |Entry| {
252 if (std.mem.eql(u8, name, Entry.name)) return true;
253 }
254 return false;
255 }
256
257 pub fn emit(
258 comptime ErrorSet: type,
259 comptime Control: anytype,
260 self: anytype,
261 op: *ir.Operation,
262 ) ErrorSet!void {
263 inline for (entries) |Entry| {
264 if (std.mem.eql(u8, op.name.name, Entry.name)) {
265 return Entry.emit(ErrorSet, Control, self, op);
266 }
267 }
268 return error.UnsupportedOperation;
269 }
270
271 comptime {
272 @setEvalBranchQuota(100_000);
273 for (entries, 0..) |Candidate, index| {
274 if (!@hasDecl(Candidate, "emit")) @compileError("SPIR-V catalog emitter is missing");
275 for (entries, 0..) |Prior, prior_index| {
276 if (prior_index >= index) continue;
277 if (std.mem.eql(u8, Candidate.name, Prior.name)) {
278 @compileError("SPIR-V catalog operation name is not unique");
279 }
280 }
281 }
282 }
283
284 test "spirv catalog resolves every operation to one emitter" {
285 try std.testing.expectEqual(@as(usize, 102), entries.len);
286 inline for (entries) |Entry| try std.testing.expect(supports(Entry.name));
287 try std.testing.expect(!supports("test.unsupported"));
288 }