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 }