lib/accy/src/preparation/kernelization/lowering/pass.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const gpu = @import("gpu");
3 const choir = @import("choir");
4 const lowering = @import("root.zig");
5 const preparation = @import("../../root.zig");
6 const product_mod = @import("../model/root.zig");
7 const common = lowering.common;
8 const elementwise = lowering.elementwise;
9 const shape_lowering = lowering.shape_lowering;
10 const dot = lowering.dot;
11 const reduction = lowering.reduction;
12 const row_pipeline = lowering.row_pipeline;
13 const iterate_lowering = lowering.iterate_lowering;
14 const flash_attention = lowering.flash_attention;
15 const scan_lowering = lowering.scan_lowering;
16 const ir = choir.ir;
17 const passes = choir.passes;
18 const bufferization = preparation.bufferization;
19 const kernel_outlining = preparation.outlining;
20 const schedule_planning = preparation.schedule;
21 const shape_analysis = preparation.shape;
22 const KernelizationAnalysis = product_mod.KernelizationAnalysis;
23 const LoweredKernel = product_mod.LoweredKernel;
24
25 pub const kernelization_analysis_name = "accy-choir-kernelization-plan";
26 pub const kernelization_pass_name = "accy-choir-plan-kernelization";
27 pub const kernelization_pass_description =
28 "Lower scheduled Accy tensor work to kernel-language programs";
29
30 pub const kernelization_analysis_descriptor = passes.AnalysisDescriptor{
31 .id = passes.analysisId(kernelization_analysis_name),
32 .name = kernelization_analysis_name,
33 .work_contract = .{
34 .identity = .{ .name = kernelization_analysis_name, .version = 1 },
35 .estimate = lowering.work.analysis,
36 },
37 };
38
39 pub fn getKernelizationAnalysis(
40 pass_ctx: *passes.PassContext,
41 op: *ir.Operation,
42 ) !*KernelizationAnalysis {
43 const ptr = try pass_ctx.getAnalysis(
44 op,
45 &kernelization_analysis_descriptor,
46 computeKernelizationAnalysis,
47 cleanupKernelizationAnalysis,
48 );
49 return @ptrCast(@alignCast(ptr));
50 }
51
52 pub fn kernelizationPass() passes.Pass {
53 return .{
54 .name = kernelization_pass_name,
55 .description = kernelization_pass_description,
56 .run_fn = runKernelizationPass,
57 .work_contract = .{
58 .identity = .{ .name = kernelization_pass_name, .version = 1 },
59 .estimate = lowering.work.pass,
60 },
61 };
62 }
63
64 fn runKernelizationPass(pass_ctx: *passes.PassContext) passes.PassResult {
65 _ = getKernelizationAnalysis(pass_ctx, pass_ctx.op) catch return .failure;
66 pass_ctx.preserveAllAnalyses();
67 return .success;
68 }
69
70 fn computeKernelizationAnalysis(
71 pass_ctx: *passes.PassContext,
72 op: *ir.Operation,
73 ) anyerror!*anyopaque {
74 if (op.context != pass_ctx.ir_ctx) return error.InvalidArtifact;
75 const schedule_plan = try schedule_planning.getSchedulePlanAnalysis(pass_ctx, op);
76 const outline_plan = try kernel_outlining.getKernelOutlinePlanAnalysis(pass_ctx, op);
77 const buffer_plan = try bufferization.getBufferPlanAnalysis(pass_ctx, op);
78 const shape_plan = try shape_analysis.getShapeLayoutAnalysis(pass_ctx, op);
79 const target_profile = preparation.readBackendTargetProfile(op);
80 const format: ?gpu.ArtifactFormat = if (target_profile) |profile| profile.artifact_format else null;
81 const math_tier: gpu.BackendMathTier = if (target_profile) |profile| profile.math_tier else .exact;
82 const scan_schedules = preparation.target.readGeneratedScanSchedules(op);
83 const row_pipeline_schedules = preparation.target.readGeneratedRowPipelineSchedules(op);
84
85 const analysis = try pass_ctx.allocator.create(KernelizationAnalysis);
86 analysis.* = KernelizationAnalysis.init(
87 pass_ctx.allocator,
88 pass_ctx.ir_ctx.capacity.asLimits(),
89 ) catch |err| {
90 pass_ctx.allocator.destroy(analysis);
91 return common.mapKernelBuildError(err);
92 };
93 errdefer {
94 analysis.deinit();
95 pass_ctx.allocator.destroy(analysis);
96 }
97
98 try analysis.reserveKernelCapacity(outline_plan.kernels.items.len);
99
100 for (outline_plan.kernels.items) |outline| {
101 const work = workItemById(schedule_plan, outline.work_item_id) orelse return error.MissingScheduleWorkItem;
102 if (try lowerKernel(
103 pass_ctx.allocator,
104 analysis.context,
105 outline,
106 work.*,
107 buffer_plan,
108 shape_plan,
109 format,
110 math_tier,
111 scan_schedules,
112 row_pipeline_schedules,
113 )) |kernel| {
114 var owned = kernel;
115 var transferred = false;
116 errdefer if (!transferred) owned.deinit(pass_ctx.allocator);
117 try product_mod.addKernel(analysis, owned);
118 transferred = true;
119 }
120 }
121
122 return @ptrCast(analysis);
123 }
124
125 fn cleanupKernelizationAnalysis(ptr: *anyopaque, allocator: std.mem.Allocator) void {
126 const analysis: *KernelizationAnalysis = @ptrCast(@alignCast(ptr));
127 analysis.deinit();
128 allocator.destroy(analysis);
129 }
130
131 fn lowerKernel(
132 allocator: std.mem.Allocator,
133 ir_ctx: *ir.Context,
134 outline: product_mod.KernelOutline,
135 work: schedule_planning.ScheduleWorkItem,
136 buffer_plan: *const bufferization.BufferPlanAnalysis,
137 shape_plan: *const shape_analysis.ShapeLayoutAnalysis,
138 format: ?gpu.ArtifactFormat,
139 math_tier: gpu.BackendMathTier,
140 scan_schedules: ?[]const u8,
141 row_pipeline_schedules: ?[]const u8,
142 ) !?LoweredKernel {
143 return lowerKernelStrict(allocator, ir_ctx, outline, work, buffer_plan, shape_plan, format, math_tier, scan_schedules, row_pipeline_schedules) catch |err| switch (err) {
144 error.UnsupportedOperation,
145 error.UnsupportedArtifactFormat,
146 error.CapabilityMismatch,
147 => if (math_tier == .exact) null else return err,
148 else => |other| return other,
149 };
150 }
151
152 fn lowerKernelStrict(
153 allocator: std.mem.Allocator,
154 ir_ctx: *ir.Context,
155 outline: product_mod.KernelOutline,
156 work: schedule_planning.ScheduleWorkItem,
157 buffer_plan: *const bufferization.BufferPlanAnalysis,
158 shape_plan: *const shape_analysis.ShapeLayoutAnalysis,
159 format: ?gpu.ArtifactFormat,
160 math_tier: gpu.BackendMathTier,
161 scan_schedules: ?[]const u8,
162 row_pipeline_schedules: ?[]const u8,
163 ) common.LoweringError!LoweredKernel {
164 return switch (outline.kind) {
165 .elementwise => elementwise.lower(allocator, ir_ctx, outline, work, buffer_plan, shape_plan, format),
166 .shape => shape_lowering.lower(allocator, ir_ctx, outline, work, buffer_plan),
167 .dot_general => dot.lower(allocator, ir_ctx, outline, work, buffer_plan, format, math_tier),
168 .reduction => reduction.lower(allocator, ir_ctx, outline, work, buffer_plan, format),
169 .row_pipeline => row_pipeline.lower(allocator, ir_ctx, outline, work, buffer_plan, format, row_pipeline_schedules),
170 .iterate => iterate_lowering.lower(allocator, ir_ctx, outline, work, buffer_plan),
171 .flash_attention => flash_attention.lower(allocator, ir_ctx, outline, work, buffer_plan, format),
172 .scan => scan_lowering.lower(allocator, ir_ctx, outline, work, buffer_plan, format, scan_schedules),
173 .kernel_call => error.UnsupportedOperation,
174 };
175 }
176
177 fn workItemById(
178 schedule_plan: *const schedule_planning.SchedulePlanAnalysis,
179 work_item_id: usize,
180 ) ?*const schedule_planning.ScheduleWorkItem {
181 if (work_item_id < schedule_plan.work_items.items.len) {
182 const work = &schedule_plan.work_items.items[work_item_id];
183 if (work.id == work_item_id) return work;
184 }
185 for (schedule_plan.work_items.items) |*work| {
186 if (work.id == work_item_id) return work;
187 }
188 return null;
189 }