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 }