lib/choir/src/backends/gpu/registration.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const choir = @import("../../root.zig");
3 const gpu = @import("../../dialects/gpu/root.zig");
4 const spirv = @import("spirv/root.zig");
5 const nvptx = @import("nvptx/root.zig");
6 const lowering = @import("lowering.zig");
7
8 fn loadGpuDialect(ctx: *choir.ir.Context) !void {
9 try choir.ir.dialects.loadDialectSpec(ctx, gpu.GpuDialect.spec);
10 }
11
12 fn loadSpirvDialect(ctx: *choir.ir.Context) !void {
13 try choir.ir.dialects.loadDialectSpec(ctx, spirv.SpirvDialect.spec);
14 }
15
16 fn loadNvptxDialect(ctx: *choir.ir.Context) !void {
17 try choir.ir.dialects.loadDialectSpec(ctx, nvptx.NvptxDialect.spec);
18 }
19
20 pub const dialect_registrations = [_]choir.extensions.DialectRegistration{
21 .{ .name = "gpu", .load = loadGpuDialect },
22 .{ .name = "spirv", .load = loadSpirvDialect, .is_backend = true },
23 .{ .name = "nvptx", .load = loadNvptxDialect, .is_backend = true },
24 };
25
26 pub const package_extension = choir.extensions.PackageExtension{
27 .name = "choir-gpu",
28 .dialects = &dialect_registrations,
29 .passes = &lowering.pass_registrations,
30 .pipelines = &lowering.pipeline_registrations,
31 };
32
33 pub fn registerTargetDialects(ctx: *choir.ir.Context) !void {
34 for (dialect_registrations) |registration| {
35 ctx.registerDialectLoader(registration.name, registration.load) catch |err| switch (err) {
36 error.DuplicateDialectLoader => {},
37 else => return err,
38 };
39 if (registration.is_backend) try ctx.registerBackendDialect(registration.name);
40 }
41 }
42
43 /// Load the native compilation dialect closure before an owned program Context activates.
44 pub fn prepareCompilationDialects(ctx: *choir.ir.Context) !void {
45 try choir.dialects.registerChoirDialect(ctx);
46 try registerTargetDialects(ctx);
47 inline for (.{ "builtin", "arith", "memref", "scf", "func" }) |name| {
48 _ = try ctx.getOrLoadDialect(name);
49 }
50 for (dialect_registrations) |registration| _ = try ctx.getOrLoadDialect(registration.name);
51 }
52
53 test "registerTargetDialects is idempotent for kernel contexts" {
54 const testing = std.testing;
55 var ctx = try choir.ir.Context.init(testing.allocator, choir.ir.Context.Limits.testing);
56 defer ctx.deinit(testing.allocator);
57
58 try registerTargetDialects(&ctx);
59 try registerTargetDialects(&ctx);
60
61 _ = try ctx.getOrLoadDialect("gpu");
62 _ = try ctx.getOrLoadDialect("spirv");
63 _ = try ctx.getOrLoadDialect("nvptx");
64 try testing.expect(ctx.isBackendDialect("spirv"));
65 try testing.expect(ctx.isBackendDialect("nvptx"));
66 }