lib/accy/src/preparation/stage.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const choir = @import("choir");
2
3 const activation_lowering = @import("activation.zig");
4 const backend_legalization = @import("backend.zig");
5 const bufferization = @import("bufferization/root.zig");
6 const canonicalization = @import("canonicalization.zig");
7 const constant_folding = @import("folding.zig");
8 const saturation = @import("saturation.zig");
9 const dtype_legalization = @import("dtype.zig");
10 const einsum_lowering = @import("einsum.zig");
11 const fusion = @import("fusion/root.zig");
12 const indexing_lowering = @import("indexing.zig");
13 const kernelization = @import("kernelization/root.zig");
14 const kernel_outlining = @import("outlining/root.zig");
15 const layout_planning = @import("layout.zig");
16 const loss_lowering = @import("loss.zig");
17 const memory_space = @import("memory.zig");
18 const schedule_planning = @import("schedule/root.zig");
19 const shape_analysis = @import("shape/root.zig");
20
21 const passes = choir.passes;
22
23 pub const ActivationLoweringOptions = activation_lowering.Options;
24 pub const EinsumLoweringOptions = einsum_lowering.Options;
25 pub const IndexingLoweringOptions = indexing_lowering.Options;
26 pub const LossLoweringOptions = loss_lowering.Options;
27 pub const TensorLoweringOptions = struct {
28 activation: activation_lowering.Options = .{},
29 einsum: einsum_lowering.Options = .{},
30 indexing: indexing_lowering.Options = .{},
31 loss: loss_lowering.Options = .{},
32
33 pub fn eql(self: TensorLoweringOptions, other: TensorLoweringOptions) bool {
34 return self.activation.eql(other.activation) and
35 self.einsum.eql(other.einsum) and
36 self.indexing.eql(other.indexing) and
37 self.loss.eql(other.loss);
38 }
39 };
40 pub const backend_legalization_pass_name = backend_legalization.backend_legalization_pass_name;
41 pub const bufferization_planning_pass_name = bufferization.bufferization_planning_pass_name;
42 pub const canonicalization_pass_name = canonicalization.canonicalization_pass_name;
43 pub const constant_folding_pass_name = constant_folding.constant_folding_pass_name;
44 pub const saturation_pass_name = saturation.saturation_pass_name;
45 pub const dtype_legalization_pass_name = dtype_legalization.dtype_legalization_pass_name;
46 pub const activation_lowering_pass_name = activation_lowering.activation_lowering_pass_name;
47 pub const einsum_lowering_pass_name = einsum_lowering.einsum_lowering_pass_name;
48 pub const indexing_lowering_pass_name = indexing_lowering.indexing_lowering_pass_name;
49 pub const loss_lowering_pass_name = loss_lowering.loss_lowering_pass_name;
50 pub const fusion_planning_pass_name = fusion.fusion_planning_pass_name;
51 pub const kernelization_pass_name = kernelization.kernelization_pass_name;
52 pub const kernel_outlining_planning_pass_name = kernel_outlining.kernel_outlining_planning_pass_name;
53 pub const layout_planning_pass_name = layout_planning.layout_planning_pass_name;
54 pub const memory_space_planning_pass_name = memory_space.memory_space_planning_pass_name;
55 pub const schedule_planning_pass_name = schedule_planning.schedule_planning_pass_name;
56 pub const shape_layout_pass_name = shape_analysis.shape_layout_pass_name;
57
58 pub fn backendLegalizationPass() passes.Pass {
59 return backend_legalization.backendLegalizationPass();
60 }
61
62 pub fn bufferizationPlanningPass() passes.Pass {
63 return bufferization.bufferizationPlanningPass();
64 }
65
66 pub fn canonicalizationPass() passes.Pass {
67 return canonicalization.canonicalizationPass();
68 }
69
70 pub fn shapeLayoutPropagationPass() passes.Pass {
71 return shape_analysis.shapeLayoutPropagationPass();
72 }
73
74 pub fn constantFoldingPass() passes.Pass {
75 return constant_folding.constantFoldingPass();
76 }
77
78 pub fn tensorSaturationPass() passes.Pass {
79 return saturation.tensorSaturationPass();
80 }
81
82 pub fn dtypeLegalizationPass() passes.Pass {
83 return dtype_legalization.dtypeLegalizationPass();
84 }
85
86 pub fn activationLoweringPass() passes.Pass {
87 return activation_lowering.activationLoweringPass();
88 }
89
90 pub fn activationLoweringPassWithOptions(options: *const activation_lowering.Options) passes.Pass {
91 return activation_lowering.activationLoweringPassWithOptions(options);
92 }
93
94 pub fn einsumLoweringPass() passes.Pass {
95 return einsum_lowering.einsumLoweringPass();
96 }
97
98 pub fn einsumLoweringPassWithOptions(options: *const einsum_lowering.Options) passes.Pass {
99 return einsum_lowering.einsumLoweringPassWithOptions(options);
100 }
101
102 pub fn indexingLoweringPass() passes.Pass {
103 return indexing_lowering.indexingLoweringPass();
104 }
105
106 pub fn indexingLoweringPassWithOptions(options: *const indexing_lowering.Options) passes.Pass {
107 return indexing_lowering.indexingLoweringPassWithOptions(options);
108 }
109
110 pub fn lossLoweringPass() passes.Pass {
111 return loss_lowering.lossLoweringPass();
112 }
113
114 pub fn lossLoweringPassWithOptions(options: *const loss_lowering.Options) passes.Pass {
115 return loss_lowering.lossLoweringPassWithOptions(options);
116 }
117
118 pub fn fusionPlanningPass() passes.Pass {
119 return fusion.fusionPlanningPass();
120 }
121
122 pub fn schedulePlanningPass() passes.Pass {
123 return schedule_planning.schedulePlanningPass();
124 }
125
126 pub fn kernelOutliningPlanningPass() passes.Pass {
127 return kernel_outlining.kernelOutliningPlanningPass();
128 }
129
130 pub fn kernelizationPass() passes.Pass {
131 return kernelization.kernelizationPass();
132 }
133
134 pub fn memorySpacePlanningPass() passes.Pass {
135 return memory_space.memorySpacePlanningPass();
136 }
137
138 pub fn layoutPlanningPass() passes.Pass {
139 return layout_planning.layoutPlanningPass();
140 }
141
142 pub const BackendPreparationStage = struct {
143 name: []const u8,
144 description: []const u8,
145 pass: passes.Pass,
146 analysis_name: ?[]const u8 = null,
147 options: []const passes.PassOptionSpec = &.{},
148 build_with_options: ?passes.PassOptionsBuilderFn = null,
149
150 fn registration(self: BackendPreparationStage) passes.PassRegistration {
151 return .{
152 .name = self.name,
153 .description = self.description,
154 .pass = self.pass,
155 .options = self.options,
156 .build_with_options = self.build_with_options,
157 };
158 }
159 };
160
161 pub const contract_pipeline_name = "accy-choir-contract";
162 pub const contract_pipeline_description =
163 "Normalize semantic Accy Choir into the backend contract product";
164 pub const tensor_pipeline_name = "accy-choir-tensor";
165 pub const tensor_pipeline_description =
166 "Materialize tensor-level Accy Choir before dispatch grouping";
167 pub const dispatch_pipeline_name = "accy-choir-dispatch";
168 pub const dispatch_pipeline_description =
169 "Plan Accy Choir dispatch groups and schedules before memory preparation";
170 pub const memory_pipeline_name = "accy-choir-memory";
171 pub const memory_pipeline_description =
172 "Plan Accy Choir buffers, memory spaces, and layouts before kernel preparation";
173 pub const kernel_pipeline_name = "accy-choir-kernel";
174 pub const kernel_pipeline_description =
175 "Outline and lower Accy Choir dispatches into kernel programs before backend legalization";
176 pub const target_pipeline_name = "accy-choir-target";
177 pub const target_pipeline_description =
178 "Legalize Accy Choir kernels for the selected backend target";
179
180 pub const contract_stages = [_]BackendPreparationStage{
181 .{
182 .name = canonicalization_pass_name,
183 .description = canonicalization.canonicalization_pass_description,
184 .pass = canonicalizationPass(),
185 },
186 .{
187 .name = saturation_pass_name,
188 .description = saturation.saturation_pass_description,
189 .pass = tensorSaturationPass(),
190 },
191 .{
192 .name = shape_layout_pass_name,
193 .description = shape_analysis.shape_layout_pass_description,
194 .pass = shapeLayoutPropagationPass(),
195 .analysis_name = shape_analysis.shape_layout_analysis_descriptor.name,
196 },
197 .{
198 .name = constant_folding_pass_name,
199 .description = constant_folding.constant_folding_pass_description,
200 .pass = constantFoldingPass(),
201 },
202 };
203
204 pub const tensor_stages = [_]BackendPreparationStage{
205 .{
206 .name = activation_lowering_pass_name,
207 .description = activation_lowering.activation_lowering_pass_description,
208 .pass = activationLoweringPass(),
209 .options = &activation_lowering.activation_lowering_pass_options,
210 .build_with_options = activation_lowering.activationLoweringPassFromOptions,
211 },
212 .{
213 .name = einsum_lowering_pass_name,
214 .description = einsum_lowering.einsum_lowering_pass_description,
215 .pass = einsumLoweringPass(),
216 .options = &einsum_lowering.einsum_lowering_pass_options,
217 .build_with_options = einsum_lowering.einsumLoweringPassFromOptions,
218 },
219 .{
220 .name = indexing_lowering_pass_name,
221 .description = indexing_lowering.indexing_lowering_pass_description,
222 .pass = indexingLoweringPass(),
223 .options = &indexing_lowering.indexing_lowering_pass_options,
224 .build_with_options = indexing_lowering.indexingLoweringPassFromOptions,
225 },
226 .{
227 .name = loss_lowering_pass_name,
228 .description = loss_lowering.loss_lowering_pass_description,
229 .pass = lossLoweringPass(),
230 .options = &loss_lowering.loss_lowering_pass_options,
231 .build_with_options = loss_lowering.lossLoweringPassFromOptions,
232 },
233 };
234
235 pub const dispatch_stages = [_]BackendPreparationStage{
236 .{
237 .name = fusion_planning_pass_name,
238 .description = fusion.fusion_planning_pass_description,
239 .pass = fusionPlanningPass(),
240 .analysis_name = fusion.fusion_plan_analysis_descriptor.name,
241 },
242 .{
243 .name = schedule_planning_pass_name,
244 .description = schedule_planning.schedule_planning_pass_description,
245 .pass = schedulePlanningPass(),
246 .analysis_name = schedule_planning.schedule_plan_analysis_descriptor.name,
247 },
248 };
249
250 pub const memory_stages = [_]BackendPreparationStage{
251 .{
252 .name = bufferization_planning_pass_name,
253 .description = bufferization.bufferization_planning_pass_description,
254 .pass = bufferizationPlanningPass(),
255 .analysis_name = bufferization.buffer_plan_analysis_descriptor.name,
256 },
257 .{
258 .name = memory_space_planning_pass_name,
259 .description = memory_space.memory_space_planning_pass_description,
260 .pass = memorySpacePlanningPass(),
261 .analysis_name = memory_space.memory_space_plan_analysis_descriptor.name,
262 },
263 .{
264 .name = layout_planning_pass_name,
265 .description = layout_planning.layout_planning_pass_description,
266 .pass = layoutPlanningPass(),
267 .analysis_name = layout_planning.layout_plan_analysis_descriptor.name,
268 },
269 };
270
271 pub const kernel_stages = [_]BackendPreparationStage{
272 .{
273 .name = kernel_outlining_planning_pass_name,
274 .description = kernel_outlining.kernel_outlining_planning_pass_description,
275 .pass = kernelOutliningPlanningPass(),
276 .analysis_name = kernel_outlining.kernel_outline_plan_analysis_descriptor.name,
277 },
278 .{
279 .name = kernelization_pass_name,
280 .description = kernelization.kernelization_pass_description,
281 .pass = kernelizationPass(),
282 .analysis_name = kernelization.kernelization_analysis_descriptor.name,
283 },
284 };
285
286 pub const target_stages = [_]BackendPreparationStage{
287 .{
288 .name = backend_legalization_pass_name,
289 .description = backend_legalization.backend_legalization_pass_description,
290 .pass = backendLegalizationPass(),
291 .analysis_name = backend_legalization.backend_legalization_analysis_descriptor.name,
292 },
293 .{
294 .name = dtype_legalization_pass_name,
295 .description = dtype_legalization.dtype_legalization_pass_description,
296 .pass = dtypeLegalizationPass(),
297 },
298 };
299
300 fn stagePassRegistrations(comptime stages: anytype) [stages.len]passes.PassRegistration {
301 comptime {
302 var registrations: [stages.len]passes.PassRegistration = undefined;
303 for (stages, 0..) |stage, index| {
304 registrations[index] = stage.registration();
305 }
306 return registrations;
307 }
308 }
309
310 fn allPassRegistrations() [contract_stages.len + tensor_stages.len + dispatch_stages.len + memory_stages.len + kernel_stages.len + target_stages.len]passes.PassRegistration {
311 comptime {
312 var registrations: [contract_stages.len + tensor_stages.len + dispatch_stages.len + memory_stages.len + kernel_stages.len + target_stages.len]passes.PassRegistration = undefined;
313 for (contract_stages, 0..) |stage, index| {
314 registrations[index] = stage.registration();
315 }
316 for (tensor_stages, 0..) |stage, index| {
317 registrations[contract_stages.len + index] = stage.registration();
318 }
319 for (dispatch_stages, 0..) |stage, index| {
320 registrations[contract_stages.len + tensor_stages.len + index] = stage.registration();
321 }
322 for (memory_stages, 0..) |stage, index| {
323 registrations[contract_stages.len + tensor_stages.len + dispatch_stages.len + index] = stage.registration();
324 }
325 for (kernel_stages, 0..) |stage, index| {
326 registrations[contract_stages.len + tensor_stages.len + dispatch_stages.len + memory_stages.len + index] = stage.registration();
327 }
328 for (target_stages, 0..) |stage, index| {
329 registrations[contract_stages.len + tensor_stages.len + dispatch_stages.len + memory_stages.len + kernel_stages.len + index] = stage.registration();
330 }
331 return registrations;
332 }
333 }
334
335 fn stageAnalysisCount(comptime stages: anytype) usize {
336 comptime {
337 var count: usize = 0;
338 for (stages) |stage| {
339 if (stage.analysis_name != null) count += 1;
340 }
341 return count;
342 }
343 }
344
345 fn stageAnalysisNames(comptime stages: anytype) [stageAnalysisCount(stages)][]const u8 {
346 comptime {
347 var names: [stageAnalysisCount(stages)][]const u8 = undefined;
348 var next: usize = 0;
349 for (stages) |stage| {
350 if (stage.analysis_name) |name| {
351 names[next] = name;
352 next += 1;
353 }
354 }
355 return names;
356 }
357 }
358
359 fn stagePlanSteps(comptime stages: anytype) [stages.len]passes.PassPlanStep {
360 comptime {
361 var steps: [stages.len]passes.PassPlanStep = undefined;
362 for (stages, 0..) |stage, index| {
363 steps[index] = passes.PassPlanStep.namedPass(stage.name);
364 }
365 return steps;
366 }
367 }
368
369 pub const contract_pipeline_registration = passes.PipelineRegistration{
370 .name = contract_pipeline_name,
371 .description = contract_pipeline_description,
372 .build = buildContractPipeline,
373 };
374
375 pub const tensor_pipeline_registration = passes.PipelineRegistration{
376 .name = tensor_pipeline_name,
377 .description = tensor_pipeline_description,
378 .build = buildTensorPipeline,
379 };
380
381 pub const dispatch_pipeline_registration = passes.PipelineRegistration{
382 .name = dispatch_pipeline_name,
383 .description = dispatch_pipeline_description,
384 .build = buildDispatchPipeline,
385 };
386
387 pub const memory_pipeline_registration = passes.PipelineRegistration{
388 .name = memory_pipeline_name,
389 .description = memory_pipeline_description,
390 .build = buildMemoryPipeline,
391 };
392
393 pub const kernel_pipeline_registration = passes.PipelineRegistration{
394 .name = kernel_pipeline_name,
395 .description = kernel_pipeline_description,
396 .build = buildKernelPipeline,
397 };
398
399 pub const target_pipeline_registration = passes.PipelineRegistration{
400 .name = target_pipeline_name,
401 .description = target_pipeline_description,
402 .build = buildTargetPipeline,
403 };
404
405 pub const contract_pass_registrations = stagePassRegistrations(contract_stages);
406 pub const tensor_pass_registrations = stagePassRegistrations(tensor_stages);
407 pub const dispatch_pass_registrations = stagePassRegistrations(dispatch_stages);
408 pub const memory_pass_registrations = stagePassRegistrations(memory_stages);
409 pub const kernel_pass_registrations = stagePassRegistrations(kernel_stages);
410 pub const target_pass_registrations = stagePassRegistrations(target_stages);
411 pub const accy_choir_pass_registrations = allPassRegistrations();
412
413 pub const contract_plan_steps = stagePlanSteps(contract_stages);
414 pub const tensor_plan_steps = stagePlanSteps(tensor_stages);
415 pub const dispatch_plan_steps = stagePlanSteps(dispatch_stages);
416 pub const memory_plan_steps = stagePlanSteps(memory_stages);
417 pub const kernel_plan_steps = stagePlanSteps(kernel_stages);
418 pub const target_plan_steps = stagePlanSteps(target_stages);
419
420 pub const contract_pass_plan = passes.PassPlan{ .steps = &contract_plan_steps };
421 pub const tensor_pass_plan = passes.PassPlan{ .steps = &tensor_plan_steps };
422 pub const dispatch_pass_plan = passes.PassPlan{ .steps = &dispatch_plan_steps };
423 pub const memory_pass_plan = passes.PassPlan{ .steps = &memory_plan_steps };
424 pub const kernel_pass_plan = passes.PassPlan{ .steps = &kernel_plan_steps };
425 pub const target_pass_plan = passes.PassPlan{ .steps = &target_plan_steps };
426
427 pub const accy_choir_package_extension = choir.extensions.PackageExtension{
428 .name = "accy-choir",
429 .passes = &accy_choir_pass_registrations,
430 .pipelines = &.{ contract_pipeline_registration, tensor_pipeline_registration, dispatch_pipeline_registration, memory_pipeline_registration, kernel_pipeline_registration, target_pipeline_registration },
431 };
432
433 const contract_analysis_names = stageAnalysisNames(contract_stages);
434 const tensor_analysis_names = stageAnalysisNames(tensor_stages);
435 const dispatch_analysis_names = stageAnalysisNames(dispatch_stages);
436 const memory_analysis_names = stageAnalysisNames(memory_stages);
437 const kernel_analysis_names = stageAnalysisNames(kernel_stages);
438 const target_analysis_names = stageAnalysisNames(target_stages);
439
440 pub const contract_pass_count = contract_stages.len;
441 pub const contract_analysis_count = contract_analysis_names.len;
442 pub const tensor_pass_count = tensor_stages.len;
443 pub const tensor_analysis_count = tensor_analysis_names.len;
444 pub const dispatch_pass_count = dispatch_stages.len;
445 pub const dispatch_analysis_count = dispatch_analysis_names.len;
446 pub const memory_pass_count = memory_stages.len;
447 pub const memory_analysis_count = memory_analysis_names.len;
448 pub const kernel_pass_count = kernel_stages.len;
449 pub const kernel_analysis_count = kernel_analysis_names.len;
450 pub const target_pass_count = target_stages.len;
451 pub const target_analysis_count = target_analysis_names.len;
452
453 pub fn contractPassName(index: usize) []const u8 {
454 return contract_stages[index].name;
455 }
456
457 pub fn contractAnalysisName(index: usize) []const u8 {
458 return contract_analysis_names[index];
459 }
460
461 pub fn tensorPassName(index: usize) []const u8 {
462 if (comptime tensor_stages.len == 0) {
463 unreachable;
464 } else {
465 return tensor_stages[index].name;
466 }
467 }
468
469 pub fn tensorAnalysisName(index: usize) []const u8 {
470 if (comptime tensor_analysis_names.len == 0) {
471 unreachable;
472 } else {
473 return tensor_analysis_names[index];
474 }
475 }
476
477 pub fn dispatchPassName(index: usize) []const u8 {
478 return dispatch_stages[index].name;
479 }
480
481 pub fn dispatchAnalysisName(index: usize) []const u8 {
482 return dispatch_analysis_names[index];
483 }
484
485 pub fn memoryPassName(index: usize) []const u8 {
486 return memory_stages[index].name;
487 }
488
489 pub fn memoryAnalysisName(index: usize) []const u8 {
490 return memory_analysis_names[index];
491 }
492
493 pub fn kernelPassName(index: usize) []const u8 {
494 return kernel_stages[index].name;
495 }
496
497 pub fn kernelAnalysisName(index: usize) []const u8 {
498 return kernel_analysis_names[index];
499 }
500
501 pub fn targetPassName(index: usize) []const u8 {
502 return target_stages[index].name;
503 }
504
505 pub fn targetAnalysisName(index: usize) []const u8 {
506 return target_analysis_names[index];
507 }
508
509 pub fn addContractPipeline(pm: *passes.PassManager) !void {
510 try contract_pipeline_registration.addTo(&pm.root);
511 }
512
513 pub fn addTensorPipeline(pm: *passes.PassManager) !void {
514 try tensor_pipeline_registration.addTo(&pm.root);
515 }
516
517 pub fn addTensorPipelineWithOptions(pm: *passes.PassManager, options: *const TensorLoweringOptions) !void {
518 try pm.addPass(activationLoweringPassWithOptions(&options.activation));
519 try pm.addPass(einsumLoweringPassWithOptions(&options.einsum));
520 try pm.addPass(indexingLoweringPassWithOptions(&options.indexing));
521 try pm.addPass(lossLoweringPassWithOptions(&options.loss));
522 }
523
524 pub fn addDispatchPipeline(pm: *passes.PassManager) !void {
525 try dispatch_pipeline_registration.addTo(&pm.root);
526 }
527
528 pub fn addMemoryPipeline(pm: *passes.PassManager) !void {
529 try memory_pipeline_registration.addTo(&pm.root);
530 }
531
532 pub fn addKernelPipeline(pm: *passes.PassManager) !void {
533 try kernel_pipeline_registration.addTo(&pm.root);
534 }
535
536 pub fn addTargetPipeline(pm: *passes.PassManager) !void {
537 try target_pipeline_registration.addTo(&pm.root);
538 }
539
540 fn buildContractPipeline(pm: *passes.OpPassManager) anyerror!void {
541 try buildStagePlan(pm, contract_pass_plan, &contract_pass_registrations);
542 }
543
544 fn buildTensorPipeline(pm: *passes.OpPassManager) anyerror!void {
545 try buildStagePlan(pm, tensor_pass_plan, &tensor_pass_registrations);
546 }
547
548 fn buildDispatchPipeline(pm: *passes.OpPassManager) anyerror!void {
549 try buildStagePlan(pm, dispatch_pass_plan, &dispatch_pass_registrations);
550 }
551
552 fn buildMemoryPipeline(pm: *passes.OpPassManager) anyerror!void {
553 try buildStagePlan(pm, memory_pass_plan, &memory_pass_registrations);
554 }
555
556 fn buildKernelPipeline(pm: *passes.OpPassManager) anyerror!void {
557 try buildStagePlan(pm, kernel_pass_plan, &kernel_pass_registrations);
558 }
559
560 fn buildTargetPipeline(pm: *passes.OpPassManager) anyerror!void {
561 try buildStagePlan(pm, target_pass_plan, &target_pass_registrations);
562 }
563
564 fn buildStagePlan(
565 pm: *passes.OpPassManager,
566 pass_plan: passes.PassPlan,
567 registrations: []const passes.PassRegistration,
568 ) anyerror!void {
569 var registry = passes.PassRegistry.init(pm.allocator);
570 defer registry.deinit();
571 for (registrations) |registration| {
572 try registry.registerPass(registration);
573 }
574 try pass_plan.addToOp(®istry, pm);
575 }