lib/accy/src/preparation/pipeline.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const gpu = @import("gpu");
3 const choir = @import("choir");
4 const backend_legalization = @import("backend.zig");
5 const canonicalization = @import("canonicalization.zig");
6 const saturation = @import("saturation.zig");
7 const accy_choir = @import("../choir/root.zig");
8 const target_product = @import("../target/root.zig");
9 const dialect_mod = accy_choir.dialect;
10 const execution_mod = @import("execution.zig");
11 const library_preparation = @import("library.zig");
12 const kernelization = @import("kernelization/root.zig");
13 const kernel_outlining = @import("outlining/root.zig");
14 const contract = accy_choir.contract;
15 const dispatch = accy_choir.dispatch;
16 const kernel_product = accy_choir.gpu;
17 const memory_product = accy_choir.memory;
18 const semantic = accy_choir.semantic;
19 const tensor = accy_choir.tensor;
20 const schedule_planning = @import("schedule/root.zig");
21 const prepared_mod = @import("prepared.zig");
22 const product_mod = @import("product.zig");
23 const run_mod = @import("run.zig");
24 const stage_mod = @import("stage.zig");
25 const target_profile = @import("target.zig");
26 const accy_root = @import("../root.zig");
27
28 const ir = choir.ir;
29 const passes = choir.passes;
30
31 pub const BackendTargetProfile = target_profile.BackendTargetProfile;
32 pub const canonicalizeModule = canonicalization.canonicalizeModule;
33 pub const saturateTensorOperations = saturation.saturateTensorOperations;
34 pub const setBackendTargetProfile = target_profile.setBackendTargetProfile;
35 pub const readBackendTargetProfile = target_profile.readBackendTargetProfile;
36 pub const KernelLibraryLowering = library_preparation.KernelLibraryLowering;
37 pub const ActivationLoweringOptions = stage_mod.ActivationLoweringOptions;
38 pub const EinsumLoweringOptions = stage_mod.EinsumLoweringOptions;
39 pub const IndexingLoweringOptions = stage_mod.IndexingLoweringOptions;
40 pub const TensorLoweringOptions = stage_mod.TensorLoweringOptions;
41 pub const backend_legalization_pass_name = stage_mod.backend_legalization_pass_name;
42 pub const bufferization_planning_pass_name = stage_mod.bufferization_planning_pass_name;
43 pub const canonicalization_pass_name = stage_mod.canonicalization_pass_name;
44 pub const constant_folding_pass_name = stage_mod.constant_folding_pass_name;
45 pub const saturation_pass_name = stage_mod.saturation_pass_name;
46 pub const dtype_legalization_pass_name = stage_mod.dtype_legalization_pass_name;
47 pub const activation_lowering_pass_name = stage_mod.activation_lowering_pass_name;
48 pub const einsum_lowering_pass_name = stage_mod.einsum_lowering_pass_name;
49 pub const indexing_lowering_pass_name = stage_mod.indexing_lowering_pass_name;
50 pub const loss_lowering_pass_name = stage_mod.loss_lowering_pass_name;
51 pub const fusion_planning_pass_name = stage_mod.fusion_planning_pass_name;
52 pub const kernelization_pass_name = stage_mod.kernelization_pass_name;
53 pub const kernel_outlining_planning_pass_name = stage_mod.kernel_outlining_planning_pass_name;
54 pub const layout_planning_pass_name = stage_mod.layout_planning_pass_name;
55 pub const memory_space_planning_pass_name = stage_mod.memory_space_planning_pass_name;
56 pub const schedule_planning_pass_name = stage_mod.schedule_planning_pass_name;
57 pub const shape_layout_pass_name = stage_mod.shape_layout_pass_name;
58
59 pub const PipelineError = execution_mod.PipelineError;
60
61 pub const BackendPreparationStats = run_mod.BackendPreparationStats;
62 pub const BackendPreparationFailureKind = run_mod.BackendPreparationFailureKind;
63 pub const BackendPreparationFailure = run_mod.BackendPreparationFailure;
64 pub const BackendPreparationTiming = run_mod.BackendPreparationTiming;
65 pub const BackendPreparationRunOptions = run_mod.BackendPreparationRunOptions;
66 pub const BackendPreparationRun = run_mod.BackendPreparationRun;
67
68 pub const BackendPreparationProductStamps = product_mod.BackendPreparationProductStamps;
69 pub const BackendPreparationProductKeys = product_mod.BackendPreparationProductKeys;
70 pub const BackendPreparedModule = product_mod.BackendPreparedModule;
71 pub const ContractPreparationResult = execution_mod.ContractPreparationResult;
72 pub const TensorPreparationResult = execution_mod.TensorPreparationResult;
73 pub const DispatchPreparationResult = execution_mod.DispatchPreparationResult;
74 pub const MemoryPreparationResult = execution_mod.MemoryPreparationResult;
75 pub const KernelPreparationResult = execution_mod.KernelPreparationResult;
76 pub const TargetPreparationResult = execution_mod.TargetPreparationResult;
77 pub const BackendPreparedJob = prepared_mod.BackendPreparedJob;
78
79 pub const backendLegalizationPass = stage_mod.backendLegalizationPass;
80 pub const bufferizationPlanningPass = stage_mod.bufferizationPlanningPass;
81 pub const canonicalizationPass = stage_mod.canonicalizationPass;
82 pub const shapeLayoutPropagationPass = stage_mod.shapeLayoutPropagationPass;
83 pub const constantFoldingPass = stage_mod.constantFoldingPass;
84 pub const tensorSaturationPass = stage_mod.tensorSaturationPass;
85 pub const dtypeLegalizationPass = stage_mod.dtypeLegalizationPass;
86 pub const activationLoweringPass = stage_mod.activationLoweringPass;
87 pub const activationLoweringPassWithOptions = stage_mod.activationLoweringPassWithOptions;
88 pub const einsumLoweringPass = stage_mod.einsumLoweringPass;
89 pub const einsumLoweringPassWithOptions = stage_mod.einsumLoweringPassWithOptions;
90 pub const indexingLoweringPass = stage_mod.indexingLoweringPass;
91 pub const indexingLoweringPassWithOptions = stage_mod.indexingLoweringPassWithOptions;
92 pub const lossLoweringPass = stage_mod.lossLoweringPass;
93 pub const lossLoweringPassWithOptions = stage_mod.lossLoweringPassWithOptions;
94 pub const fusionPlanningPass = stage_mod.fusionPlanningPass;
95 pub const schedulePlanningPass = stage_mod.schedulePlanningPass;
96 pub const kernelOutliningPlanningPass = stage_mod.kernelOutliningPlanningPass;
97 pub const kernelizationPass = stage_mod.kernelizationPass;
98 pub const memorySpacePlanningPass = stage_mod.memorySpacePlanningPass;
99 pub const layoutPlanningPass = stage_mod.layoutPlanningPass;
100 pub const BackendPreparationStage = stage_mod.BackendPreparationStage;
101 pub const contract_pipeline_name = stage_mod.contract_pipeline_name;
102 pub const contract_pipeline_description = stage_mod.contract_pipeline_description;
103 pub const tensor_pipeline_name = stage_mod.tensor_pipeline_name;
104 pub const tensor_pipeline_description = stage_mod.tensor_pipeline_description;
105 pub const dispatch_pipeline_name = stage_mod.dispatch_pipeline_name;
106 pub const dispatch_pipeline_description = stage_mod.dispatch_pipeline_description;
107 pub const memory_pipeline_name = stage_mod.memory_pipeline_name;
108 pub const memory_pipeline_description = stage_mod.memory_pipeline_description;
109 pub const kernel_pipeline_name = stage_mod.kernel_pipeline_name;
110 pub const kernel_pipeline_description = stage_mod.kernel_pipeline_description;
111 pub const target_pipeline_name = stage_mod.target_pipeline_name;
112 pub const target_pipeline_description = stage_mod.target_pipeline_description;
113 pub const contract_stages = stage_mod.contract_stages;
114 pub const tensor_stages = stage_mod.tensor_stages;
115 pub const dispatch_stages = stage_mod.dispatch_stages;
116 pub const memory_stages = stage_mod.memory_stages;
117 pub const kernel_stages = stage_mod.kernel_stages;
118 pub const target_stages = stage_mod.target_stages;
119 pub const contract_pipeline_registration = stage_mod.contract_pipeline_registration;
120 pub const tensor_pipeline_registration = stage_mod.tensor_pipeline_registration;
121 pub const dispatch_pipeline_registration = stage_mod.dispatch_pipeline_registration;
122 pub const memory_pipeline_registration = stage_mod.memory_pipeline_registration;
123 pub const kernel_pipeline_registration = stage_mod.kernel_pipeline_registration;
124 pub const target_pipeline_registration = stage_mod.target_pipeline_registration;
125 pub const contract_pass_registrations = stage_mod.contract_pass_registrations;
126 pub const tensor_pass_registrations = stage_mod.tensor_pass_registrations;
127 pub const dispatch_pass_registrations = stage_mod.dispatch_pass_registrations;
128 pub const memory_pass_registrations = stage_mod.memory_pass_registrations;
129 pub const kernel_pass_registrations = stage_mod.kernel_pass_registrations;
130 pub const target_pass_registrations = stage_mod.target_pass_registrations;
131 pub const accy_choir_pass_registrations = stage_mod.accy_choir_pass_registrations;
132 pub const accy_choir_package_extension = stage_mod.accy_choir_package_extension;
133 pub const contract_pass_plan = stage_mod.contract_pass_plan;
134 pub const tensor_pass_plan = stage_mod.tensor_pass_plan;
135 pub const dispatch_pass_plan = stage_mod.dispatch_pass_plan;
136 pub const memory_pass_plan = stage_mod.memory_pass_plan;
137 pub const kernel_pass_plan = stage_mod.kernel_pass_plan;
138 pub const target_pass_plan = stage_mod.target_pass_plan;
139 pub const contract_pass_count = stage_mod.contract_pass_count;
140 pub const contract_analysis_count = stage_mod.contract_analysis_count;
141 pub const tensor_pass_count = stage_mod.tensor_pass_count;
142 pub const tensor_analysis_count = stage_mod.tensor_analysis_count;
143 pub const dispatch_pass_count = stage_mod.dispatch_pass_count;
144 pub const dispatch_analysis_count = stage_mod.dispatch_analysis_count;
145 pub const memory_pass_count = stage_mod.memory_pass_count;
146 pub const memory_analysis_count = stage_mod.memory_analysis_count;
147 pub const kernel_pass_count = stage_mod.kernel_pass_count;
148 pub const kernel_analysis_count = stage_mod.kernel_analysis_count;
149 pub const target_pass_count = stage_mod.target_pass_count;
150 pub const target_analysis_count = stage_mod.target_analysis_count;
151 pub const contractPassName = stage_mod.contractPassName;
152 pub const contractAnalysisName = stage_mod.contractAnalysisName;
153 pub const tensorPassName = stage_mod.tensorPassName;
154 pub const tensorAnalysisName = stage_mod.tensorAnalysisName;
155 pub const dispatchPassName = stage_mod.dispatchPassName;
156 pub const dispatchAnalysisName = stage_mod.dispatchAnalysisName;
157 pub const memoryPassName = stage_mod.memoryPassName;
158 pub const memoryAnalysisName = stage_mod.memoryAnalysisName;
159 pub const kernelPassName = stage_mod.kernelPassName;
160 pub const kernelAnalysisName = stage_mod.kernelAnalysisName;
161 pub const targetPassName = stage_mod.targetPassName;
162 pub const targetAnalysisName = stage_mod.targetAnalysisName;
163 pub const addContractPipeline = stage_mod.addContractPipeline;
164 pub const addTensorPipeline = stage_mod.addTensorPipeline;
165 pub const addTensorPipelineWithOptions = stage_mod.addTensorPipelineWithOptions;
166 pub const addDispatchPipeline = stage_mod.addDispatchPipeline;
167 pub const addMemoryPipeline = stage_mod.addMemoryPipeline;
168 pub const addKernelPipeline = stage_mod.addKernelPipeline;
169 pub const addTargetPipeline = stage_mod.addTargetPipeline;
170
171 pub const runTargetPipeline = execution_mod.runTargetPipeline;
172 pub const runTargetPipelineWithOptions = execution_mod.runTargetPipelineWithOptions;
173 pub const runTargetPipelineWithDiagnostics = execution_mod.runTargetPipelineWithDiagnostics;
174
175 pub const runBackendPreparationPipelineFromSemanticModule =
176 prepared_mod.runBackendPreparationPipelineFromSemanticModule;
177 pub const prepareBackendJobFromSemanticModule =
178 prepared_mod.prepareBackendJobFromSemanticModule;
179
180 pub const prepareContractJobFromSemanticModule = execution_mod.prepareContractJobFromSemanticModule;
181 pub const prepareTensorJobFromContractJob = execution_mod.prepareTensorJobFromContractJob;
182 pub const prepareDispatchJobFromTensorJob = execution_mod.prepareDispatchJobFromTensorJob;
183 pub const prepareMemoryJobFromDispatchJob = execution_mod.prepareMemoryJobFromDispatchJob;
184 pub const prepareKernelJobFromMemoryJob = execution_mod.prepareKernelJobFromMemoryJob;
185 pub const prepareTargetJobFromKernelJob = execution_mod.prepareTargetJobFromKernelJob;
186
187 pub const prepareBackendJobFromContractJob =
188 prepared_mod.prepareBackendJobFromContractJob;
189 pub const prepareBackendJobFromContractPreparationResult =
190 prepared_mod.prepareBackendJobFromContractPreparationResult;
191 pub const prepareBackendJobFromTensorPreparationResult =
192 prepared_mod.prepareBackendJobFromTensorPreparationResult;
193 pub const prepareBackendJobFromDispatchPreparationResult =
194 prepared_mod.prepareBackendJobFromDispatchPreparationResult;
195 pub const prepareBackendJobFromMemoryPreparationResult =
196 prepared_mod.prepareBackendJobFromMemoryPreparationResult;
197 pub const prepareBackendJobFromKernelPreparationResult =
198 prepared_mod.prepareBackendJobFromKernelPreparationResult;
199 pub const prepareBackendJobFromTargetPreparationResult =
200 prepared_mod.prepareBackendJobFromTargetPreparationResult;
201
202 pub const prepareContractJobFromSemanticModuleWithRun =
203 execution_mod.prepareContractJobFromSemanticModuleWithRun;
204 pub const prepareTensorJobFromContractJobWithRun =
205 execution_mod.prepareTensorJobFromContractJobWithRun;
206 pub const prepareDispatchJobFromTensorJobWithRun =
207 execution_mod.prepareDispatchJobFromTensorJobWithRun;
208 pub const prepareMemoryJobFromDispatchJobWithRun =
209 execution_mod.prepareMemoryJobFromDispatchJobWithRun;
210 pub const prepareKernelJobFromMemoryJobWithRun =
211 execution_mod.prepareKernelJobFromMemoryJobWithRun;
212 pub const prepareTargetJobFromKernelJobWithRun =
213 execution_mod.prepareTargetJobFromKernelJobWithRun;
214
215 const testing = std.testing;
216
217 pub fn buildBackendPreparationContext(
218 allocator: std.mem.Allocator,
219 context_limits: ir.Context.Limits,
220 ) !ir.Context {
221 return try semantic.buildSemanticContext(allocator, context_limits);
222 }
223
224 const countOperationTree = execution_mod.countOperationTree;
225
226 fn expectProductStamp(stamp: choir.product.incremental.ProductStamp, name: []const u8, fingerprint: u64) !void {
227 try testing.expectEqualStrings(name, stamp.name);
228 try testing.expectEqual(fingerprint, stamp.fingerprint);
229 }
230
231 fn readSymbolName(func: *ir.Operation) ?[]const u8 {
232 return ir.SymbolTable.getSymbolName(func);
233 }
234
235 fn hasFunctionNameSuffix(module_body: *ir.Block, suffix: []const u8) bool {
236 var iter = module_body.operations.head;
237 while (iter) |op_ptr| {
238 const op: *ir.Operation = @ptrCast(@alignCast(op_ptr));
239 if (std.mem.eql(u8, op.name.name, "func.func")) {
240 if (readSymbolName(op)) |existing_name| {
241 if (std.mem.endsWith(u8, existing_name, suffix)) return true;
242 }
243 }
244 iter = op.next_op;
245 }
246 return false;
247 }
248
249 fn expectPreparedKernelPlan(
250 allocator: std.mem.Allocator,
251 choir_mod: *ir.Operation,
252 ctx: *ir.Context,
253 expected_kernel_count: usize,
254 ) !void {
255 var cache = passes.AnalysisCache.init(allocator, null);
256 defer cache.deinit();
257 var pass_ctx = passes.PassContext.init(choir_mod, ctx, allocator, &cache);
258 defer pass_ctx.deinit();
259
260 const schedule_plan = try schedule_planning.getSchedulePlanAnalysis(&pass_ctx, choir_mod);
261 const outline_plan = try kernel_outlining.getKernelOutlinePlanAnalysis(&pass_ctx, choir_mod);
262 const kernel_plan = try kernelization.getKernelizationAnalysis(&pass_ctx, choir_mod);
263 const legal_plan = try backend_legalization.getBackendLegalizationAnalysis(&pass_ctx, choir_mod);
264
265 try testing.expectEqual(expected_kernel_count, schedule_plan.workItemCount());
266 try testing.expectEqual(expected_kernel_count, outline_plan.kernelCount());
267 try testing.expectEqual(expected_kernel_count, kernel_plan.kernelCount());
268 try testing.expectEqual(expected_kernel_count, legal_plan.kernelCount());
269 try testing.expect(legal_plan.isLegal());
270 }
271
272 fn expectTargetStagePipeline(pm: *const passes.PassManager) !void {
273 try testing.expectEqual(target_pass_count, pm.root.pipeline.items.len);
274 inline for (target_stages, 0..) |stage, index| {
275 switch (pm.root.pipeline.items[index]) {
276 .pass => |pass| try testing.expectEqualStrings(stage.name, pass.name),
277 .nested => try testing.expect(false),
278 }
279 }
280 }
281
282 fn expectTensorStagePipeline(pm: *const passes.PassManager) !void {
283 try testing.expectEqual(tensor_pass_count, pm.root.pipeline.items.len);
284 inline for (tensor_stages, 0..) |stage, index| {
285 switch (pm.root.pipeline.items[index]) {
286 .pass => |pass| try testing.expectEqualStrings(stage.name, pass.name),
287 .nested => try testing.expect(false),
288 }
289 }
290 }
291
292 fn expectDispatchStagePipeline(pm: *const passes.PassManager) !void {
293 try testing.expectEqual(dispatch_pass_count, pm.root.pipeline.items.len);
294 inline for (dispatch_stages, 0..) |stage, index| {
295 switch (pm.root.pipeline.items[index]) {
296 .pass => |pass| try testing.expectEqualStrings(stage.name, pass.name),
297 .nested => try testing.expect(false),
298 }
299 }
300 }
301
302 fn expectMemoryStagePipeline(pm: *const passes.PassManager) !void {
303 try testing.expectEqual(memory_pass_count, pm.root.pipeline.items.len);
304 inline for (memory_stages, 0..) |stage, index| {
305 switch (pm.root.pipeline.items[index]) {
306 .pass => |pass| try testing.expectEqualStrings(stage.name, pass.name),
307 .nested => try testing.expect(false),
308 }
309 }
310 }
311
312 fn expectKernelStagePipeline(pm: *const passes.PassManager) !void {
313 try testing.expectEqual(kernel_pass_count, pm.root.pipeline.items.len);
314 inline for (kernel_stages, 0..) |stage, index| {
315 switch (pm.root.pipeline.items[index]) {
316 .pass => |pass| try testing.expectEqualStrings(stage.name, pass.name),
317 .nested => try testing.expect(false),
318 }
319 }
320 }
321
322 fn expectTargetStageText(text: []const u8) !void {
323 var iter = std.mem.splitScalar(u8, text, ',');
324 inline for (target_stages) |stage| {
325 const name = iter.next() orelse return error.MissingTargetStage;
326 try testing.expectEqualStrings(stage.name, name);
327 }
328 try testing.expect(iter.next() == null);
329 }
330
331 fn expectStagePlanText(
332 allocator: std.mem.Allocator,
333 comptime stages: anytype,
334 pass_plan: passes.PassPlan,
335 ) !void {
336 const text = try pass_plan.formatAlloc(allocator);
337 defer allocator.free(text);
338 var iter = std.mem.splitScalar(u8, text, ',');
339 inline for (stages) |stage| {
340 const name = iter.next() orelse return error.MissingStage;
341 try testing.expectEqualStrings(stage.name, name);
342 }
343 try testing.expect(iter.next() == null);
344 }
345
346 fn buildSemanticAddModule(
347 allocator: std.mem.Allocator,
348 name: []const u8,
349 dims: []const i64,
350 ) !*semantic.SemanticModule {
351 var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
352 errdefer builder.deinit();
353 const ty = try builder.tensor(.f32, dims);
354 var fb = try builder.beginFunction(name, &.{ ty, ty }, &.{ty});
355 const sum = try fb.add(fb.parameter(0), fb.parameter(1));
356 try fb.return_(&.{sum});
357 try fb.finish();
358 return try builder.finish();
359 }
360
361 fn buildSemanticAddMulModule(
362 allocator: std.mem.Allocator,
363 name: []const u8,
364 dims: []const i64,
365 ) !*semantic.SemanticModule {
366 var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
367 errdefer builder.deinit();
368 const ty = try builder.tensor(.f32, dims);
369 var fb = try builder.beginFunction(name, &.{ ty, ty, ty }, &.{ ty, ty });
370 const sum = try fb.add(fb.parameter(0), fb.parameter(1));
371 const product = try fb.mul(sum, fb.parameter(2));
372 try fb.return_(&.{ sum, product });
373 try fb.finish();
374 return try builder.finish();
375 }
376
377 fn prepareDispatchResultForContinuationTest(
378 allocator: std.mem.Allocator,
379 name: []const u8,
380 ) !DispatchPreparationResult {
381 const semantic_module = try buildSemanticAddMulModule(allocator, name, &.{4});
382 var contract_result = try prepareContractJobFromSemanticModuleWithRun(allocator, semantic_module, .{});
383 var contract_owned = true;
384 errdefer if (contract_owned) contract_result.deinit();
385
386 contract_owned = false;
387 var tensor_result = try prepareTensorJobFromContractJobWithRun(allocator, contract_result.module, .{});
388 var tensor_owned = true;
389 errdefer if (tensor_owned) tensor_result.deinit();
390
391 tensor_owned = false;
392 return try prepareDispatchJobFromTensorJobWithRun(allocator, tensor_result.module, .{});
393 }
394
395 fn prepareMemoryResultForContinuationTest(
396 allocator: std.mem.Allocator,
397 name: []const u8,
398 ) !MemoryPreparationResult {
399 var dispatch_result = try prepareDispatchResultForContinuationTest(allocator, name);
400 var dispatch_owned = true;
401 errdefer if (dispatch_owned) dispatch_result.deinit();
402
403 dispatch_owned = false;
404 return try prepareMemoryJobFromDispatchJobWithRun(allocator, dispatch_result.module, .{});
405 }
406
407 fn prepareKernelResultForContinuationTest(
408 allocator: std.mem.Allocator,
409 name: []const u8,
410 ) !KernelPreparationResult {
411 var memory_result = try prepareMemoryResultForContinuationTest(allocator, name);
412 var memory_owned = true;
413 errdefer if (memory_owned) memory_result.deinit();
414
415 memory_owned = false;
416 return try prepareKernelJobFromMemoryJobWithRun(allocator, memory_result.module, .{});
417 }
418
419 fn prepareTargetResultForContinuationTest(
420 allocator: std.mem.Allocator,
421 name: []const u8,
422 ) !TargetPreparationResult {
423 var kernel_result = try prepareKernelResultForContinuationTest(allocator, name);
424 var kernel_owned = true;
425 errdefer if (kernel_owned) kernel_result.deinit();
426
427 kernel_owned = false;
428 return try prepareTargetJobFromKernelJobWithRun(allocator, kernel_result.module, .{});
429 }
430
431 fn expectPreparedRunFingerprints(prepared: *const BackendPreparedJob) !void {
432 const target_module = try prepared.targetModule();
433 const tensor_module = target_module.kernel_module.memory_module.dispatch_module.tensor_module;
434 try testing.expectEqual(
435 tensor_module.contract_module.fingerprint(),
436 prepared.run.contract_fingerprint,
437 );
438 try testing.expectEqual(tensor_module.fingerprint(), prepared.run.tensor_fingerprint);
439 try testing.expectEqual(target_module.fingerprint(), prepared.run.target_fingerprint);
440 try testing.expect(prepared.run.dispatch_fingerprint != prepared.run.tensor_fingerprint);
441 try testing.expect(prepared.run.memory_fingerprint != prepared.run.tensor_fingerprint);
442 try testing.expect(prepared.run.kernel_fingerprint != prepared.run.tensor_fingerprint);
443 }
444
445 fn buildSemanticAddMulSubMaxModule(
446 allocator: std.mem.Allocator,
447 name: []const u8,
448 dims: []const i64,
449 ) !*semantic.SemanticModule {
450 var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
451 errdefer builder.deinit();
452 const ty = try builder.tensor(.f32, dims);
453 var fb = try builder.beginFunction(name, &.{ ty, ty }, &.{ty});
454 const rhs = fb.parameter(1);
455 var current = try fb.add(fb.parameter(0), rhs);
456 current = try fb.mul(current, rhs);
457 current = try fb.sub(current, rhs);
458 current = try fb.max(current, rhs);
459 try fb.return_(&.{current});
460 try fb.finish();
461 return try builder.finish();
462 }
463
464 test "target pipeline validates kernel candidates without static cloned IR" {
465 const allocator = testing.allocator;
466
467 const module = try buildSemanticAddModule(allocator, "pass_add4", &.{4});
468 defer module.deinit();
469
470 const choir_mod = module.choir_module;
471 const ctx = module.context();
472 var pm = passes.PassManager.init(allocator);
473 defer pm.deinit();
474 try addTargetPipeline(&pm);
475
476 try testing.expectEqual(passes.PassResult.success, pm.run(choir_mod, ctx));
477 try testing.expectEqual(@as(u64, target_pass_count), pm.stats.pass_runs);
478 try testing.expect(pm.stats.analysis_misses > 0);
479
480 const body = choir_mod.getRegion(0).?.getEntryBlock().?;
481 try testing.expect(ir.inspection.functionByNameInBlock(body, "pass_add4") != null);
482 try testing.expect(!hasFunctionNameSuffix(body, "lowered"));
483 try expectPreparedKernelPlan(allocator, choir_mod, ctx, 1);
484 }
485
486 test "prepareContractJobFromSemanticModule materializes normalized contract job" {
487 const allocator = testing.allocator;
488
489 var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
490 defer builder.deinit();
491
492 const ty = try builder.tensor(.f32, &.{1});
493 var fb = try builder.beginFunction("contract_fold_add", &.{}, &.{ty});
494 const one: [1]f32 = .{1};
495 const two: [1]f32 = .{2};
496 const lhs = try fb.constant(ty, std.mem.sliceAsBytes(one[0..]));
497 const rhs = try fb.constant(ty, std.mem.sliceAsBytes(two[0..]));
498 const sum = try fb.add(lhs, rhs);
499 try fb.return_(&.{sum});
500 try fb.finish();
501
502 const semantic_module = try builder.finish();
503 const contract_module = try prepareContractJobFromSemanticModule(allocator, semantic_module, .{});
504 defer contract_module.deinit();
505
506 try testing.expectEqualStrings(contract.product_name, accy_choir.contract_product_name);
507 try contract_module.verify();
508 try testing.expectEqual(@as(usize, 0), ir.inspection.countOperationsNamed(contract_module.choir_module, dialect_mod.AccyDialect.AddOp.operation_name));
509 }
510
511 test "prepareTensorJobFromContractJob materializes tensor job" {
512 const allocator = testing.allocator;
513
514 const semantic_module = try buildSemanticAddMulModule(allocator, "tensor_product_add_mul", &.{4});
515 const contract_module = try prepareContractJobFromSemanticModule(allocator, semantic_module, .{});
516 const tensor_module = try prepareTensorJobFromContractJob(allocator, contract_module, .{});
517 defer tensor_module.deinit();
518
519 try testing.expectEqualStrings(tensor.product_name, accy_choir.tensor_product_name);
520 try tensor_module.verify();
521 }
522
523 test "prepareDispatchJobFromTensorJob materializes dispatch job" {
524 const allocator = testing.allocator;
525
526 const semantic_module = try buildSemanticAddMulModule(allocator, "dispatch_product_add_mul", &.{4});
527 const contract_module = try prepareContractJobFromSemanticModule(allocator, semantic_module, .{});
528 const tensor_module = try prepareTensorJobFromContractJob(allocator, contract_module, .{});
529 const dispatch_module = try prepareDispatchJobFromTensorJob(allocator, tensor_module, .{});
530 defer dispatch_module.deinit();
531
532 try testing.expectEqualStrings(dispatch.product_name, accy_choir.dispatch_product_name);
533 try dispatch_module.verify();
534 try testing.expect(dispatch_module.tensor_module.analysis_cache.entries.count() > 0);
535 var pass_ctx = dispatch_module.passContext();
536 defer pass_ctx.deinit();
537 const plan = try schedule_planning.getSchedulePlanAnalysis(&pass_ctx, pass_ctx.op);
538 try testing.expectEqual(@as(usize, 2), plan.workItemCount());
539 }
540
541 test "prepareDispatchJobFromTensorJob fingerprints dispatch analysis product" {
542 const allocator = testing.allocator;
543
544 const first_semantic_module = try buildSemanticAddMulModule(allocator, "dispatch_fingerprint_first", &.{4});
545 const first_contract_module = try prepareContractJobFromSemanticModule(allocator, first_semantic_module, .{});
546 const first_tensor_module = try prepareTensorJobFromContractJob(allocator, first_contract_module, .{});
547 var first_dispatch = try prepareDispatchJobFromTensorJobWithRun(
548 allocator,
549 first_tensor_module,
550 .{},
551 );
552 defer first_dispatch.deinit();
553
554 const second_semantic_module = try buildSemanticAddMulModule(allocator, "dispatch_fingerprint_second", &.{4});
555 const second_contract_module = try prepareContractJobFromSemanticModule(allocator, second_semantic_module, .{});
556 const second_tensor_module = try prepareTensorJobFromContractJob(allocator, second_contract_module, .{});
557 var second_dispatch = try prepareDispatchJobFromTensorJobWithRun(
558 allocator,
559 second_tensor_module,
560 .{},
561 );
562 defer second_dispatch.deinit();
563
564 const changed_semantic_module = try buildSemanticAddMulSubMaxModule(allocator, "dispatch_fingerprint_changed", &.{4});
565 const changed_contract_module = try prepareContractJobFromSemanticModule(allocator, changed_semantic_module, .{});
566 const changed_tensor_module = try prepareTensorJobFromContractJob(allocator, changed_contract_module, .{});
567 var changed_dispatch = try prepareDispatchJobFromTensorJobWithRun(
568 allocator,
569 changed_tensor_module,
570 .{},
571 );
572 defer changed_dispatch.deinit();
573
574 try testing.expectEqual(first_dispatch.plan_fingerprint, second_dispatch.plan_fingerprint);
575 try testing.expect(first_dispatch.plan_fingerprint != changed_dispatch.plan_fingerprint);
576 }
577
578 test "prepareMemoryJobFromDispatchJob materializes memory job" {
579 const allocator = testing.allocator;
580
581 const semantic_module = try buildSemanticAddMulModule(allocator, "memory_product_add_mul", &.{4});
582 const contract_module = try prepareContractJobFromSemanticModule(allocator, semantic_module, .{});
583 const tensor_module = try prepareTensorJobFromContractJob(allocator, contract_module, .{});
584 const dispatch_module = try prepareDispatchJobFromTensorJob(allocator, tensor_module, .{});
585 const memory_module = try prepareMemoryJobFromDispatchJob(allocator, dispatch_module, .{});
586 defer memory_module.deinit();
587
588 try testing.expectEqualStrings(memory_product.product_name, accy_choir.memory_product_name);
589 try memory_module.verify();
590 try testing.expect(memory_module.dispatch_module.tensor_module.analysis_cache.entries.count() > 0);
591 var pass_ctx = memory_module.passContext();
592 defer pass_ctx.deinit();
593 const buffers = @import("bufferization/root.zig");
594 const plan = try buffers.getBufferPlanAnalysis(&pass_ctx, pass_ctx.op);
595 try testing.expectEqual(@as(usize, 5), plan.slotCount());
596 }
597
598 test "prepareMemoryJobFromDispatchJob fingerprints memory analysis product" {
599 const allocator = testing.allocator;
600
601 const first_semantic_module = try buildSemanticAddMulModule(allocator, "memory_fingerprint_first", &.{4});
602 const first_contract_module = try prepareContractJobFromSemanticModule(allocator, first_semantic_module, .{});
603 const first_tensor_module = try prepareTensorJobFromContractJob(allocator, first_contract_module, .{});
604 const first_dispatch_module = try prepareDispatchJobFromTensorJob(allocator, first_tensor_module, .{});
605 var first_memory = try prepareMemoryJobFromDispatchJobWithRun(
606 allocator,
607 first_dispatch_module,
608 .{},
609 );
610 defer first_memory.deinit();
611
612 const second_semantic_module = try buildSemanticAddMulModule(allocator, "memory_fingerprint_second", &.{4});
613 const second_contract_module = try prepareContractJobFromSemanticModule(allocator, second_semantic_module, .{});
614 const second_tensor_module = try prepareTensorJobFromContractJob(allocator, second_contract_module, .{});
615 const second_dispatch_module = try prepareDispatchJobFromTensorJob(allocator, second_tensor_module, .{});
616 var second_memory = try prepareMemoryJobFromDispatchJobWithRun(
617 allocator,
618 second_dispatch_module,
619 .{},
620 );
621 defer second_memory.deinit();
622
623 const changed_semantic_module = try buildSemanticAddMulSubMaxModule(allocator, "memory_fingerprint_changed", &.{4});
624 const changed_contract_module = try prepareContractJobFromSemanticModule(allocator, changed_semantic_module, .{});
625 const changed_tensor_module = try prepareTensorJobFromContractJob(allocator, changed_contract_module, .{});
626 const changed_dispatch_module = try prepareDispatchJobFromTensorJob(allocator, changed_tensor_module, .{});
627 var changed_memory = try prepareMemoryJobFromDispatchJobWithRun(
628 allocator,
629 changed_dispatch_module,
630 .{},
631 );
632 defer changed_memory.deinit();
633
634 try testing.expectEqual(first_memory.plan_fingerprint, second_memory.plan_fingerprint);
635 try testing.expect(first_memory.plan_fingerprint != changed_memory.plan_fingerprint);
636 }
637
638 test "prepareKernelJobFromMemoryJob materializes kernel job" {
639 const allocator = testing.allocator;
640
641 const semantic_module = try buildSemanticAddMulModule(allocator, "kernel_product_add_mul", &.{4});
642 const contract_module = try prepareContractJobFromSemanticModule(allocator, semantic_module, .{});
643 const tensor_module = try prepareTensorJobFromContractJob(allocator, contract_module, .{});
644 const dispatch_module = try prepareDispatchJobFromTensorJob(allocator, tensor_module, .{});
645 const memory_module = try prepareMemoryJobFromDispatchJob(allocator, dispatch_module, .{});
646 const kernel_module = try prepareKernelJobFromMemoryJob(allocator, memory_module, .{});
647 defer kernel_module.deinit();
648
649 try testing.expectEqualStrings(kernel_product.product_name, accy_choir.kernel_product_name);
650 try kernel_module.verify();
651 try testing.expect(kernel_module.memory_module.dispatch_module.tensor_module.analysis_cache.entries.count() > 0);
652
653 var pass_ctx = kernel_module.passContext();
654 defer pass_ctx.deinit();
655 const kernel_plan = try kernelization.getKernelizationAnalysis(&pass_ctx, kernel_module.choir_module);
656 try testing.expectEqual(@as(usize, 2), kernel_plan.kernelCount());
657 }
658
659 test "prepareKernelJobFromMemoryJob fingerprints kernel analysis product" {
660 const allocator = testing.allocator;
661
662 const first_semantic_module = try buildSemanticAddMulModule(allocator, "kernel_fingerprint_first", &.{4});
663 const first_contract_module = try prepareContractJobFromSemanticModule(allocator, first_semantic_module, .{});
664 const first_tensor_module = try prepareTensorJobFromContractJob(allocator, first_contract_module, .{});
665 const first_dispatch_module = try prepareDispatchJobFromTensorJob(allocator, first_tensor_module, .{});
666 const first_memory_module = try prepareMemoryJobFromDispatchJob(allocator, first_dispatch_module, .{});
667 var first_kernel = try prepareKernelJobFromMemoryJobWithRun(
668 allocator,
669 first_memory_module,
670 .{},
671 );
672 defer first_kernel.deinit();
673
674 const second_semantic_module = try buildSemanticAddMulModule(allocator, "kernel_fingerprint_second", &.{4});
675 const second_contract_module = try prepareContractJobFromSemanticModule(allocator, second_semantic_module, .{});
676 const second_tensor_module = try prepareTensorJobFromContractJob(allocator, second_contract_module, .{});
677 const second_dispatch_module = try prepareDispatchJobFromTensorJob(allocator, second_tensor_module, .{});
678 const second_memory_module = try prepareMemoryJobFromDispatchJob(allocator, second_dispatch_module, .{});
679 var second_kernel = try prepareKernelJobFromMemoryJobWithRun(
680 allocator,
681 second_memory_module,
682 .{},
683 );
684 defer second_kernel.deinit();
685
686 const changed_semantic_module = try buildSemanticAddMulSubMaxModule(allocator, "kernel_fingerprint_changed", &.{4});
687 const changed_contract_module = try prepareContractJobFromSemanticModule(allocator, changed_semantic_module, .{});
688 const changed_tensor_module = try prepareTensorJobFromContractJob(allocator, changed_contract_module, .{});
689 const changed_dispatch_module = try prepareDispatchJobFromTensorJob(allocator, changed_tensor_module, .{});
690 const changed_memory_module = try prepareMemoryJobFromDispatchJob(allocator, changed_dispatch_module, .{});
691 var changed_kernel = try prepareKernelJobFromMemoryJobWithRun(
692 allocator,
693 changed_memory_module,
694 .{},
695 );
696 defer changed_kernel.deinit();
697
698 try testing.expectEqual(first_kernel.plan_fingerprint, second_kernel.plan_fingerprint);
699 try testing.expect(first_kernel.plan_fingerprint != changed_kernel.plan_fingerprint);
700 }
701
702 test "kernel preparation lowers fused elementwise benchmark chain" {
703 const allocator = testing.allocator;
704
705 const semantic_module = try buildSemanticAddMulSubMaxModule(allocator, "kernel_product_add_mul_sub_max", &.{16});
706 const contract_module = try prepareContractJobFromSemanticModule(allocator, semantic_module, .{});
707 const tensor_module = try prepareTensorJobFromContractJob(allocator, contract_module, .{});
708 const dispatch_module = try prepareDispatchJobFromTensorJob(allocator, tensor_module, .{});
709 const memory_module = try prepareMemoryJobFromDispatchJob(allocator, dispatch_module, .{});
710 const kernel_module = try prepareKernelJobFromMemoryJob(allocator, memory_module, .{});
711 defer kernel_module.deinit();
712
713 var pass_ctx = kernel_module.passContext();
714 defer pass_ctx.deinit();
715 const kernel_plan = try kernelization.getKernelizationAnalysis(&pass_ctx, kernel_module.choir_module);
716 try testing.expectEqual(@as(usize, 1), kernel_plan.kernelCount());
717 const kernel = kernel_plan.kernels.items[0];
718 try testing.expect(std.mem.indexOf(u8, kernel.entry_name, "add") != null);
719 try testing.expect(std.mem.indexOf(u8, kernel.entry_name, "mul") != null);
720 try testing.expect(std.mem.indexOf(u8, kernel.entry_name, "sub") != null);
721 try testing.expect(std.mem.indexOf(u8, kernel.entry_name, "max") != null);
722 }
723
724 test "kernel preparation handles einsum operand-local reductions" {
725 const allocator = testing.allocator;
726
727 var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
728 defer builder.deinit();
729 const lhs_ty = try builder.tensor(.f32, &.{ 16, 16 });
730 const rhs_ty = try builder.tensor(.f32, &.{ 16, 16 });
731 const out_ty = try builder.tensor(.f32, &.{ 16, 16 });
732 var fb = try builder.beginFunction("kernel_product_einsum_local_reduction", &.{ lhs_ty, rhs_ty }, &.{out_ty});
733 const out = try fb.einsum(&.{ fb.parameter(0), fb.parameter(1) }, out_ty, "ab,cd->ac");
734 try fb.return_(&.{out});
735 try fb.finish();
736 const semantic_module = try builder.finish();
737
738 const contract_module = try prepareContractJobFromSemanticModule(allocator, semantic_module, .{});
739 const tensor_module = try prepareTensorJobFromContractJob(allocator, contract_module, .{});
740 const dispatch_module = try prepareDispatchJobFromTensorJob(allocator, tensor_module, .{});
741 const memory_module = try prepareMemoryJobFromDispatchJob(allocator, dispatch_module, .{});
742 const kernel_module = try prepareKernelJobFromMemoryJob(allocator, memory_module, .{});
743 defer kernel_module.deinit();
744
745 try kernel_module.verify();
746 var pass_ctx = kernel_module.passContext();
747 defer pass_ctx.deinit();
748 const kernel_plan = try kernelization.getKernelizationAnalysis(&pass_ctx, kernel_module.choir_module);
749 try testing.expectEqual(@as(usize, 3), kernel_plan.kernelCount());
750 }
751
752 test "kernel preparation handles three operand einsum contractions" {
753 const allocator = testing.allocator;
754
755 var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
756 defer builder.deinit();
757 const q_ty = try builder.tensor(.f32, &.{ 2, 16, 16 });
758 const k_ty = try builder.tensor(.f32, &.{ 2, 16, 16 });
759 const v_ty = try builder.tensor(.f32, &.{ 2, 16, 16 });
760 const out_ty = try builder.tensor(.f32, &.{ 2, 16, 16 });
761 var fb = try builder.beginFunction("kernel_product_einsum_three_operand", &.{ q_ty, k_ty, v_ty }, &.{out_ty});
762 const out = try fb.einsum(&.{ fb.parameter(0), fb.parameter(1), fb.parameter(2) }, out_ty, "bqd,bkd,bkv->bqv");
763 try fb.return_(&.{out});
764 try fb.finish();
765 const semantic_module = try builder.finish();
766
767 const contract_module = try prepareContractJobFromSemanticModule(allocator, semantic_module, .{});
768 const tensor_module = try prepareTensorJobFromContractJob(allocator, contract_module, .{});
769 const dispatch_module = try prepareDispatchJobFromTensorJob(allocator, tensor_module, .{});
770 const memory_module = try prepareMemoryJobFromDispatchJob(allocator, dispatch_module, .{});
771 const kernel_module = try prepareKernelJobFromMemoryJob(allocator, memory_module, .{});
772 defer kernel_module.deinit();
773
774 try kernel_module.verify();
775 var pass_ctx = kernel_module.passContext();
776 defer pass_ctx.deinit();
777 const kernel_plan = try kernelization.getKernelizationAnalysis(&pass_ctx, kernel_module.choir_module);
778 try testing.expect(kernel_plan.kernelCount() > 0);
779 }
780
781 test "prepareTargetJobFromKernelJob materializes target job" {
782 const allocator = testing.allocator;
783
784 const semantic_module = try buildSemanticAddMulModule(allocator, "target_product_add_mul", &.{4});
785 const contract_module = try prepareContractJobFromSemanticModule(allocator, semantic_module, .{});
786 const tensor_module = try prepareTensorJobFromContractJob(allocator, contract_module, .{});
787 const dispatch_module = try prepareDispatchJobFromTensorJob(allocator, tensor_module, .{});
788 const memory_module = try prepareMemoryJobFromDispatchJob(allocator, dispatch_module, .{});
789 const kernel_module = try prepareKernelJobFromMemoryJob(allocator, memory_module, .{});
790 const target_module = try prepareTargetJobFromKernelJob(allocator, kernel_module, .{});
791 defer target_module.deinit();
792
793 try target_module.verify();
794 try testing.expect(target_module.choir_module != target_module.kernel_module.choir_module);
795 try testing.expect(target_module.kernel_module.memory_module.dispatch_module.tensor_module.analysis_cache.entries.count() > 0);
796 try testing.expect(target_module.analysis_cache.entries.count() > 0);
797 try testing.expectEqual(@as(usize, 2), target_module.kernelizationProduct().kernelCount());
798
799 var pass_ctx = target_module.passContext();
800 defer pass_ctx.deinit();
801 const legal_plan = try backend_legalization.getBackendLegalizationAnalysis(&pass_ctx, target_module.choir_module);
802 try testing.expectEqual(@as(usize, 2), legal_plan.kernelCount());
803 try testing.expect(legal_plan.isLegal());
804 }
805
806 test "prepareBackendJobFromDispatchPreparationResult continues prepared pipeline" {
807 const allocator = testing.allocator;
808
809 var dispatch_result = try prepareDispatchResultForContinuationTest(allocator, "dispatch_continuation_add_mul");
810 var dispatch_owned = true;
811 errdefer if (dispatch_owned) dispatch_result.deinit();
812
813 const dispatch_ns = dispatch_result.elapsed_ns;
814 const dispatch_stats = dispatch_result.stats;
815 const dispatch_fingerprint = dispatch_result.plan_fingerprint;
816 dispatch_owned = false;
817 var prepared = try prepareBackendJobFromDispatchPreparationResult(
818 allocator,
819 dispatch_result,
820 .{},
821 11,
822 .{ .pass_runs = 101, .analysis_hits = 102 },
823 13,
824 .{ .pass_runs = 103, .analysis_misses = 104 },
825 17,
826 null,
827 0,
828 );
829 defer prepared.deinit();
830
831 try testing.expectEqual(@as(u64, 11), prepared.run.contract_ns);
832 try testing.expectEqual(@as(u64, 13), prepared.run.tensor_ns);
833 try testing.expectEqual(dispatch_ns, prepared.run.dispatch_ns);
834 try testing.expectEqual(@as(u64, 101), prepared.run.contract_stats.pass_runs);
835 try testing.expectEqual(@as(u64, 102), prepared.run.contract_stats.analysis_hits);
836 try testing.expectEqual(@as(u64, 103), prepared.run.tensor_stats.pass_runs);
837 try testing.expectEqual(@as(u64, 104), prepared.run.tensor_stats.analysis_misses);
838 try testing.expectEqual(dispatch_stats.pass_runs, prepared.run.dispatch_stats.pass_runs);
839 try testing.expectEqual(dispatch_stats.analysis_misses, prepared.run.dispatch_stats.analysis_misses);
840 try testing.expectEqual(@as(u64, memory_pass_count), prepared.run.memory_stats.pass_runs);
841 try testing.expectEqual(@as(u64, kernel_pass_count), prepared.run.kernel_stats.pass_runs);
842 try testing.expectEqual(@as(u64, target_pass_count), prepared.run.target_stats.pass_runs);
843 try testing.expectEqual(@as(u64, 17), prepared.run.initial_choir_ops);
844 try testing.expectEqual(dispatch_fingerprint, prepared.run.dispatch_fingerprint);
845 try expectPreparedRunFingerprints(&prepared);
846 try testing.expectEqual(@as(usize, 2), try prepared.generatedKernelCount());
847 }
848
849 test "prepareBackendJobFromMemoryPreparationResult continues prepared pipeline" {
850 const allocator = testing.allocator;
851
852 var memory_result = try prepareMemoryResultForContinuationTest(allocator, "memory_continuation_add_mul");
853 var memory_owned = true;
854 errdefer if (memory_owned) memory_result.deinit();
855
856 const memory_ns = memory_result.elapsed_ns;
857 const memory_stats = memory_result.stats;
858 const memory_fingerprint = memory_result.plan_fingerprint;
859 memory_owned = false;
860 var prepared = try prepareBackendJobFromMemoryPreparationResult(
861 allocator,
862 memory_result,
863 .{},
864 19,
865 .{ .pass_runs = 107 },
866 23,
867 .{ .pass_runs = 109 },
868 29,
869 .{ .pass_runs = 113, .analysis_hits = 114 },
870 127,
871 31,
872 null,
873 0,
874 );
875 defer prepared.deinit();
876
877 try testing.expectEqual(@as(u64, 19), prepared.run.contract_ns);
878 try testing.expectEqual(@as(u64, 23), prepared.run.tensor_ns);
879 try testing.expectEqual(@as(u64, 29), prepared.run.dispatch_ns);
880 try testing.expectEqual(memory_ns, prepared.run.memory_ns);
881 try testing.expectEqual(@as(u64, 107), prepared.run.contract_stats.pass_runs);
882 try testing.expectEqual(@as(u64, 109), prepared.run.tensor_stats.pass_runs);
883 try testing.expectEqual(@as(u64, 113), prepared.run.dispatch_stats.pass_runs);
884 try testing.expectEqual(@as(u64, 114), prepared.run.dispatch_stats.analysis_hits);
885 try testing.expectEqual(memory_stats.pass_runs, prepared.run.memory_stats.pass_runs);
886 try testing.expectEqual(memory_stats.analysis_misses, prepared.run.memory_stats.analysis_misses);
887 try testing.expectEqual(@as(u64, kernel_pass_count), prepared.run.kernel_stats.pass_runs);
888 try testing.expectEqual(@as(u64, target_pass_count), prepared.run.target_stats.pass_runs);
889 try testing.expectEqual(@as(u64, 31), prepared.run.initial_choir_ops);
890 try testing.expectEqual(@as(u64, 127), prepared.run.dispatch_fingerprint);
891 try testing.expectEqual(memory_fingerprint, prepared.run.memory_fingerprint);
892 try expectPreparedRunFingerprints(&prepared);
893 try testing.expectEqual(@as(usize, 2), try prepared.generatedKernelCount());
894 }
895
896 test "prepareBackendJobFromKernelPreparationResult continues prepared pipeline" {
897 const allocator = testing.allocator;
898
899 var kernel_result = try prepareKernelResultForContinuationTest(allocator, "kernel_continuation_add_mul");
900 var kernel_owned = true;
901 errdefer if (kernel_owned) kernel_result.deinit();
902
903 const kernel_ns = kernel_result.elapsed_ns;
904 const kernel_stats = kernel_result.stats;
905 const kernel_fingerprint = kernel_result.plan_fingerprint;
906 kernel_owned = false;
907 var prepared = try prepareBackendJobFromKernelPreparationResult(
908 allocator,
909 kernel_result,
910 .{},
911 37,
912 .{ .pass_runs = 127 },
913 41,
914 .{ .pass_runs = 131 },
915 43,
916 .{ .pass_runs = 137 },
917 149,
918 47,
919 .{ .pass_runs = 139, .analysis_misses = 140 },
920 151,
921 53,
922 null,
923 0,
924 );
925 defer prepared.deinit();
926
927 try testing.expectEqual(@as(u64, 37), prepared.run.contract_ns);
928 try testing.expectEqual(@as(u64, 41), prepared.run.tensor_ns);
929 try testing.expectEqual(@as(u64, 43), prepared.run.dispatch_ns);
930 try testing.expectEqual(@as(u64, 47), prepared.run.memory_ns);
931 try testing.expectEqual(kernel_ns, prepared.run.kernel_ns);
932 try testing.expectEqual(@as(u64, 139), prepared.run.memory_stats.pass_runs);
933 try testing.expectEqual(@as(u64, 140), prepared.run.memory_stats.analysis_misses);
934 try testing.expectEqual(kernel_stats.pass_runs, prepared.run.kernel_stats.pass_runs);
935 try testing.expectEqual(kernel_stats.analysis_misses, prepared.run.kernel_stats.analysis_misses);
936 try testing.expectEqual(@as(u64, target_pass_count), prepared.run.target_stats.pass_runs);
937 try testing.expectEqual(@as(u64, 53), prepared.run.initial_choir_ops);
938 try testing.expectEqual(@as(u64, 149), prepared.run.dispatch_fingerprint);
939 try testing.expectEqual(@as(u64, 151), prepared.run.memory_fingerprint);
940 try testing.expectEqual(kernel_fingerprint, prepared.run.kernel_fingerprint);
941 try expectPreparedRunFingerprints(&prepared);
942 try testing.expectEqual(@as(usize, 2), try prepared.generatedKernelCount());
943 }
944
945 test "prepareBackendJobFromTargetPreparationResult owns prepared product" {
946 const allocator = testing.allocator;
947
948 var target_result = try prepareTargetResultForContinuationTest(allocator, "target_continuation_add_mul");
949 var target_owned = true;
950 errdefer if (target_owned) target_result.deinit();
951
952 const target_ns = target_result.elapsed_ns;
953 const target_stats = target_result.stats;
954 const target_fingerprint = target_result.module.fingerprint();
955 target_owned = false;
956 var prepared = try prepareBackendJobFromTargetPreparationResult(
957 allocator,
958 target_result,
959 .{},
960 59,
961 .{ .pass_runs = 149 },
962 61,
963 .{ .pass_runs = 151 },
964 67,
965 .{ .pass_runs = 157 },
966 173,
967 71,
968 .{ .pass_runs = 163 },
969 179,
970 73,
971 .{ .pass_runs = 167, .analysis_hits = 168 },
972 181,
973 79,
974 null,
975 0,
976 );
977 defer prepared.deinit();
978
979 try testing.expectEqual(@as(u64, 59), prepared.run.contract_ns);
980 try testing.expectEqual(@as(u64, 61), prepared.run.tensor_ns);
981 try testing.expectEqual(@as(u64, 67), prepared.run.dispatch_ns);
982 try testing.expectEqual(@as(u64, 71), prepared.run.memory_ns);
983 try testing.expectEqual(@as(u64, 73), prepared.run.kernel_ns);
984 try testing.expectEqual(target_ns, prepared.run.target_ns);
985 try testing.expectEqual(@as(u64, 167), prepared.run.kernel_stats.pass_runs);
986 try testing.expectEqual(@as(u64, 168), prepared.run.kernel_stats.analysis_hits);
987 try testing.expectEqual(target_stats.pass_runs, prepared.run.target_stats.pass_runs);
988 try testing.expectEqual(target_stats.analysis_misses, prepared.run.target_stats.analysis_misses);
989 try testing.expectEqual(@as(u64, 79), prepared.run.initial_choir_ops);
990 try testing.expectEqual(@as(u64, 173), prepared.run.dispatch_fingerprint);
991 try testing.expectEqual(@as(u64, 179), prepared.run.memory_fingerprint);
992 try testing.expectEqual(@as(u64, 181), prepared.run.kernel_fingerprint);
993 try testing.expectEqual(target_fingerprint, prepared.run.target_fingerprint);
994 try expectPreparedRunFingerprints(&prepared);
995 try testing.expectEqual(@as(usize, 2), try prepared.generatedKernelCount());
996 }
997
998 test "prepareBackendJobFromSemanticModule owns its transient pipeline and diagnostic summaries" {
999 const allocator = testing.allocator;
1000
1001 const target = try BackendTargetProfile.init(.{
1002 .identity = .{
1003 .backend = .cuda,
1004 .family = .nvidia_cuda,
1005 },
1006 .dtypes = gpu.DTypeSet.init(&.{ .f32, .i32 }),
1007 .artifact_formats = gpu.ArtifactFormatSet.init(&.{.cuda_ptx}),
1008 }, .cuda, .cuda_ptx);
1009
1010 const module = try buildSemanticAddModule(allocator, "prepared_product_add4", &.{4});
1011 var prepared = try prepareBackendJobFromSemanticModule(allocator, module, .{ .target_profile = target });
1012 defer prepared.deinit();
1013 try expectPreparedJobState(&prepared, target);
1014 try expectPreparedJobProgram(&prepared, target);
1015 try expectPreparedJobDiagnostics(&prepared);
1016 }
1017
1018 fn expectPreparedJobStamps(prepared: *const BackendPreparedJob) !void {
1019 const stamps = try prepared.productStamps();
1020 try expectProductStamp(stamps.semantic.?, semantic.product_name, prepared.run.semantic_fingerprint.?);
1021 try expectProductStamp(stamps.contract, contract.product_name, prepared.run.contract_fingerprint);
1022 try expectProductStamp(stamps.tensor, tensor.product_name, prepared.run.tensor_fingerprint);
1023 try expectProductStamp(stamps.dispatch, dispatch.product_name, prepared.run.dispatch_fingerprint);
1024 try expectProductStamp(stamps.memory, memory_product.product_name, prepared.run.memory_fingerprint);
1025 try expectProductStamp(stamps.kernel, kernel_product.product_name, prepared.run.kernel_fingerprint);
1026 try expectProductStamp(stamps.target, target_product.product_name, prepared.run.target_fingerprint);
1027 }
1028
1029 fn expectPreparedJobState(prepared: *BackendPreparedJob, target: BackendTargetProfile) !void {
1030 const allocator = testing.allocator;
1031 const target_module = try prepared.targetModule();
1032
1033 const read_target = target_profile.readBackendTargetProfile(prepared.choir_module) orelse return error.MissingTargetProfile;
1034 try testing.expectEqual(target.backend_kind, read_target.backend_kind);
1035 try testing.expectEqual(target.artifact_format, read_target.artifact_format);
1036 try testing.expectEqual(target.math_tier, read_target.math_tier);
1037 try testing.expectEqual(target.dtype_bits, read_target.dtype_bits);
1038 try testing.expect(prepared.choir_module == target_module.choir_module);
1039 try testing.expect(target_module.choir_module != target_module.kernel_module.choir_module);
1040 const kernel_stage_target = target_profile.readBackendTargetProfile(target_module.kernel_module.choir_module) orelse return error.MissingTargetProfile;
1041 try testing.expectEqual(target.artifact_format, kernel_stage_target.artifact_format);
1042 try testing.expectEqual(target.math_tier, kernel_stage_target.math_tier);
1043 try testing.expect(target_module.kernel_module.fingerprint() != 0);
1044 try testing.expectEqual(
1045 try choir.operationFingerprint(allocator, target_module.kernel_module.choir_module),
1046 target_module.kernel_module.fingerprint(),
1047 );
1048 try testing.expectEqual(target.backend_kind, prepared.run.target_profile.?.backend_kind);
1049 try testing.expectEqual(@as(u64, contract_pass_count), prepared.run.contract_stats.pass_runs);
1050 try testing.expectEqual(@as(u64, tensor_pass_count), prepared.run.tensor_stats.pass_runs);
1051 try testing.expectEqual(@as(u64, dispatch_pass_count), prepared.run.dispatch_stats.pass_runs);
1052 try testing.expectEqual(@as(u64, memory_pass_count), prepared.run.memory_stats.pass_runs);
1053 try testing.expectEqual(@as(u64, kernel_pass_count), prepared.run.kernel_stats.pass_runs);
1054 try testing.expectEqual(@as(u64, target_pass_count), prepared.run.target_stats.pass_runs);
1055 try testing.expect(prepared.run.contract_fingerprint != 0);
1056 try testing.expect(prepared.run.tensor_fingerprint != 0);
1057 try testing.expect(prepared.run.dispatch_fingerprint != 0);
1058 try testing.expect(prepared.run.memory_fingerprint != 0);
1059 try testing.expect(prepared.run.kernel_fingerprint != 0);
1060 try testing.expect(prepared.run.target_fingerprint != 0);
1061 try testing.expect(prepared.run.semantic_fingerprint != null);
1062 try expectPreparedJobStamps(prepared);
1063 try testing.expectEqual(try choir.operationFingerprint(allocator, prepared.choir_module), prepared.run.target_fingerprint);
1064 try testing.expect(prepared.run.dispatch_stats.analysis_misses > 0);
1065 try testing.expect(prepared.run.memory_stats.analysis_misses > 0);
1066 try testing.expect(prepared.run.kernel_stats.analysis_misses > 0);
1067 try testing.expect(prepared.run.target_stats.analysis_misses > 0);
1068 try testing.expect(prepared.run.initial_choir_ops > 0);
1069 try testing.expect(prepared.run.final_choir_ops > 0);
1070 try testing.expect(prepared.ctx.isDialectLoaded("memref"));
1071 try testing.expect(prepared.ctx.isDialectLoaded("scf"));
1072 }
1073
1074 fn expectPreparedJobProgram(prepared: *BackendPreparedJob, target: BackendTargetProfile) !void {
1075 const allocator = testing.allocator;
1076 const body = prepared.choir_module.getRegion(0).?.getEntryBlock().?;
1077 try testing.expect(ir.inspection.functionByNameInBlock(body, "prepared_product_add4") != null);
1078 try testing.expect(!hasFunctionNameSuffix(body, "lowered"));
1079
1080 var pass_ctx = prepared.passContext();
1081 defer pass_ctx.deinit();
1082 const kernel_plan = try prepared.kernelizationProduct();
1083 const legal_plan = try backend_legalization.getBackendLegalizationAnalysis(&pass_ctx, prepared.choir_module);
1084
1085 try testing.expectEqual(@as(usize, 1), kernel_plan.kernelCount());
1086 try testing.expectEqual(@as(usize, 1), legal_plan.kernelCount());
1087 try testing.expect(legal_plan.isLegal());
1088 try testing.expectEqual(target.dtype_bits, legal_plan.target.?.dtype_bits);
1089
1090 try testing.expectEqual(@as(usize, 1), try prepared.generatedKernelCount());
1091 const summary = try prepared.generatedKernelSummary(0);
1092 try testing.expect(summary.entry_name.len > 0);
1093 try testing.expectEqualStrings(kernel_plan.kernels.items[0].entry_name, summary.entry_name);
1094 try testing.expectEqual(kernel_plan.kernels.items[0].argument_count, summary.argument_count);
1095 try testing.expectEqual(kernel_plan.kernels.items[0].body_fingerprint, summary.body_fingerprint);
1096 try testing.expectEqual(kernelization.GeneratedScheduleKind.flat, summary.schedule.kind);
1097 try testing.expect(summary.launch_geometry != null);
1098 const program = try prepared.generatedKernelProgram(0);
1099 try testing.expect(prepared.ctx != program.kernelModule().context);
1100 try ir.verifyOperation(program.kernelModule(), ir.verify.default_options);
1101 const launch = try program.launch();
1102 try testing.expectEqual(launch.block[0], summary.launch_geometry.?.threadgroup[0]);
1103 const work_summary = try prepared.generatedKernelSummaryForWork(summary.work_item_id);
1104 try testing.expectEqualStrings(summary.entry_name, work_summary.entry_name);
1105 try testing.expectEqual(summary.body_fingerprint, work_summary.body_fingerprint);
1106 const work_program = try prepared.generatedKernelProgramForWork(summary.work_item_id);
1107 const work_launch = try work_program.launch();
1108 try testing.expectEqual(launch.block[0], work_launch.block[0]);
1109 var summaries = try prepared.copyGeneratedKernelSummaries(allocator);
1110 defer summaries.deinit();
1111 try testing.expectEqual(@as(usize, 1), summaries.len());
1112 const copied_summary = try summaries.summaryForWork(summary.work_item_id);
1113 try testing.expectEqualStrings(summary.entry_name, copied_summary.entry_name);
1114 try testing.expectEqual(summary.body_fingerprint, copied_summary.body_fingerprint);
1115 try testing.expectError(error.InvalidIndex, prepared.generatedKernelSummary(1));
1116 try testing.expectError(error.MissingKernelization, prepared.generatedKernelSummaryForWork(std.math.maxInt(usize)));
1117 try testing.expectError(error.MissingKernelization, summaries.summaryForWork(std.math.maxInt(usize)));
1118 }
1119
1120 fn expectPreparedJobDiagnostics(prepared: *BackendPreparedJob) !void {
1121 const target_module = try prepared.targetModule();
1122 prepared.run.semantic_fingerprint.? +%= 13;
1123 prepared.run.contract_fingerprint +%= 17;
1124 prepared.run.tensor_fingerprint +%= 19;
1125 prepared.run.dispatch_fingerprint +%= 23;
1126 prepared.run.memory_fingerprint +%= 29;
1127 prepared.run.kernel_fingerprint +%= 31;
1128 prepared.run.target_fingerprint +%= 37;
1129
1130 const live_contract_fingerprint = target_module.kernel_module.memory_module.dispatch_module.tensor_module.contract_module.fingerprint();
1131 try testing.expect(live_contract_fingerprint != prepared.run.contract_fingerprint);
1132 try testing.expect(target_module.fingerprint() != prepared.run.target_fingerprint);
1133
1134 try expectPreparedJobStamps(prepared);
1135 }
1136
1137 test "backend preparation target profile constrains backend legalization dtypes" {
1138 const allocator = testing.allocator;
1139
1140 const target = try BackendTargetProfile.init(.{
1141 .identity = .{
1142 .backend = .vulkan,
1143 .family = .vulkan,
1144 },
1145 .dtypes = gpu.DTypeSet.init(&.{.i32}),
1146 .artifact_formats = gpu.ArtifactFormatSet.init(&.{.vulkan_spirv}),
1147 }, .vulkan, .vulkan_spirv);
1148
1149 const module = try buildSemanticAddModule(allocator, "target_profile_rejects_f32", &.{4});
1150 var prepared = try prepareBackendJobFromSemanticModule(allocator, module, .{ .target_profile = target });
1151 defer prepared.deinit();
1152
1153 var pass_ctx = prepared.passContext();
1154 defer pass_ctx.deinit();
1155 const legal_plan = try backend_legalization.getBackendLegalizationAnalysis(&pass_ctx, prepared.choir_module);
1156 const status = legal_plan.firstIllegalStatus() orelse return error.MissingIllegalKernel;
1157
1158 try testing.expect(!legal_plan.isLegal());
1159 try testing.expect(!legal_plan.hasPipelineFailure());
1160 try testing.expectEqual(backend_legalization.BackendKernelStatus.unsupported_dtype, status);
1161 try testing.expectEqual(error.CapabilityMismatch, backend_legalization.backendKernelStatusError(status));
1162 }
1163
1164 test "backend preparation pipelines register through Choir extensions" {
1165 const allocator = testing.allocator;
1166
1167 var ctx = try buildBackendPreparationContext(allocator, ir.Context.Limits.testing);
1168 defer ctx.deinit(allocator);
1169
1170 var extension_registry = choir.extensions.ExtensionRegistry.init(allocator);
1171 defer extension_registry.deinit();
1172 try extension_registry.registerPackage(&ctx, accy_choir_package_extension);
1173
1174 var pipeline_registry = passes.PipelineRegistry.init(allocator);
1175 defer pipeline_registry.deinit();
1176 try extension_registry.registerPipelinesTo(&pipeline_registry);
1177 try testing.expect(pipeline_registry.lookup(contract_pipeline_name) != null);
1178 try testing.expect(pipeline_registry.lookup(tensor_pipeline_name) != null);
1179 try testing.expect(pipeline_registry.lookup(dispatch_pipeline_name) != null);
1180 try testing.expect(pipeline_registry.lookup(memory_pipeline_name) != null);
1181 try testing.expect(pipeline_registry.lookup(kernel_pipeline_name) != null);
1182 try testing.expect(pipeline_registry.lookup(target_pipeline_name) != null);
1183
1184 var contract_pm = passes.PassManager.init(allocator);
1185 defer contract_pm.deinit();
1186 try pipeline_registry.addPipelineTo(contract_pipeline_name, &contract_pm);
1187 try testing.expectEqual(contract_pass_count, contract_pm.root.pipeline.items.len);
1188
1189 var tensor_pm = passes.PassManager.init(allocator);
1190 defer tensor_pm.deinit();
1191 try pipeline_registry.addPipelineTo(tensor_pipeline_name, &tensor_pm);
1192 try expectTensorStagePipeline(&tensor_pm);
1193
1194 var dispatch_pm = passes.PassManager.init(allocator);
1195 defer dispatch_pm.deinit();
1196 try pipeline_registry.addPipelineTo(dispatch_pipeline_name, &dispatch_pm);
1197 try expectDispatchStagePipeline(&dispatch_pm);
1198
1199 var memory_pm = passes.PassManager.init(allocator);
1200 defer memory_pm.deinit();
1201 try pipeline_registry.addPipelineTo(memory_pipeline_name, &memory_pm);
1202 try expectMemoryStagePipeline(&memory_pm);
1203
1204 var kernel_pm = passes.PassManager.init(allocator);
1205 defer kernel_pm.deinit();
1206 try pipeline_registry.addPipelineTo(kernel_pipeline_name, &kernel_pm);
1207 try expectKernelStagePipeline(&kernel_pm);
1208
1209 var target_pm = passes.PassManager.init(allocator);
1210 defer target_pm.deinit();
1211 try pipeline_registry.addPipelineTo(target_pipeline_name, &target_pm);
1212 try expectTargetStagePipeline(&target_pm);
1213 }
1214
1215 test "individual Accy Choir passes register through Choir extensions" {
1216 const allocator = testing.allocator;
1217
1218 var ctx = try buildBackendPreparationContext(allocator, ir.Context.Limits.testing);
1219 defer ctx.deinit(allocator);
1220
1221 var extension_registry = choir.extensions.ExtensionRegistry.init(allocator);
1222 defer extension_registry.deinit();
1223 try extension_registry.registerPackage(&ctx, accy_choir_package_extension);
1224
1225 var pass_registry = passes.PassRegistry.init(allocator);
1226 defer pass_registry.deinit();
1227 try extension_registry.registerPassEntriesTo(&pass_registry);
1228
1229 try testing.expectEqual(
1230 @as(usize, accy_choir_pass_registrations.len),
1231 pass_registry.passes.items.len,
1232 );
1233 inline for (accy_choir_pass_registrations) |registration| {
1234 try testing.expect(pass_registry.lookupPass(registration.name) != null);
1235 }
1236
1237 var pm = passes.PassManager.init(allocator);
1238 defer pm.deinit();
1239 try passes.parsePassPipeline(
1240 &pass_registry,
1241 canonicalization_pass_name ++ "," ++ saturation_pass_name,
1242 &pm,
1243 );
1244 try testing.expectEqual(@as(usize, 2), pm.root.pipeline.items.len);
1245
1246 const text = try passes.formatPassManagerPipelineAlloc(allocator, &pm);
1247 defer allocator.free(text);
1248 try testing.expectEqualStrings(canonicalization_pass_name ++ "," ++ saturation_pass_name, text);
1249 }
1250
1251 test "backend preparation stage table defines registries and profiling products" {
1252 try testing.expectEqual(contract_stages.len, contract_pass_count);
1253 try testing.expectEqual(tensor_stages.len, tensor_pass_count);
1254 try testing.expectEqual(dispatch_stages.len, dispatch_pass_count);
1255 try testing.expectEqual(memory_stages.len, memory_pass_count);
1256 try testing.expectEqual(kernel_stages.len, kernel_pass_count);
1257 try testing.expectEqual(target_stages.len, target_pass_count);
1258 try testing.expectEqual(contract_stages.len + tensor_stages.len + dispatch_stages.len + memory_stages.len + kernel_stages.len + target_stages.len, accy_choir_pass_registrations.len);
1259 try testing.expectEqualStrings(kernelization.kernelization_analysis_descriptor.name, kernel_stages[1].analysis_name.?);
1260 try testing.expectEqual(contract_stages.len, contract_pass_plan.steps.len);
1261 try testing.expectEqual(tensor_stages.len, tensor_pass_plan.steps.len);
1262 try testing.expectEqual(dispatch_stages.len, dispatch_pass_plan.steps.len);
1263 try testing.expectEqual(memory_stages.len, memory_pass_plan.steps.len);
1264 try testing.expectEqual(kernel_stages.len, kernel_pass_plan.steps.len);
1265 try testing.expectEqual(target_stages.len, target_pass_plan.steps.len);
1266
1267 var contract_analysis_index: usize = 0;
1268 inline for (contract_stages, 0..) |stage, index| {
1269 try testing.expectEqualStrings(stage.name, contractPassName(index));
1270 try testing.expectEqualStrings(stage.name, accy_choir_pass_registrations[index].name);
1271 try testing.expectEqualStrings(stage.description, accy_choir_pass_registrations[index].description);
1272 try testing.expectEqualStrings(stage.pass.name, accy_choir_pass_registrations[index].pass.name);
1273
1274 if (stage.analysis_name) |name| {
1275 try testing.expectEqualStrings(name, contractAnalysisName(contract_analysis_index));
1276 contract_analysis_index += 1;
1277 }
1278 }
1279
1280 var tensor_analysis_index: usize = 0;
1281 inline for (tensor_stages, 0..) |stage, index| {
1282 const registration_index = contract_stages.len + index;
1283 try testing.expectEqualStrings(stage.name, tensorPassName(index));
1284 try testing.expectEqualStrings(stage.name, accy_choir_pass_registrations[registration_index].name);
1285 try testing.expectEqualStrings(stage.description, accy_choir_pass_registrations[registration_index].description);
1286 try testing.expectEqualStrings(stage.pass.name, accy_choir_pass_registrations[registration_index].pass.name);
1287 try testing.expectEqual(stage.options.len, accy_choir_pass_registrations[registration_index].options.len);
1288 try testing.expectEqual(stage.build_with_options != null, accy_choir_pass_registrations[registration_index].build_with_options != null);
1289
1290 if (stage.analysis_name) |name| {
1291 try testing.expectEqualStrings(name, tensorAnalysisName(tensor_analysis_index));
1292 tensor_analysis_index += 1;
1293 }
1294 }
1295
1296 var dispatch_analysis_index: usize = 0;
1297 inline for (dispatch_stages, 0..) |stage, index| {
1298 const registration_index = contract_stages.len + tensor_stages.len + index;
1299 try testing.expectEqualStrings(stage.name, dispatchPassName(index));
1300 try testing.expectEqualStrings(stage.name, accy_choir_pass_registrations[registration_index].name);
1301 try testing.expectEqualStrings(stage.description, accy_choir_pass_registrations[registration_index].description);
1302 try testing.expectEqualStrings(stage.pass.name, accy_choir_pass_registrations[registration_index].pass.name);
1303
1304 if (stage.analysis_name) |name| {
1305 try testing.expectEqualStrings(name, dispatchAnalysisName(dispatch_analysis_index));
1306 dispatch_analysis_index += 1;
1307 }
1308 }
1309
1310 var memory_analysis_index: usize = 0;
1311 inline for (memory_stages, 0..) |stage, index| {
1312 const registration_index = contract_stages.len + tensor_stages.len + dispatch_stages.len + index;
1313 try testing.expectEqualStrings(stage.name, memoryPassName(index));
1314 try testing.expectEqualStrings(stage.name, accy_choir_pass_registrations[registration_index].name);
1315 try testing.expectEqualStrings(stage.description, accy_choir_pass_registrations[registration_index].description);
1316 try testing.expectEqualStrings(stage.pass.name, accy_choir_pass_registrations[registration_index].pass.name);
1317
1318 if (stage.analysis_name) |name| {
1319 try testing.expectEqualStrings(name, memoryAnalysisName(memory_analysis_index));
1320 memory_analysis_index += 1;
1321 }
1322 }
1323
1324 var kernel_analysis_index: usize = 0;
1325 inline for (kernel_stages, 0..) |stage, index| {
1326 const registration_index = contract_stages.len + tensor_stages.len + dispatch_stages.len + memory_stages.len + index;
1327 try testing.expectEqualStrings(stage.name, kernelPassName(index));
1328 try testing.expectEqualStrings(stage.name, accy_choir_pass_registrations[registration_index].name);
1329 try testing.expectEqualStrings(stage.description, accy_choir_pass_registrations[registration_index].description);
1330 try testing.expectEqualStrings(stage.pass.name, accy_choir_pass_registrations[registration_index].pass.name);
1331
1332 if (stage.analysis_name) |name| {
1333 try testing.expectEqualStrings(name, kernelAnalysisName(kernel_analysis_index));
1334 kernel_analysis_index += 1;
1335 }
1336 }
1337
1338 var target_analysis_index: usize = 0;
1339 inline for (target_stages, 0..) |stage, index| {
1340 const registration_index = contract_stages.len + tensor_stages.len + dispatch_stages.len + memory_stages.len + kernel_stages.len + index;
1341 try testing.expectEqualStrings(stage.name, targetPassName(index));
1342 try testing.expectEqualStrings(stage.name, accy_choir_pass_registrations[registration_index].name);
1343 try testing.expectEqualStrings(stage.description, accy_choir_pass_registrations[registration_index].description);
1344 try testing.expectEqualStrings(stage.pass.name, accy_choir_pass_registrations[registration_index].pass.name);
1345
1346 if (stage.analysis_name) |name| {
1347 try testing.expectEqualStrings(name, targetAnalysisName(target_analysis_index));
1348 target_analysis_index += 1;
1349 }
1350 }
1351
1352 try testing.expectEqual(contract_analysis_index, contract_analysis_count);
1353 try testing.expectEqual(tensor_analysis_index, tensor_analysis_count);
1354 try testing.expectEqual(dispatch_analysis_index, dispatch_analysis_count);
1355 try testing.expectEqual(memory_analysis_index, memory_analysis_count);
1356 try testing.expectEqual(kernel_analysis_index, kernel_analysis_count);
1357 try testing.expectEqual(target_analysis_index, target_analysis_count);
1358 }
1359
1360 test "backend preparation stage plans materialize through Choir registry" {
1361 const allocator = testing.allocator;
1362
1363 var ctx = try buildBackendPreparationContext(allocator, ir.Context.Limits.testing);
1364 defer ctx.deinit(allocator);
1365
1366 var extension_registry = choir.extensions.ExtensionRegistry.init(allocator);
1367 defer extension_registry.deinit();
1368 try extension_registry.registerPackage(&ctx, accy_choir_package_extension);
1369
1370 var pass_registry = passes.PassRegistry.init(allocator);
1371 defer pass_registry.deinit();
1372 try extension_registry.registerPassEntriesTo(&pass_registry);
1373
1374 try expectStagePlanText(allocator, contract_stages, contract_pass_plan);
1375 try expectStagePlanText(allocator, tensor_stages, tensor_pass_plan);
1376 try expectStagePlanText(allocator, dispatch_stages, dispatch_pass_plan);
1377 try expectStagePlanText(allocator, memory_stages, memory_pass_plan);
1378 try expectStagePlanText(allocator, kernel_stages, kernel_pass_plan);
1379 try expectStagePlanText(allocator, target_stages, target_pass_plan);
1380
1381 var contract_pm = passes.PassManager.init(allocator);
1382 defer contract_pm.deinit();
1383 try contract_pass_plan.addTo(&pass_registry, &contract_pm);
1384 try testing.expectEqual(contract_pass_count, contract_pm.root.pipeline.items.len);
1385
1386 var tensor_pm = passes.PassManager.init(allocator);
1387 defer tensor_pm.deinit();
1388 try tensor_pass_plan.addTo(&pass_registry, &tensor_pm);
1389 try expectTensorStagePipeline(&tensor_pm);
1390
1391 var dispatch_pm = passes.PassManager.init(allocator);
1392 defer dispatch_pm.deinit();
1393 try dispatch_pass_plan.addTo(&pass_registry, &dispatch_pm);
1394 try expectDispatchStagePipeline(&dispatch_pm);
1395
1396 var memory_pm = passes.PassManager.init(allocator);
1397 defer memory_pm.deinit();
1398 try memory_pass_plan.addTo(&pass_registry, &memory_pm);
1399 try expectMemoryStagePipeline(&memory_pm);
1400
1401 var kernel_pm = passes.PassManager.init(allocator);
1402 defer kernel_pm.deinit();
1403 try kernel_pass_plan.addTo(&pass_registry, &kernel_pm);
1404 try expectKernelStagePipeline(&kernel_pm);
1405
1406 var target_pm = passes.PassManager.init(allocator);
1407 defer target_pm.deinit();
1408 try target_pass_plan.addTo(&pass_registry, &target_pm);
1409 try expectTargetStagePipeline(&target_pm);
1410 }
1411
1412 test "backend preparation run options expose product controls only" {
1413 inline for (
1414 @typeInfo(BackendPreparationRunOptions).@"struct".field_names,
1415 @typeInfo(BackendPreparationRunOptions).@"struct".field_types,
1416 @typeInfo(BackendPreparationRunOptions).@"struct".field_attrs,
1417 ) |field_name, field_name_type, field_name_attrs| {
1418 const field = .{ .name = field_name, .type = field_name_type, .attrs = field_name_attrs };
1419 try testing.expect(!std.mem.eql(u8, field.name, "pass_manager"));
1420 try testing.expect(field.type != passes.PassManagerRunOptions);
1421 }
1422 }
1423
1424 test "target pipeline materializes from textual Choir registry" {
1425 const allocator = testing.allocator;
1426
1427 var ctx = try buildBackendPreparationContext(allocator, ir.Context.Limits.testing);
1428 defer ctx.deinit(allocator);
1429
1430 var extension_registry = choir.extensions.ExtensionRegistry.init(allocator);
1431 defer extension_registry.deinit();
1432 try extension_registry.registerPackage(&ctx, accy_choir_package_extension);
1433
1434 var pass_registry = passes.PassRegistry.init(allocator);
1435 defer pass_registry.deinit();
1436 try extension_registry.registerPassEntriesTo(&pass_registry);
1437
1438 var pm = passes.PassManager.init(allocator);
1439 defer pm.deinit();
1440 try passes.parsePassPipeline(&pass_registry, target_pipeline_name, &pm);
1441 try expectTargetStagePipeline(&pm);
1442
1443 const text = try passes.formatPassManagerPipelineAlloc(allocator, &pm);
1444 defer allocator.free(text);
1445 try expectTargetStageText(text);
1446 }
1447
1448 test "contract pipeline materializes from textual Choir registry" {
1449 const allocator = testing.allocator;
1450
1451 var ctx = try buildBackendPreparationContext(allocator, ir.Context.Limits.testing);
1452 defer ctx.deinit(allocator);
1453
1454 var extension_registry = choir.extensions.ExtensionRegistry.init(allocator);
1455 defer extension_registry.deinit();
1456 try extension_registry.registerPackage(&ctx, accy_choir_package_extension);
1457
1458 var pass_registry = passes.PassRegistry.init(allocator);
1459 defer pass_registry.deinit();
1460 try extension_registry.registerPassEntriesTo(&pass_registry);
1461
1462 var pm = passes.PassManager.init(allocator);
1463 defer pm.deinit();
1464 try passes.parsePassPipeline(&pass_registry, contract_pipeline_name, &pm);
1465
1466 try testing.expectEqual(contract_pass_count, pm.root.pipeline.items.len);
1467 inline for (contract_stages, 0..) |stage, index| {
1468 switch (pm.root.pipeline.items[index]) {
1469 .pass => |pass| try testing.expectEqualStrings(stage.name, pass.name),
1470 .nested => try testing.expect(false),
1471 }
1472 }
1473 }
1474
1475 test "tensor pipeline materializes from textual Choir registry" {
1476 const allocator = testing.allocator;
1477
1478 var ctx = try buildBackendPreparationContext(allocator, ir.Context.Limits.testing);
1479 defer ctx.deinit(allocator);
1480
1481 var extension_registry = choir.extensions.ExtensionRegistry.init(allocator);
1482 defer extension_registry.deinit();
1483 try extension_registry.registerPackage(&ctx, accy_choir_package_extension);
1484
1485 var pass_registry = passes.PassRegistry.init(allocator);
1486 defer pass_registry.deinit();
1487 try extension_registry.registerPassEntriesTo(&pass_registry);
1488
1489 var pm = passes.PassManager.init(allocator);
1490 defer pm.deinit();
1491 try passes.parsePassPipeline(&pass_registry, tensor_pipeline_name, &pm);
1492
1493 try expectTensorStagePipeline(&pm);
1494 }
1495
1496 test "tensor lowering pass options materialize from textual Choir registry" {
1497 const allocator = testing.allocator;
1498
1499 var ctx = try buildBackendPreparationContext(allocator, ir.Context.Limits.testing);
1500 defer ctx.deinit(allocator);
1501
1502 var extension_registry = choir.extensions.ExtensionRegistry.init(allocator);
1503 defer extension_registry.deinit();
1504 try extension_registry.registerPackage(&ctx, accy_choir_package_extension);
1505
1506 var pass_registry = passes.PassRegistry.init(allocator);
1507 defer pass_registry.deinit();
1508 try extension_registry.registerPassEntriesTo(&pass_registry);
1509
1510 const pipeline_text =
1511 activation_lowering_pass_name ++ "{kernel-library=enabled}," ++
1512 einsum_lowering_pass_name ++ "{strategy=beam,exact-state-limit=64,beam-width=8,auto-beam-width=16,kernel-library=enabled}," ++
1513 indexing_lowering_pass_name ++ "{kernel-library=enabled,gather-thread-blocks=4,scatter-thread-blocks=5}," ++
1514 loss_lowering_pass_name ++ "{kernel-library=enabled,row-sparse-cross-entropy-thread-blocks=4}";
1515
1516 var pm = passes.PassManager.init(allocator);
1517 defer pm.deinit();
1518 try passes.parsePassPipeline(&pass_registry, pipeline_text, &pm);
1519
1520 try expectTensorStagePipeline(&pm);
1521 const text = try passes.formatPassManagerPipelineAlloc(allocator, &pm);
1522 defer allocator.free(text);
1523 try testing.expectEqualStrings(pipeline_text, text);
1524 }
1525
1526 test "dispatch pipeline materializes from textual Choir registry" {
1527 const allocator = testing.allocator;
1528
1529 var ctx = try buildBackendPreparationContext(allocator, ir.Context.Limits.testing);
1530 defer ctx.deinit(allocator);
1531
1532 var extension_registry = choir.extensions.ExtensionRegistry.init(allocator);
1533 defer extension_registry.deinit();
1534 try extension_registry.registerPackage(&ctx, accy_choir_package_extension);
1535
1536 var pass_registry = passes.PassRegistry.init(allocator);
1537 defer pass_registry.deinit();
1538 try extension_registry.registerPassEntriesTo(&pass_registry);
1539
1540 var pm = passes.PassManager.init(allocator);
1541 defer pm.deinit();
1542 try passes.parsePassPipeline(&pass_registry, dispatch_pipeline_name, &pm);
1543
1544 try expectDispatchStagePipeline(&pm);
1545 }
1546
1547 test "memory pipeline materializes from textual Choir registry" {
1548 const allocator = testing.allocator;
1549
1550 var ctx = try buildBackendPreparationContext(allocator, ir.Context.Limits.testing);
1551 defer ctx.deinit(allocator);
1552
1553 var extension_registry = choir.extensions.ExtensionRegistry.init(allocator);
1554 defer extension_registry.deinit();
1555 try extension_registry.registerPackage(&ctx, accy_choir_package_extension);
1556
1557 var pass_registry = passes.PassRegistry.init(allocator);
1558 defer pass_registry.deinit();
1559 try extension_registry.registerPassEntriesTo(&pass_registry);
1560
1561 var pm = passes.PassManager.init(allocator);
1562 defer pm.deinit();
1563 try passes.parsePassPipeline(&pass_registry, memory_pipeline_name, &pm);
1564
1565 try expectMemoryStagePipeline(&pm);
1566 }
1567
1568 test "kernel pipeline materializes from textual Choir registry" {
1569 const allocator = testing.allocator;
1570
1571 var ctx = try buildBackendPreparationContext(allocator, ir.Context.Limits.testing);
1572 defer ctx.deinit(allocator);
1573
1574 var extension_registry = choir.extensions.ExtensionRegistry.init(allocator);
1575 defer extension_registry.deinit();
1576 try extension_registry.registerPackage(&ctx, accy_choir_package_extension);
1577
1578 var pass_registry = passes.PassRegistry.init(allocator);
1579 defer pass_registry.deinit();
1580 try extension_registry.registerPassEntriesTo(&pass_registry);
1581
1582 var pm = passes.PassManager.init(allocator);
1583 defer pm.deinit();
1584 try passes.parsePassPipeline(&pass_registry, kernel_pipeline_name, &pm);
1585
1586 try expectKernelStagePipeline(&pm);
1587 }
1588
1589 test "runTargetPipeline exposes a direct pipeline helper" {
1590 const allocator = testing.allocator;
1591
1592 const module = try buildSemanticAddModule(allocator, "pass_add3", &.{3});
1593 defer module.deinit();
1594
1595 const choir_mod = module.choir_module;
1596 const ctx = module.context();
1597 try runTargetPipeline(allocator, choir_mod, ctx);
1598
1599 const body = choir_mod.getRegion(0).?.getEntryBlock().?;
1600 try testing.expect(ir.inspection.functionByNameInBlock(body, "pass_add3") != null);
1601 try testing.expect(!hasFunctionNameSuffix(body, "lowered"));
1602 try expectPreparedKernelPlan(allocator, choir_mod, ctx, 1);
1603 }
1604
1605 test "runTargetPipelineWithOptions forwards pass manager options" {
1606 const allocator = testing.allocator;
1607
1608 const module = try buildSemanticAddMulModule(allocator, "pass_threaded_add_mul", &.{4});
1609 defer module.deinit();
1610
1611 const choir_mod = module.choir_module;
1612 const ctx = module.context();
1613
1614 var worker_gpa = std.heap.DebugAllocator(.{}){};
1615 defer {
1616 const status = worker_gpa.deinit();
1617 testing.expect(status == .ok) catch @panic("target pipeline worker allocator leaked allocations");
1618 }
1619
1620 try runTargetPipelineWithOptions(allocator, choir_mod, ctx, .{
1621 .max_threads = 2,
1622 .worker_allocator = worker_gpa.allocator(),
1623 });
1624
1625 const body = choir_mod.getRegion(0).?.getEntryBlock().?;
1626 try testing.expect(ir.inspection.functionByNameInBlock(body, "pass_threaded_add_mul") != null);
1627 try testing.expect(!hasFunctionNameSuffix(body, "lowered"));
1628 try expectPreparedKernelPlan(allocator, choir_mod, ctx, 2);
1629 }