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(&registry, pm);
575 }