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 }