lib/accy/src/preparation/fusion/pass.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const alloc_fixed = @import("alloc_fixed");
3 const choir = @import("choir");
4 const accy_choir = @import("../../choir/root.zig");
5 const dialect_mod = accy_choir.dialect;
6 const shape_analysis = @import("../shape/root.zig");
7
8 const ir = choir.ir;
9 const passes = choir.passes;
10 const work = passes.pass.work;
11
12 pub const fusion_plan_analysis_name = "accy-choir-fusion-plan";
13 pub const fusion_planning_pass_name = "accy-choir-plan-fusion";
14 pub const fusion_planning_pass_description =
15 "Plan block-local Accy Choir elementwise fusion clusters";
16
17 pub const ClaimKind = enum {
18 recompute,
19 slot,
20 };
21
22 pub const ClaimMap = std.AutoHashMap(*ir.Operation, ClaimKind);
23
24 pub const FusionClusterKind = accy_choir.record.dispatch.FusionClusterKind;
25
26 pub const FusionCluster = struct {
27 ops: []*ir.Operation,
28 kind: FusionClusterKind = .elementwise,
29
30 pub fn len(self: FusionCluster) usize {
31 return self.ops.len;
32 }
33
34 pub fn root(self: FusionCluster) ?*ir.Operation {
35 if (self.ops.len == 0) return null;
36 return self.ops[self.ops.len - 1];
37 }
38
39 fn deinit(self: *FusionCluster, allocator: std.mem.Allocator) void {
40 allocator.free(self.ops);
41 self.* = undefined;
42 }
43 };
44
45 pub const FusionPlanAnalysis = struct {
46 allocator: std.mem.Allocator,
47 clusters: std.ArrayListUnmanaged(FusionCluster),
48 elided: std.ArrayListUnmanaged(*ir.Operation),
49 fused_op_count: usize = 0,
50 max_cluster_len: usize = 0,
51
52 pub fn init(allocator: std.mem.Allocator) FusionPlanAnalysis {
53 return .{
54 .allocator = allocator,
55 .clusters = .empty,
56 .elided = .empty,
57 };
58 }
59
60 pub fn deinit(self: *FusionPlanAnalysis) void {
61 for (self.clusters.items) |*cluster| {
62 cluster.deinit(self.allocator);
63 }
64 self.clusters.deinit(self.allocator);
65 self.elided.deinit(self.allocator);
66 self.* = undefined;
67 }
68
69 pub fn clusterCount(self: FusionPlanAnalysis) usize {
70 return self.clusters.items.len;
71 }
72
73 fn addCluster(self: *FusionPlanAnalysis, ops: []const *ir.Operation) !void {
74 try self.addClusterOfKind(ops, .elementwise);
75 }
76
77 fn addClusterOfKind(self: *FusionPlanAnalysis, ops: []const *ir.Operation, kind: FusionClusterKind) !void {
78 if (ops.len < 2) return;
79 try self.addOwnedCluster(ops, kind);
80 }
81
82 fn addReductionCluster(self: *FusionPlanAnalysis, ops: []const *ir.Operation) !void {
83 if (ops.len == 0) return;
84 try self.addOwnedCluster(ops, .reduction_input);
85 }
86
87 fn addOwnedCluster(self: *FusionPlanAnalysis, ops: []const *ir.Operation, kind: FusionClusterKind) !void {
88 const owned = try self.allocator.alloc(*ir.Operation, ops.len);
89 @memcpy(owned, ops);
90 errdefer self.allocator.free(owned);
91 try self.clusters.append(self.allocator, .{ .ops = owned, .kind = kind });
92 self.fused_op_count += owned.len;
93 self.max_cluster_len = @max(self.max_cluster_len, owned.len);
94 }
95 };
96
97 const FusionWork = struct {
98 input: work.Census = .{},
99 scratch: u64 = 0,
100 copies: u64 = 0,
101 rounds: u64 = 1,
102
103 fn inspect(op: *ir.Operation) !FusionWork {
104 var counts = FusionWork{ .input = try work.Census.inspect(op) };
105 _ = try op.walk(.{ .order = .pre_order }, &counts, visit);
106 return counts;
107 }
108
109 fn visit(self: *FusionWork, op: *ir.Operation) !ir.Operation.WalkResult {
110 for (op.regions.items) |*region| {
111 var blocks = region.getBlocks();
112 while (blocks.next()) |block| {
113 var count: u64 = 0;
114 var current: ?*ir.Operation = @ptrCast(@alignCast(block.operations.head));
115 while (current) |child| : (current = child.next_op) {
116 count = try work.add(count, 1);
117 }
118 try self.includeBlock(count);
119 }
120 }
121 return .advance;
122 }
123
124 fn includeBlock(self: *FusionWork, count: u64) !void {
125 if (count == 0) return;
126 const maps = try work.add(try work.multiply(6, count), 4);
127 const lists = try work.add(try work.multiply(4, count), 4);
128 const map_bytes = try work.multiply(
129 maps,
130 try work.hashMapGrowth(*ir.Operation, void, count),
131 );
132 const list_bytes = try work.multiply(lists, try work.arrayListGrowth(*ir.Operation, count));
133 self.scratch = try work.add(self.scratch, try work.add(map_bytes, list_bytes));
134 const output_count = try work.multiply(7, count);
135 const output_bytes = try work.add(try work.multiply(count, @sizeOf(*ir.Operation)), 8);
136 self.copies = try work.add(self.copies, try work.multiply(output_count, output_bytes));
137 const rounds = try work.multiply(try work.multiply(7, count), try work.add(count, 1));
138 self.rounds = try work.add(self.rounds, rounds);
139 }
140
141 fn bounds(self: FusionWork) !work.Bounds {
142 var retained: u64 = @sizeOf(FusionPlanAnalysis) + @alignOf(FusionPlanAnalysis);
143 retained = try work.add(retained, self.copies);
144 const clusters = try work.arrayListGrowth(
145 FusionCluster,
146 try work.multiply(7, self.input.operations),
147 );
148 retained = try work.add(retained, clusters);
149 retained = try work.add(
150 retained,
151 try work.arrayListGrowth(*ir.Operation, self.input.operations),
152 );
153 const temporary = try work.add(
154 self.scratch,
155 try work.hashMapGrowth(*ir.Operation, ClaimKind, self.input.operations),
156 );
157 const bytes = try work.add(retained, temporary);
158 if (bytes > std.math.maxInt(usize)) return error.WorkOverflow;
159 const units = try work.add(try work.add(self.input.atoms, self.input.input_bytes), 1);
160 const slots = try work.hashMapCapacity(
161 try work.add(self.input.operations, self.input.values),
162 );
163 const probes = try work.add(slots, units);
164 const visits = try work.multiply(try work.multiply(self.rounds, units), probes);
165 return .{
166 .work = .{
167 .input_bytes = self.input.input_bytes,
168 .structural_visits = try work.multiply(64, visits),
169 .analysis_computations = 1,
170 .allocation_capacity = bytes,
171 },
172 .workspace = bytes,
173 .retained_storage = bytes,
174 };
175 }
176 };
177
178 fn fusionAnalysisWork(input: work.Input) !work.Bounds {
179 return (try FusionWork.inspect(input.operation)).bounds();
180 }
181
182 fn fusionPassWork(_: work.Input) !work.Bounds {
183 return .{ .work = .{ .structural_visits = 1 } };
184 }
185
186 pub const fusion_plan_analysis_descriptor = passes.AnalysisDescriptor{
187 .id = passes.analysisId(fusion_plan_analysis_name),
188 .name = fusion_plan_analysis_name,
189 .work_contract = .{
190 .identity = .{ .name = fusion_plan_analysis_name, .version = 1 },
191 .estimate = fusionAnalysisWork,
192 },
193 };
194
195 pub fn getFusionPlanAnalysis(
196 pass_ctx: *passes.PassContext,
197 op: *ir.Operation,
198 ) !*FusionPlanAnalysis {
199 const ptr = try pass_ctx.getAnalysis(
200 op,
201 &fusion_plan_analysis_descriptor,
202 computeFusionPlanAnalysis,
203 cleanupFusionPlanAnalysis,
204 );
205 return @ptrCast(@alignCast(ptr));
206 }
207
208 pub fn fusionPlanningPass() passes.Pass {
209 return .{
210 .name = fusion_planning_pass_name,
211 .description = fusion_planning_pass_description,
212 .run_fn = runFusionPlanningPass,
213 .work_contract = .{
214 .identity = .{ .name = fusion_planning_pass_name, .version = 1 },
215 .estimate = fusionPassWork,
216 },
217 };
218 }
219
220 fn runFusionPlanningPass(pass_ctx: *passes.PassContext) passes.PassResult {
221 _ = getFusionPlanAnalysis(pass_ctx, pass_ctx.op) catch return .failure;
222 pass_ctx.preserveAllAnalyses();
223 return .success;
224 }
225
226 fn computeFusionPlanAnalysis(
227 pass_ctx: *passes.PassContext,
228 op: *ir.Operation,
229 ) anyerror!*anyopaque {
230 const shapes = try shape_analysis.getShapeLayoutAnalysis(pass_ctx, op);
231 const analysis = try pass_ctx.allocator.create(FusionPlanAnalysis);
232 analysis.* = FusionPlanAnalysis.init(pass_ctx.allocator);
233 errdefer {
234 analysis.deinit();
235 pass_ctx.allocator.destroy(analysis);
236 }
237
238 var claimed = ClaimMap.init(pass_ctx.allocator);
239 defer claimed.deinit();
240
241 try collectFlashAttentionPlansInRegions(pass_ctx.allocator, op.regions.items, shapes, analysis, &claimed);
242 try collectRowPipelinePlansInRegions(pass_ctx.allocator, op.regions.items, shapes, analysis, &claimed);
243 try collectDotEpiloguePlansInRegions(pass_ctx.allocator, op.regions.items, shapes, analysis, &claimed);
244 try collectReductionInputPlansInRegions(pass_ctx.allocator, op.regions.items, shapes, analysis, &claimed);
245 try collectIterateElisionsInRegions(pass_ctx.allocator, op.regions.items, shapes, analysis, &claimed);
246 try collectFusionPlansInRegions(pass_ctx.allocator, op.regions.items, shapes, analysis, &claimed);
247 return @ptrCast(analysis);
248 }
249
250 fn cleanupFusionPlanAnalysis(ptr: *anyopaque, allocator: std.mem.Allocator) void {
251 const analysis: *FusionPlanAnalysis = @ptrCast(@alignCast(ptr));
252 analysis.deinit();
253 allocator.destroy(analysis);
254 }
255
256 fn collectFusionPlansInRegions(
257 allocator: std.mem.Allocator,
258 regions: []ir.Region,
259 shapes: *const shape_analysis.ShapeLayoutAnalysis,
260 analysis: *FusionPlanAnalysis,
261 claimed: *ClaimMap,
262 ) anyerror!void {
263 for (regions) |*region| {
264 var block_iter = region.getBlocks();
265 while (block_iter.next()) |block| {
266 var current: ?*ir.Operation = @ptrCast(@alignCast(block.operations.head));
267 while (current) |op| {
268 const next = op.next_op;
269
270 if (op.regions.items.len > 0 and !regionsAreOpaque(op)) {
271 try collectFusionPlansInRegions(allocator, op.regions.items, shapes, analysis, claimed);
272 }
273
274 if (isFusableElementwiseOp(op) and !claimed.contains(op) and !hasFusableOutputUser(op, shapes, claimed)) {
275 try claimClusterRootedAt(allocator, op, shapes, analysis, claimed);
276 }
277
278 current = next;
279 }
280 var leftover: ?*ir.Operation = @ptrCast(@alignCast(block.operations.tail));
281 while (leftover) |op| {
282 const previous = op.prev_op;
283 if (isFusableElementwiseOp(op) and !claimed.contains(op)) {
284 try claimClusterRootedAt(allocator, op, shapes, analysis, claimed);
285 }
286 leftover = previous;
287 }
288 }
289 }
290 }
291
292 fn claimClusterRootedAt(
293 allocator: std.mem.Allocator,
294 root: *ir.Operation,
295 shapes: *const shape_analysis.ShapeLayoutAnalysis,
296 analysis: *FusionPlanAnalysis,
297 claimed: *ClaimMap,
298 ) anyerror!void {
299 var cluster: std.ArrayListUnmanaged(*ir.Operation) = .empty;
300 defer cluster.deinit(allocator);
301 var visiting = std.AutoHashMap(*ir.Operation, void).init(allocator);
302 defer visiting.deinit();
303
304 try collectProducerDAG(allocator, root, shapes, claimed, &visiting, &cluster);
305 try retainClosedProducers(allocator, root, &cluster, shapes, claimed);
306 if (cluster.items.len >= 2) {
307 try analysis.addCluster(cluster.items);
308 for (cluster.items) |node| {
309 try claimed.put(node, .recompute);
310 }
311 }
312 }
313
314 fn hasFusableOutputUser(
315 op: *ir.Operation,
316 shapes: *const shape_analysis.ShapeLayoutAnalysis,
317 claimed: *ClaimMap,
318 ) bool {
319 const result = op.getResult(0) orelse return false;
320 var use = result.first_use;
321 while (use) |current_use| : (use = current_use.next_use) {
322 const user: *ir.Operation = @ptrCast(@alignCast(current_use.owner));
323 if (user.getBlock() != op.getBlock()) continue;
324 if (claimed.contains(user)) continue;
325 if (!isFusableElementwiseOp(user)) continue;
326 if (!compatibleProducerConsumer(op, user, shapes)) continue;
327 return true;
328 }
329 return false;
330 }
331
332 fn collectProducerDAG(
333 allocator: std.mem.Allocator,
334 op: *ir.Operation,
335 shapes: *const shape_analysis.ShapeLayoutAnalysis,
336 claimed: *ClaimMap,
337 visiting: *std.AutoHashMap(*ir.Operation, void),
338 cluster: *std.ArrayListUnmanaged(*ir.Operation),
339 ) anyerror!void {
340 if (claimed.contains(op)) return;
341 if (visiting.contains(op)) return;
342 try visiting.put(op, {});
343
344 if (isSeeThroughShapeOp(op)) {
345 const source = op.getOperand(0) orelse return;
346 if (source.getDefiningOp()) |source_def_any| {
347 const source_def: *ir.Operation = @ptrCast(@alignCast(source_def_any));
348 if (isSeeThroughShapeOp(source_def) and !claimed.contains(source_def) and source_def.getBlock() == op.getBlock()) {
349 try collectProducerDAG(allocator, source_def, shapes, claimed, visiting, cluster);
350 }
351 }
352 try cluster.append(allocator, op);
353 return;
354 }
355
356 for (op.getOperandValues()) |operand| {
357 const producer = fusableProducerForOperand(operand, op, shapes, claimed) orelse continue;
358 try collectProducerDAG(allocator, producer, shapes, claimed, visiting, cluster);
359 }
360
361 try cluster.append(allocator, op);
362 }
363
364 fn retainClosedProducers(
365 allocator: std.mem.Allocator,
366 root: *ir.Operation,
367 cluster: *std.ArrayListUnmanaged(*ir.Operation),
368 shapes: *const shape_analysis.ShapeLayoutAnalysis,
369 claimed: *const ClaimMap,
370 ) anyerror!void {
371 var members = std.AutoHashMap(*ir.Operation, void).init(allocator);
372 defer members.deinit();
373 for (cluster.items) |op| try members.put(op, {});
374
375 var changed = true;
376 while (changed) {
377 changed = false;
378 var index: usize = 0;
379 while (index < cluster.items.len) {
380 const op = cluster.items[index];
381 if (op != root and !allUsesRecomputableOrInside(op, &members, shapes, claimed)) {
382 _ = members.remove(op);
383 _ = cluster.orderedRemove(index);
384 changed = true;
385 continue;
386 }
387 index += 1;
388 }
389 }
390 }
391
392 fn allUsesRecomputableOrInside(
393 op: *ir.Operation,
394 members: *const std.AutoHashMap(*ir.Operation, void),
395 shapes: *const shape_analysis.ShapeLayoutAnalysis,
396 claimed: *const ClaimMap,
397 ) bool {
398 const result = op.getResult(0) orelse return false;
399 var use = result.first_use;
400 while (use) |current_use| : (use = current_use.next_use) {
401 const user: *ir.Operation = @ptrCast(@alignCast(current_use.owner));
402 if (members.contains(user)) continue;
403 if (userCanRecompute(op, user, shapes, claimed)) continue;
404 return false;
405 }
406 return true;
407 }
408
409 fn userCanRecompute(
410 op: *ir.Operation,
411 user: *ir.Operation,
412 shapes: *const shape_analysis.ShapeLayoutAnalysis,
413 claimed: *const ClaimMap,
414 ) bool {
415 if (claimed.get(user)) |kind| {
416 if (kind == .slot) return false;
417 }
418 if (user.getBlock() != op.getBlock()) {
419 return iterateRegionCanRecompute(op, user, shapes, claimed);
420 }
421 if (isName(user.name.name, dialect_mod.AccyDialect.IterateOp.operation_name)) {
422 const op_result = op.getResult(0) orelse return false;
423 const user_result = user.getResult(0) orelse return false;
424 const op_info = shapes.get(op_result) orelse return false;
425 const user_info = shapes.get(user_result) orelse return false;
426 if (!op_info.hasStaticLayout() or !user_info.hasStaticLayout()) return false;
427 return op_info.element_count.? == user_info.element_count.?;
428 }
429 if (isFlatGatherOp(user)) {
430 const op_result = op.getResult(0) orelse return false;
431 const user_operands = user.getOperandValues();
432 if (user_operands[0] == op_result) return false;
433 return resultShapesMatch(op, user, shapes);
434 }
435 if (!isFusableElementwiseOp(user)) return false;
436 const op_result = op.getResult(0) orelse return false;
437 const user_result = user.getResult(0) orelse return false;
438 const op_info = shapes.get(op_result) orelse return false;
439 const user_info = shapes.get(user_result) orelse return false;
440 return sameStaticShape(op_info, user_info);
441 }
442
443 fn iterateRegionCanRecompute(
444 op: *ir.Operation,
445 user: *ir.Operation,
446 shapes: *const shape_analysis.ShapeLayoutAnalysis,
447 claimed: *const ClaimMap,
448 ) bool {
449 const user_block = user.getBlock() orelse return false;
450 const parent = user_block.getParentOperation() orelse return false;
451 if (!isName(parent.name.name, dialect_mod.AccyDialect.IterateOp.operation_name)) return false;
452 if (parent.getBlock() != op.getBlock()) return false;
453 if (claimed.contains(parent)) return false;
454 const op_result = op.getResult(0) orelse return false;
455 const parent_result = parent.getResult(0) orelse return false;
456 const op_info = shapes.get(op_result) orelse return false;
457 const parent_info = shapes.get(parent_result) orelse return false;
458 if (!op_info.hasStaticLayout() or !parent_info.hasStaticLayout()) return false;
459 return op_info.element_count.? == parent_info.element_count.?;
460 }
461
462 fn fusableProducerForOperand(
463 operand: *ir.Value,
464 consumer: *ir.Operation,
465 shapes: *const shape_analysis.ShapeLayoutAnalysis,
466 claimed: *ClaimMap,
467 ) ?*ir.Operation {
468 const def_any = operand.getDefiningOp() orelse return null;
469 const def_op: *ir.Operation = @ptrCast(@alignCast(def_any));
470 if (claimed.contains(def_op)) return null;
471 if (def_op.getBlock() != consumer.getBlock()) return null;
472 if (isFlatGatherOp(consumer)) {
473 const consumer_operands = consumer.getOperandValues();
474 if (consumer_operands[0] == operand) return null;
475 if (!isFusableElementwiseOp(def_op) and !isFlatGatherOp(def_op)) return null;
476 if (!resultShapesMatch(def_op, consumer, shapes)) return null;
477 return def_op;
478 }
479 if (isSeeThroughShapeOp(def_op)) {
480 if (!compatibleProducerConsumerShapes(def_op, consumer, shapes)) return null;
481 return def_op;
482 }
483 if (isFlatGatherOp(def_op)) {
484 if (!compatibleProducerConsumer(def_op, consumer, shapes)) return null;
485 return def_op;
486 }
487 if (!isFusableElementwiseOp(def_op)) return null;
488 if (!compatibleProducerConsumer(def_op, consumer, shapes)) return null;
489 return def_op;
490 }
491
492 pub fn isFlatGatherOp(op: *ir.Operation) bool {
493 if (!isName(op.name.name, dialect_mod.AccyDialect.GatherOp.operation_name)) return false;
494 if (op.getNumResults() != 1) return false;
495 const operands = op.getOperandValues();
496 if (operands.len != 2) return false;
497 const attr = op.getAttr("axis") orelse return false;
498 const int_attr = attr.cast(ir.Attribute.IntegerAttr) orelse return false;
499 if (int_attr.getValue() != 0) return false;
500 var arena_buffer: [256]u8 = undefined;
501 var arena = alloc_fixed.FixedBuffer.init(arena_buffer[0..]);
502 const src_type = dialect_mod.decodeTensorType(arena.allocator(), operands[0].type) catch return false;
503 return src_type.dims.len == 1;
504 }
505
506 fn resultShapesMatch(
507 producer: *ir.Operation,
508 consumer: *ir.Operation,
509 shapes: *const shape_analysis.ShapeLayoutAnalysis,
510 ) bool {
511 const producer_result = producer.getResult(0) orelse return false;
512 const consumer_result = consumer.getResult(0) orelse return false;
513 const producer_info = shapes.get(producer_result) orelse return false;
514 const consumer_info = shapes.get(consumer_result) orelse return false;
515 return sameStaticShape(producer_info, consumer_info);
516 }
517
518 fn collectIterateElisionsInRegions(
519 allocator: std.mem.Allocator,
520 regions: []ir.Region,
521 shapes: *const shape_analysis.ShapeLayoutAnalysis,
522 analysis: *FusionPlanAnalysis,
523 claimed: *ClaimMap,
524 ) anyerror!void {
525 for (regions) |*region| {
526 var block_iter = region.getBlocks();
527 while (block_iter.next()) |block| {
528 var current: ?*ir.Operation = @ptrCast(@alignCast(block.operations.head));
529 while (current) |op| {
530 const next = op.next_op;
531 if (op.regions.items.len > 0 and !regionsAreOpaque(op)) {
532 try collectIterateElisionsInRegions(allocator, op.regions.items, shapes, analysis, claimed);
533 }
534 current = next;
535 }
536 try collectBlockIterateElisions(allocator, block, shapes, analysis, claimed);
537 }
538 }
539 }
540
541 fn collectBlockIterateElisions(
542 allocator: std.mem.Allocator,
543 block: *ir.Block,
544 shapes: *const shape_analysis.ShapeLayoutAnalysis,
545 analysis: *FusionPlanAnalysis,
546 claimed: *ClaimMap,
547 ) anyerror!void {
548 var has_recompute_sink = false;
549 var scan: ?*ir.Operation = @ptrCast(@alignCast(block.operations.head));
550 while (scan) |op| : (scan = op.next_op) {
551 if (claimed.contains(op)) continue;
552 if (isName(op.name.name, dialect_mod.AccyDialect.IterateOp.operation_name) or isFlatGatherOp(op)) {
553 has_recompute_sink = true;
554 break;
555 }
556 }
557 if (!has_recompute_sink) return;
558
559 var candidates = std.AutoHashMap(*ir.Operation, void).init(allocator);
560 defer candidates.deinit();
561 var current: ?*ir.Operation = @ptrCast(@alignCast(block.operations.head));
562 while (current) |op| : (current = op.next_op) {
563 if (claimed.contains(op)) continue;
564 if (!isFusableElementwiseOp(op) and
565 !isName(op.name.name, dialect_mod.AccyDialect.IotaOp.operation_name) and
566 !isFlatGatherOp(op)) continue;
567 if (op.getNumResults() != 1) continue;
568 try candidates.put(op, {});
569 }
570 if (candidates.count() == 0) return;
571
572 var removals: std.ArrayListUnmanaged(*ir.Operation) = .empty;
573 defer removals.deinit(allocator);
574 var changed = true;
575 while (changed) {
576 changed = false;
577 removals.clearRetainingCapacity();
578 var candidate_iter = candidates.keyIterator();
579 while (candidate_iter.next()) |candidate| {
580 if (!iterateElisionUsesOk(candidate.*, &candidates, shapes, claimed)) {
581 try removals.append(allocator, candidate.*);
582 }
583 }
584 for (removals.items) |candidate| {
585 _ = candidates.remove(candidate);
586 changed = true;
587 }
588 }
589
590 var claim_iter = candidates.keyIterator();
591 while (claim_iter.next()) |candidate| {
592 try claimed.put(candidate.*, .recompute);
593 try analysis.elided.append(analysis.allocator, candidate.*);
594 }
595 }
596
597 fn iterateElisionUsesOk(
598 op: *ir.Operation,
599 candidates: *const std.AutoHashMap(*ir.Operation, void),
600 shapes: *const shape_analysis.ShapeLayoutAnalysis,
601 claimed: *const ClaimMap,
602 ) bool {
603 const result = op.getResult(0) orelse return false;
604 var use = result.first_use;
605 while (use) |current_use| : (use = current_use.next_use) {
606 const user: *ir.Operation = @ptrCast(@alignCast(current_use.owner));
607 if (candidates.contains(user)) continue;
608 if (userCanRecompute(op, user, shapes, claimed)) continue;
609 return false;
610 }
611 return true;
612 }
613
614 pub fn regionsAreOpaque(op: *ir.Operation) bool {
615 return isName(op.name.name, dialect_mod.AccyDialect.IterateOp.operation_name);
616 }
617
618 pub fn isSeeThroughShapeOp(op: *ir.Operation) bool {
619 if (op.getNumResults() != 1) return false;
620 if (isName(op.name.name, dialect_mod.AccyDialect.SliceOp.operation_name)) return true;
621 if (isName(op.name.name, dialect_mod.AccyDialect.PadOp.operation_name)) {
622 return padHasZeroInterior(op);
623 }
624 return false;
625 }
626
627 fn padHasZeroInterior(op: *ir.Operation) bool {
628 const attr = op.getAttrAs(ir.Attribute.DialectAttr, "interior") orelse return false;
629 if (attr.payload.len % @sizeOf(i64) != 0) return false;
630 var index: usize = 0;
631 while (index < attr.payload.len) : (index += @sizeOf(i64)) {
632 var dilation: i64 = undefined;
633 @memcpy(std.mem.asBytes(&dilation), attr.payload[index..][0..@sizeOf(i64)]);
634 if (dilation != 0) return false;
635 }
636 return true;
637 }
638
639 fn compatibleProducerConsumer(
640 producer: *ir.Operation,
641 consumer: *ir.Operation,
642 shapes: *const shape_analysis.ShapeLayoutAnalysis,
643 ) bool {
644 const producer_result = producer.getResult(0) orelse return false;
645 const consumer_result = consumer.getResult(0) orelse return false;
646 const producer_info = shapes.get(producer_result) orelse return false;
647 const consumer_info = shapes.get(consumer_result) orelse return false;
648 if (!sameStaticShape(producer_info, consumer_info)) return false;
649
650 for (consumer.getOperandValues()) |operand| {
651 const operand_info = shapes.get(operand) orelse return false;
652 if (!sameStaticShape(producer_info, operand_info)) return false;
653 }
654 return true;
655 }
656
657 fn sameStaticShape(lhs: shape_analysis.TensorInfo, rhs: shape_analysis.TensorInfo) bool {
658 if (!lhs.hasStaticLayout() or !rhs.hasStaticLayout()) return false;
659 if (lhs.element_count.? != rhs.element_count.?) return false;
660 return std.mem.eql(i64, lhs.dims, rhs.dims);
661 }
662
663 pub const flash_attention_dim = 64;
664 pub const flash_attention_min_seq = 128;
665
666 fn collectFlashAttentionPlansInRegions(
667 allocator: std.mem.Allocator,
668 regions: []ir.Region,
669 shapes: *const shape_analysis.ShapeLayoutAnalysis,
670 analysis: *FusionPlanAnalysis,
671 claimed: *ClaimMap,
672 ) anyerror!void {
673 for (regions) |*region| {
674 var block_iter = region.getBlocks();
675 while (block_iter.next()) |block| {
676 var current: ?*ir.Operation = @ptrCast(@alignCast(block.operations.head));
677 while (current) |op| {
678 const next = op.next_op;
679 if (op.regions.items.len > 0 and !regionsAreOpaque(op)) {
680 try collectFlashAttentionPlansInRegions(allocator, op.regions.items, shapes, analysis, claimed);
681 }
682 current = next;
683 }
684 try collectBlockFlashAttentionPlans(allocator, block, shapes, analysis, claimed);
685 }
686 }
687 }
688
689 fn collectBlockFlashAttentionPlans(
690 allocator: std.mem.Allocator,
691 block: *ir.Block,
692 shapes: *const shape_analysis.ShapeLayoutAnalysis,
693 analysis: *FusionPlanAnalysis,
694 claimed: *ClaimMap,
695 ) anyerror!void {
696 var current: ?*ir.Operation = @ptrCast(@alignCast(block.operations.head));
697 while (current) |dot2| : (current = dot2.next_op) {
698 if (!isName(dot2.name.name, dialect_mod.AccyDialect.DotGeneralOp.operation_name)) continue;
699 if (claimed.contains(dot2)) continue;
700 if (!dotIsCanonicalRank2(dot2)) continue;
701
702 const probs = dot2.getOperand(0) orelse continue;
703 const probs_any = probs.getDefiningOp() orelse continue;
704 const terminal: *ir.Operation = @ptrCast(@alignCast(probs_any));
705 if (terminal.getBlock() != block or claimed.contains(terminal)) continue;
706 if (!isFusableElementwiseOp(terminal)) continue;
707 if (!probs.hasOneUse()) continue;
708
709 const dot2_result = dot2.getResult(0) orelse continue;
710 const out_info = shapes.get(dot2_result) orelse continue;
711 if (!out_info.hasStaticLayout() or out_info.dims.len != 2) continue;
712 if (out_info.dims[1] != flash_attention_dim) continue;
713 const seq = out_info.dims[0];
714 if (seq < flash_attention_min_seq or @rem(seq, 64) != 0) continue;
715 if (out_info.dtype != .f32) continue;
716
717 var probe = RowPipelineProbe{ .members = std.AutoHashMap(*ir.Operation, void).init(allocator) };
718 defer probe.deinit();
719 probe.rows = seq;
720 probe.cols = seq;
721 if (!try collectRowPipelineMembers(terminal, block, shapes, claimed, &probe)) continue;
722 if (probe.reduce_count != 2) continue;
723 if (probe.members.count() > max_row_pipeline_ops) continue;
724 if (!flashPipeClosed(terminal, dot2, &probe)) continue;
725
726 const dot1 = flashScoresDot(block, shapes, claimed, &probe, seq) orelse continue;
727
728 var ops: std.ArrayListUnmanaged(*ir.Operation) = .empty;
729 defer ops.deinit(allocator);
730 try ops.append(allocator, dot1);
731 var walk: ?*ir.Operation = @ptrCast(@alignCast(block.operations.head));
732 while (walk) |member| : (walk = member.next_op) {
733 if (member == dot1 or member == dot2) continue;
734 if (probe.members.contains(member)) try ops.append(allocator, member);
735 }
736 try ops.append(allocator, dot2);
737
738 try analysis.addClusterOfKind(ops.items, .flash_attention);
739 for (ops.items) |member| try claimed.put(member, .slot);
740 }
741 }
742
743 fn dotIsCanonicalRank2(op: *ir.Operation) bool {
744 const lhs_contract = op.getAttrAs(ir.Attribute.DialectAttr, dialect_mod.AccyDialect.DotGeneralOp.dialectAttrName("lhs_contract")) orelse
745 op.getAttrAs(ir.Attribute.DialectAttr, "lhs_contract") orelse return false;
746 const rhs_contract = op.getAttrAs(ir.Attribute.DialectAttr, dialect_mod.AccyDialect.DotGeneralOp.dialectAttrName("rhs_contract")) orelse
747 op.getAttrAs(ir.Attribute.DialectAttr, "rhs_contract") orelse return false;
748 const lhs_batch = op.getAttrAs(ir.Attribute.DialectAttr, dialect_mod.AccyDialect.DotGeneralOp.dialectAttrName("lhs_batch")) orelse
749 op.getAttrAs(ir.Attribute.DialectAttr, "lhs_batch") orelse return false;
750 if (lhs_batch.payload.len != 0) return false;
751 if (lhs_contract.payload.len != @sizeOf(i64) or rhs_contract.payload.len != @sizeOf(i64)) return false;
752 var lhs_dim: i64 = undefined;
753 var rhs_dim: i64 = undefined;
754 @memcpy(std.mem.asBytes(&lhs_dim), lhs_contract.payload[0..@sizeOf(i64)]);
755 @memcpy(std.mem.asBytes(&rhs_dim), rhs_contract.payload[0..@sizeOf(i64)]);
756 return lhs_dim == 1 and rhs_dim == 0;
757 }
758
759 fn flashPipeClosed(terminal: *ir.Operation, dot2: *ir.Operation, probe: *const RowPipelineProbe) bool {
760 var member_iter = probe.members.keyIterator();
761 while (member_iter.next()) |member| {
762 const result = member.*.getResult(0) orelse return false;
763 var use = result.first_use;
764 while (use) |current_use| : (use = current_use.next_use) {
765 const user: *ir.Operation = @ptrCast(@alignCast(current_use.owner));
766 if (probe.members.contains(user)) continue;
767 if (member.* == terminal and user == dot2) continue;
768 return false;
769 }
770 }
771 return true;
772 }
773
774 fn flashScoresDot(
775 block: *ir.Block,
776 shapes: *const shape_analysis.ShapeLayoutAnalysis,
777 claimed: *const ClaimMap,
778 probe: *const RowPipelineProbe,
779 seq: i64,
780 ) ?*ir.Operation {
781 var scores_dot: ?*ir.Operation = null;
782 var member_iter = probe.members.keyIterator();
783 while (member_iter.next()) |member| {
784 for (member.*.getOperandValues()) |operand| {
785 const def_any = operand.getDefiningOp() orelse continue;
786 const def_op: *ir.Operation = @ptrCast(@alignCast(def_any));
787 if (probe.members.contains(def_op)) continue;
788 if (!isName(def_op.name.name, dialect_mod.AccyDialect.DotGeneralOp.operation_name)) continue;
789 if (def_op.getBlock() != block or claimed.contains(def_op)) continue;
790 if (!dotIsCanonicalRank2(def_op)) continue;
791 const result = def_op.getResult(0) orelse continue;
792 const info = shapes.get(result) orelse continue;
793 if (!info.hasStaticLayout() or info.dims.len != 2) continue;
794 if (info.dims[0] != seq or info.dims[1] != seq) continue;
795 var use = result.first_use;
796 var all_inside = true;
797 while (use) |current_use| : (use = current_use.next_use) {
798 const user: *ir.Operation = @ptrCast(@alignCast(current_use.owner));
799 if (!probe.members.contains(user)) {
800 all_inside = false;
801 break;
802 }
803 }
804 if (!all_inside) continue;
805 const lhs = def_op.getOperand(0) orelse continue;
806 const lhs_info = shapes.get(lhs) orelse continue;
807 if (!lhs_info.hasStaticLayout() or lhs_info.dims.len != 2) continue;
808 if (lhs_info.dims[1] != flash_attention_dim) continue;
809 if (scores_dot != null and scores_dot != def_op) return null;
810 scores_dot = def_op;
811 }
812 }
813 return scores_dot;
814 }
815
816 pub const row_pipeline_threads = 256;
817 pub const max_row_pipeline_ops = 24;
818 pub const max_row_pipeline_reduces = 4;
819 pub const max_row_pipeline_cols = 8192;
820
821 fn collectRowPipelinePlansInRegions(
822 allocator: std.mem.Allocator,
823 regions: []ir.Region,
824 shapes: *const shape_analysis.ShapeLayoutAnalysis,
825 analysis: *FusionPlanAnalysis,
826 claimed: *ClaimMap,
827 ) anyerror!void {
828 for (regions) |*region| {
829 var block_iter = region.getBlocks();
830 while (block_iter.next()) |block| {
831 var current: ?*ir.Operation = @ptrCast(@alignCast(block.operations.head));
832 while (current) |op| {
833 const next = op.next_op;
834 if (op.regions.items.len > 0 and !regionsAreOpaque(op)) {
835 try collectRowPipelinePlansInRegions(allocator, op.regions.items, shapes, analysis, claimed);
836 }
837 current = next;
838 }
839 try collectBlockRowPipelinePlans(allocator, block, shapes, analysis, claimed);
840 }
841 }
842 }
843
844 const RowPipelineProbe = struct {
845 members: std.AutoHashMap(*ir.Operation, void),
846 reduce_count: usize = 0,
847 rows: i64 = 0,
848 cols: i64 = 0,
849
850 fn deinit(self: *RowPipelineProbe) void {
851 self.members.deinit();
852 }
853 };
854
855 fn collectBlockRowPipelinePlans(
856 allocator: std.mem.Allocator,
857 block: *ir.Block,
858 shapes: *const shape_analysis.ShapeLayoutAnalysis,
859 analysis: *FusionPlanAnalysis,
860 claimed: *ClaimMap,
861 ) anyerror!void {
862 var current: ?*ir.Operation = @ptrCast(@alignCast(block.operations.head));
863 while (current) |op| : (current = op.next_op) {
864 if (!rowPipelineTerminalCandidate(op, shapes, claimed)) continue;
865
866 var probe = RowPipelineProbe{ .members = std.AutoHashMap(*ir.Operation, void).init(allocator) };
867 defer probe.deinit();
868 const info = rowPipelineRowInfo(op, shapes) orelse continue;
869 probe.rows = info.rows;
870 probe.cols = info.cols;
871
872 if (!try collectRowPipelineMembers(op, block, shapes, claimed, &probe)) continue;
873 if (probe.reduce_count == 0) continue;
874 if (probe.members.count() > max_row_pipeline_ops) continue;
875 if (!rowPipelineIsClosed(op, &probe)) continue;
876
877 var ops: std.ArrayListUnmanaged(*ir.Operation) = .empty;
878 defer ops.deinit(allocator);
879 var walk: ?*ir.Operation = @ptrCast(@alignCast(block.operations.head));
880 while (walk) |member| : (walk = member.next_op) {
881 if (member == op) continue;
882 if (probe.members.contains(member)) try ops.append(allocator, member);
883 }
884 try ops.append(allocator, op);
885
886 try analysis.addClusterOfKind(ops.items, .row_pipeline);
887 for (ops.items) |member| try claimed.put(member, .slot);
888 }
889 }
890
891 const RowPipelineRowInfo = struct {
892 rows: i64,
893 cols: i64,
894 };
895
896 const RowPipelineQuery = union(enum) {
897 member: *ir.Operation,
898 operand: *ir.Value,
899 broadcast: *ir.Operation,
900 stat: *ir.Value,
901 reduce: *ir.Operation,
902 };
903
904 fn rowPipelineRowInfo(op: *ir.Operation, shapes: *const shape_analysis.ShapeLayoutAnalysis) ?RowPipelineRowInfo {
905 const result = op.getResult(0) orelse return null;
906 const info = shapes.get(result) orelse return null;
907 if (!info.hasStaticLayout()) return null;
908 if (info.dims.len != 2) return null;
909 const rows = info.dims[0];
910 const cols = info.dims[1];
911 if (rows < 1 or cols < 1) return null;
912 if (cols > max_row_pipeline_cols) return null;
913 if (@rem(cols, @as(i64, row_pipeline_threads) * 4) != 0) return null;
914 return .{ .rows = rows, .cols = cols };
915 }
916
917 fn rowPipelineTerminalCandidate(
918 op: *ir.Operation,
919 shapes: *const shape_analysis.ShapeLayoutAnalysis,
920 claimed: *ClaimMap,
921 ) bool {
922 if (claimed.contains(op)) return false;
923 if (!isFusableElementwiseOp(op)) return false;
924 const result = op.getResult(0) orelse return false;
925 const info = shapes.get(result) orelse return false;
926 if (!info.hasStaticLayout() or info.dims.len != 2) return false;
927 var use = result.first_use;
928 while (use) |current_use| : (use = current_use.next_use) {
929 const user: *ir.Operation = @ptrCast(@alignCast(current_use.owner));
930 if (isFusableElementwiseOp(user) and user.getBlock() == op.getBlock()) return false;
931 if (isName(user.name.name, dialect_mod.AccyDialect.ReduceOp.operation_name)) return false;
932 }
933 return true;
934 }
935
936 fn rowPipelineValueIsF32Rows(
937 value: *ir.Value,
938 shapes: *const shape_analysis.ShapeLayoutAnalysis,
939 probe: *const RowPipelineProbe,
940 ) bool {
941 const info = shapes.get(value) orelse return false;
942 if (!info.hasStaticLayout() or info.dims.len != 2) return false;
943 if (info.dims[0] != probe.rows or info.dims[1] != probe.cols) return false;
944 return info.dtype == .f32;
945 }
946
947 fn collectRowPipelineMembers(
948 op: *ir.Operation,
949 block: *ir.Block,
950 shapes: *const shape_analysis.ShapeLayoutAnalysis,
951 claimed: *ClaimMap,
952 probe: *RowPipelineProbe,
953 ) anyerror!bool {
954 return rowPipelineQueryOk(.{ .member = op }, block, shapes, claimed, probe);
955 }
956
957 fn rowPipelineQueryOk(
958 query: RowPipelineQuery,
959 block: *ir.Block,
960 shapes: *const shape_analysis.ShapeLayoutAnalysis,
961 claimed: *ClaimMap,
962 probe: *RowPipelineProbe,
963 ) anyerror!bool {
964 switch (query) {
965 .member => |op| {
966 if (probe.members.contains(op)) return true;
967 if (claimed.contains(op)) return false;
968 if (op.getBlock() != block) return false;
969 if (!isFusableElementwiseOp(op)) return false;
970 const result = op.getResult(0) orelse return false;
971 if (!rowPipelineValueIsF32Rows(result, shapes, probe)) return false;
972 try probe.members.put(op, {});
973
974 for (op.getOperandValues()) |operand| {
975 if (try rowPipelineQueryOk(.{ .operand = operand }, block, shapes, claimed, probe)) continue;
976 return false;
977 }
978 return true;
979 },
980 .operand => |operand| {
981 const def_any = operand.getDefiningOp() orelse {
982 return rowPipelineValueIsF32Rows(operand, shapes, probe);
983 };
984 const def_op: *ir.Operation = @ptrCast(@alignCast(def_any));
985 if (probe.members.contains(def_op)) return true;
986 if (isName(def_op.name.name, dialect_mod.AccyDialect.ConstantOp.operation_name)) {
987 return rowPipelineValueIsF32Rows(operand, shapes, probe) and constantIsSplat(def_op);
988 }
989
990 if (isName(def_op.name.name, dialect_mod.AccyDialect.BroadcastInDimOp.operation_name) and
991 def_op.getBlock() == block and !claimed.contains(def_op))
992 {
993 if (try rowPipelineQueryOk(.{ .broadcast = def_op }, block, shapes, claimed, probe)) return true;
994 }
995
996 if (def_op.getBlock() == block and !claimed.contains(def_op) and isFusableElementwiseOp(def_op)) {
997 return rowPipelineQueryOk(.{ .member = def_op }, block, shapes, claimed, probe);
998 }
999
1000 return rowPipelineValueIsF32Rows(operand, shapes, probe);
1001 },
1002 .broadcast => |broadcast| {
1003 const broadcast_result = broadcast.getResult(0) orelse return false;
1004 if (!rowPipelineValueIsF32Rows(broadcast_result, shapes, probe)) return false;
1005 var dims_buffer: [8]i64 = undefined;
1006 const dims = readBroadcastDims(broadcast, dims_buffer[0..]) orelse return false;
1007 if (dims.len != 1) return false;
1008
1009 const operands = broadcast.getOperandValues();
1010 if (operands.len != 1) return false;
1011
1012 if (dims[0] == 1) {
1013 const source_info = shapes.get(operands[0]) orelse return false;
1014 if (!source_info.hasStaticLayout() or source_info.dims.len != 1) return false;
1015 if (source_info.dims[0] != probe.cols) return false;
1016 if (source_info.dtype != .f32) return false;
1017 if (operands[0].getDefiningOp()) |source_any| {
1018 const source: *ir.Operation = @ptrCast(@alignCast(source_any));
1019 if (probe.members.contains(source)) return false;
1020 }
1021 try probe.members.put(broadcast, {});
1022 return true;
1023 }
1024
1025 if (dims[0] != 0) return false;
1026 if (!try rowPipelineQueryOk(.{ .stat = operands[0] }, block, shapes, claimed, probe)) return false;
1027 try probe.members.put(broadcast, {});
1028 return true;
1029 },
1030 .stat => |value| {
1031 const info = shapes.get(value) orelse return false;
1032 if (!info.hasStaticLayout() or info.dims.len != 1) return false;
1033 if (info.dims[0] != probe.rows) return false;
1034 if (info.dtype != .f32) return false;
1035
1036 const def_any = value.getDefiningOp() orelse return false;
1037 const def_op: *ir.Operation = @ptrCast(@alignCast(def_any));
1038 if (probe.members.contains(def_op)) return true;
1039 if (def_op.getBlock() != block or claimed.contains(def_op)) return false;
1040
1041 if (isName(def_op.name.name, dialect_mod.AccyDialect.ConstantOp.operation_name)) {
1042 return constantIsSplat(def_op);
1043 }
1044
1045 if (isName(def_op.name.name, dialect_mod.AccyDialect.ReduceOp.operation_name)) {
1046 return rowPipelineQueryOk(.{ .reduce = def_op }, block, shapes, claimed, probe);
1047 }
1048
1049 if (!isFusableElementwiseOp(def_op)) return false;
1050 try probe.members.put(def_op, {});
1051 for (def_op.getOperandValues()) |operand| {
1052 if (!try rowPipelineQueryOk(.{ .stat = operand }, block, shapes, claimed, probe)) return false;
1053 }
1054 return true;
1055 },
1056 .reduce => |reduce| {
1057 if (probe.members.contains(reduce)) return true;
1058 if (probe.reduce_count >= max_row_pipeline_reduces) return false;
1059 if (!reduceHasConstantInit(reduce)) return false;
1060
1061 var axes_buffer: [8]i64 = undefined;
1062 const axes = readReduceDimensions(reduce, axes_buffer[0..]) orelse return false;
1063 if (axes.len != 1 or axes[0] != 1) return false;
1064
1065 const reduce_result = reduce.getResult(0) orelse return false;
1066 const reduce_info = shapes.get(reduce_result) orelse return false;
1067 if (!reduce_info.hasStaticLayout() or reduce_info.dims.len != 1) return false;
1068 if (reduce_info.dims[0] != probe.rows) return false;
1069
1070 const input = reduce.getOperand(0) orelse return false;
1071 try probe.members.put(reduce, {});
1072 probe.reduce_count += 1;
1073 if (!try rowPipelineQueryOk(.{ .operand = input }, block, shapes, claimed, probe)) return false;
1074 return true;
1075 },
1076 }
1077 }
1078
1079 fn constantIsSplat(op: *ir.Operation) bool {
1080 const constant = dialect_mod.AccyDialect.ConstantOp{ .op = op };
1081 const payload = constant.getPayload() orelse return false;
1082 if (payload.len < 4 or payload.len % 4 != 0) return false;
1083 const first = payload[0..4];
1084 var offset: usize = 4;
1085 while (offset < payload.len) : (offset += 4) {
1086 if (!std.mem.eql(u8, payload[offset..][0..4], first)) return false;
1087 }
1088 return true;
1089 }
1090
1091 fn readBroadcastDims(op: *ir.Operation, buffer: []i64) ?[]const i64 {
1092 if (readI64Payload(op.getAttrAs(ir.Attribute.DialectAttr, "broadcast_dims"), buffer)) |dims| return dims;
1093 return readI64Payload(op.getAttrAs(ir.Attribute.DialectAttr, dialect_mod.AccyDialect.BroadcastInDimOp.dialectAttrName("broadcast_dims")), buffer);
1094 }
1095
1096 fn readReduceDimensions(op: *ir.Operation, buffer: []i64) ?[]const i64 {
1097 if (readI64Payload(op.getAttrAs(ir.Attribute.DialectAttr, "dimensions"), buffer)) |dims| return dims;
1098 return readI64Payload(op.getAttrAs(ir.Attribute.DialectAttr, dialect_mod.AccyDialect.ReduceOp.dialectAttrName("dimensions")), buffer);
1099 }
1100
1101 fn readI64Payload(attr: ?*const ir.Attribute.DialectAttr, buffer: []i64) ?[]const i64 {
1102 const dialect_attr = attr orelse return null;
1103 if (dialect_attr.payload.len % @sizeOf(i64) != 0) return null;
1104 const count = dialect_attr.payload.len / @sizeOf(i64);
1105 if (count > buffer.len) return null;
1106 for (buffer[0..count], 0..) |*slot, index| {
1107 var value: i64 = undefined;
1108 @memcpy(std.mem.asBytes(&value), dialect_attr.payload[index * @sizeOf(i64) ..][0..@sizeOf(i64)]);
1109 slot.* = value;
1110 }
1111 return buffer[0..count];
1112 }
1113
1114 fn rowPipelineIsClosed(root: *ir.Operation, probe: *const RowPipelineProbe) bool {
1115 var member_iter = probe.members.keyIterator();
1116 while (member_iter.next()) |member| {
1117 if (member.* == root) continue;
1118 const result = member.*.getResult(0) orelse return false;
1119 var use = result.first_use;
1120 while (use) |current_use| : (use = current_use.next_use) {
1121 const user: *ir.Operation = @ptrCast(@alignCast(current_use.owner));
1122 if (!probe.members.contains(user) and user != root) return false;
1123 }
1124 }
1125 return true;
1126 }
1127
1128 pub const max_dot_epilogue_ops = 6;
1129 pub const max_reduction_prologue_ops = 48;
1130 const max_reduction_prologue_depth = 24;
1131
1132 fn collectReductionInputPlansInRegions(
1133 allocator: std.mem.Allocator,
1134 regions: []ir.Region,
1135 shapes: *const shape_analysis.ShapeLayoutAnalysis,
1136 analysis: *FusionPlanAnalysis,
1137 claimed: *ClaimMap,
1138 ) anyerror!void {
1139 for (regions) |*region| {
1140 var block_iter = region.getBlocks();
1141 while (block_iter.next()) |block| {
1142 var current: ?*ir.Operation = @ptrCast(@alignCast(block.operations.head));
1143 while (current) |op| {
1144 const next = op.next_op;
1145 if (op.regions.items.len > 0 and !regionsAreOpaque(op)) {
1146 try collectReductionInputPlansInRegions(allocator, op.regions.items, shapes, analysis, claimed);
1147 }
1148 current = next;
1149 }
1150 try collectBlockReductionInputPlans(allocator, block, shapes, analysis, claimed);
1151 }
1152 }
1153 }
1154
1155 fn collectBlockReductionInputPlans(
1156 allocator: std.mem.Allocator,
1157 block: *ir.Block,
1158 shapes: *const shape_analysis.ShapeLayoutAnalysis,
1159 analysis: *FusionPlanAnalysis,
1160 claimed: *ClaimMap,
1161 ) anyerror!void {
1162 var reduces: std.ArrayListUnmanaged(*ir.Operation) = .empty;
1163 defer reduces.deinit(allocator);
1164
1165 var current: ?*ir.Operation = @ptrCast(@alignCast(block.operations.head));
1166 while (current) |op| : (current = op.next_op) {
1167 if (!isName(op.name.name, dialect_mod.AccyDialect.ReduceOp.operation_name)) continue;
1168 if (claimed.contains(op)) continue;
1169 if (!reduceHasConstantInit(op)) continue;
1170 try reduces.append(allocator, op);
1171 }
1172 if (reduces.items.len == 0) return;
1173
1174 var members = std.AutoHashMap(*ir.Operation, void).init(allocator);
1175 defer members.deinit();
1176
1177 for (reduces.items) |reduce| {
1178 const input = reduce.getOperand(0) orelse continue;
1179 const baseline = shapes.get(input) orelse continue;
1180 try collectReductionProducerDAG(input, baseline, block, shapes, claimed, &members, max_reduction_prologue_depth);
1181 }
1182 if (members.count() == 0) return;
1183 if (members.count() > max_reduction_prologue_ops) return;
1184
1185 var removals: std.ArrayListUnmanaged(*ir.Operation) = .empty;
1186 defer removals.deinit(allocator);
1187 var changed = true;
1188 while (changed) {
1189 changed = false;
1190 removals.clearRetainingCapacity();
1191 var member_iter = members.keyIterator();
1192 while (member_iter.next()) |member| {
1193 if (!reductionUsesInside(member.*, &members, reduces.items, shapes, claimed)) {
1194 try removals.append(allocator, member.*);
1195 }
1196 }
1197 for (removals.items) |member| {
1198 _ = members.remove(member);
1199 changed = true;
1200 }
1201 }
1202 if (members.count() == 0) return;
1203
1204 var owned = std.AutoHashMap(*ir.Operation, void).init(allocator);
1205 defer owned.deinit();
1206 var cluster: std.ArrayListUnmanaged(*ir.Operation) = .empty;
1207 defer cluster.deinit(allocator);
1208
1209 var grouped = std.AutoHashMap(*ir.Operation, void).init(allocator);
1210 defer grouped.deinit();
1211 try collectConcatReductionClusters(allocator, block, shapes, reduces.items, &members, &owned, &cluster, &grouped, analysis, claimed);
1212
1213 for (reduces.items) |reduce| {
1214 if (grouped.contains(reduce)) continue;
1215 const input = reduce.getOperand(0) orelse continue;
1216 cluster.clearRetainingCapacity();
1217 try collectOwnedPrologue(allocator, input, &members, &owned, &cluster);
1218 const input_inlined = inputIsMember(input, &members);
1219 if (cluster.items.len == 0 and !input_inlined) continue;
1220 try cluster.append(allocator, reduce);
1221 try analysis.addReductionCluster(cluster.items);
1222 for (cluster.items) |cluster_op| try claimed.put(cluster_op, .recompute);
1223 }
1224 }
1225
1226 pub const max_concat_reduction_group = 8;
1227
1228 fn collectConcatReductionClusters(
1229 allocator: std.mem.Allocator,
1230 block: *ir.Block,
1231 shapes: *const shape_analysis.ShapeLayoutAnalysis,
1232 reduces: []const *ir.Operation,
1233 members: *const std.AutoHashMap(*ir.Operation, void),
1234 owned: *std.AutoHashMap(*ir.Operation, void),
1235 cluster: *std.ArrayListUnmanaged(*ir.Operation),
1236 grouped: *std.AutoHashMap(*ir.Operation, void),
1237 analysis: *FusionPlanAnalysis,
1238 claimed: *ClaimMap,
1239 ) anyerror!void {
1240 var current: ?*ir.Operation = @ptrCast(@alignCast(block.operations.head));
1241 while (current) |op| : (current = op.next_op) {
1242 if (!isName(op.name.name, dialect_mod.AccyDialect.ConcatenateOp.operation_name)) continue;
1243 if (claimed.contains(op)) continue;
1244 const operands = op.getOperandValues();
1245 if (operands.len < 2 or operands.len > max_concat_reduction_group) continue;
1246 if (!concatIsAxisZeroRankOne(op, shapes)) continue;
1247
1248 var group_buffer: [max_concat_reduction_group]*ir.Operation = undefined;
1249 const group = concatReduceGroup(operands, reduces, shapes, group_buffer[0..]) orelse continue;
1250
1251 cluster.clearRetainingCapacity();
1252 var any_inlined = false;
1253 for (group) |reduce| {
1254 const input = reduce.getOperand(0) orelse break;
1255 try collectOwnedPrologue(allocator, input, members, owned, cluster);
1256 if (inputIsMember(input, members)) any_inlined = true;
1257 }
1258 if (cluster.items.len == 0 and !any_inlined) continue;
1259 for (group) |reduce| try cluster.append(allocator, reduce);
1260 try cluster.append(allocator, op);
1261 try analysis.addReductionCluster(cluster.items);
1262 for (cluster.items) |cluster_op| try claimed.put(cluster_op, .recompute);
1263 for (group) |reduce| try grouped.put(reduce, {});
1264 }
1265 }
1266
1267 fn concatIsAxisZeroRankOne(op: *ir.Operation, shapes: *const shape_analysis.ShapeLayoutAnalysis) bool {
1268 const attr = op.getAttrAs(ir.Attribute.IntegerAttr, "dimension") orelse return false;
1269 if (attr.getValue() != 0) return false;
1270 const result = op.getResult(0) orelse return false;
1271 const info = shapes.get(result) orelse return false;
1272 return info.dims.len == 1;
1273 }
1274
1275 fn concatReduceGroup(
1276 operands: []const *ir.Value,
1277 reduces: []const *ir.Operation,
1278 shapes: *const shape_analysis.ShapeLayoutAnalysis,
1279 buffer: []*ir.Operation,
1280 ) ?[]*ir.Operation {
1281 var group_shape: ?shape_analysis.TensorInfo = null;
1282 var input_shape: ?shape_analysis.TensorInfo = null;
1283 for (operands, 0..) |operand, index| {
1284 if (!operand.hasOneUse()) return null;
1285 const def_any = operand.getDefiningOp() orelse return null;
1286 const def_op: *ir.Operation = @ptrCast(@alignCast(def_any));
1287 var collected = false;
1288 for (reduces) |reduce| {
1289 if (reduce == def_op) {
1290 collected = true;
1291 break;
1292 }
1293 }
1294 if (!collected) return null;
1295 const operand_info = shapes.get(operand) orelse return null;
1296 if (group_shape) |existing| {
1297 if (!sameStaticShape(existing, operand_info)) return null;
1298 } else {
1299 group_shape = operand_info;
1300 }
1301 const reduce_input = def_op.getOperand(0) orelse return null;
1302 const reduce_input_info = shapes.get(reduce_input) orelse return null;
1303 if (input_shape) |existing| {
1304 if (!sameStaticShape(existing, reduce_input_info)) return null;
1305 } else {
1306 input_shape = reduce_input_info;
1307 }
1308 buffer[index] = def_op;
1309 }
1310 return buffer[0..operands.len];
1311 }
1312
1313 fn inputIsMember(input: *ir.Value, members: *const std.AutoHashMap(*ir.Operation, void)) bool {
1314 const def_any = input.getDefiningOp() orelse return false;
1315 const def_op: *ir.Operation = @ptrCast(@alignCast(def_any));
1316 return members.contains(def_op);
1317 }
1318
1319 fn reduceHasConstantInit(reduce: *ir.Operation) bool {
1320 const init = reduce.getOperand(1) orelse return false;
1321 const def_any = init.getDefiningOp() orelse return false;
1322 const def_op: *ir.Operation = @ptrCast(@alignCast(def_any));
1323 return isName(def_op.name.name, dialect_mod.AccyDialect.ConstantOp.operation_name);
1324 }
1325
1326 fn collectReductionProducerDAG(
1327 value: *ir.Value,
1328 baseline: shape_analysis.TensorInfo,
1329 block: *ir.Block,
1330 shapes: *const shape_analysis.ShapeLayoutAnalysis,
1331 claimed: *ClaimMap,
1332 members: *std.AutoHashMap(*ir.Operation, void),
1333 depth: usize,
1334 ) anyerror!void {
1335 if (depth == 0) return;
1336 const def_any = value.getDefiningOp() orelse return;
1337 const op: *ir.Operation = @ptrCast(@alignCast(def_any));
1338 if (members.contains(op)) return;
1339 if (claimed.contains(op)) return;
1340 if (op.getBlock() != block) return;
1341 if (!isFusableElementwiseOp(op)) return;
1342
1343 const result = op.getResult(0) orelse return;
1344 const result_info = shapes.get(result) orelse return;
1345 if (!sameStaticShape(result_info, baseline)) return;
1346
1347 for (op.getOperandValues()) |operand| {
1348 if (broadcastLeafSource(operand) != null) continue;
1349 const operand_info = shapes.get(operand) orelse return;
1350 if (!sameStaticShape(operand_info, baseline)) return;
1351 }
1352
1353 try members.put(op, {});
1354 for (op.getOperandValues()) |operand| {
1355 if (broadcastLeafSource(operand) != null) continue;
1356 try collectReductionProducerDAG(operand, baseline, block, shapes, claimed, members, depth - 1);
1357 }
1358 }
1359
1360 fn reductionUsesInside(
1361 op: *ir.Operation,
1362 members: *const std.AutoHashMap(*ir.Operation, void),
1363 reduces: []const *ir.Operation,
1364 shapes: *const shape_analysis.ShapeLayoutAnalysis,
1365 claimed: *const ClaimMap,
1366 ) bool {
1367 const result = op.getResult(0) orelse return false;
1368 var use = result.first_use;
1369 while (use) |current_use| : (use = current_use.next_use) {
1370 const user: *ir.Operation = @ptrCast(@alignCast(current_use.owner));
1371 if (members.contains(user)) continue;
1372 var is_fused_reduce = false;
1373 for (reduces) |reduce| {
1374 if (user == reduce) {
1375 is_fused_reduce = true;
1376 break;
1377 }
1378 }
1379 if (is_fused_reduce) continue;
1380 if (userCanRecompute(op, user, shapes, claimed)) continue;
1381 return false;
1382 }
1383 return true;
1384 }
1385
1386 fn collectOwnedPrologue(
1387 allocator: std.mem.Allocator,
1388 value: *ir.Value,
1389 members: *const std.AutoHashMap(*ir.Operation, void),
1390 owned: *std.AutoHashMap(*ir.Operation, void),
1391 cluster: *std.ArrayListUnmanaged(*ir.Operation),
1392 ) anyerror!void {
1393 const def_any = value.getDefiningOp() orelse return;
1394 const op: *ir.Operation = @ptrCast(@alignCast(def_any));
1395 if (!members.contains(op)) return;
1396 if (owned.contains(op)) return;
1397 try owned.put(op, {});
1398 for (op.getOperandValues()) |operand| {
1399 if (broadcastLeafSource(operand) != null) continue;
1400 try collectOwnedPrologue(allocator, operand, members, owned, cluster);
1401 }
1402 try cluster.append(allocator, op);
1403 }
1404
1405 fn collectDotEpiloguePlansInRegions(
1406 allocator: std.mem.Allocator,
1407 regions: []ir.Region,
1408 shapes: *const shape_analysis.ShapeLayoutAnalysis,
1409 analysis: *FusionPlanAnalysis,
1410 claimed: *ClaimMap,
1411 ) anyerror!void {
1412 for (regions) |*region| {
1413 var block_iter = region.getBlocks();
1414 while (block_iter.next()) |block| {
1415 var current: ?*ir.Operation = @ptrCast(@alignCast(block.operations.head));
1416 while (current) |op| {
1417 const next = op.next_op;
1418 if (op.regions.items.len > 0 and !regionsAreOpaque(op)) {
1419 try collectDotEpiloguePlansInRegions(allocator, op.regions.items, shapes, analysis, claimed);
1420 }
1421 if (isName(op.name.name, dialect_mod.AccyDialect.DotGeneralOp.operation_name) and !claimed.contains(op)) {
1422 var chain_buffer: [1 + max_dot_epilogue_ops]*ir.Operation = undefined;
1423 const chain = collectDotEpilogueChain(op, shapes, claimed, chain_buffer[0..]);
1424 if (chain.len >= 2) {
1425 try analysis.addClusterOfKind(chain, .dot_epilogue);
1426 for (chain) |chain_op| try claimed.put(chain_op, .slot);
1427 }
1428 }
1429 current = next;
1430 }
1431 }
1432 }
1433 }
1434
1435 fn collectDotEpilogueChain(
1436 dot: *ir.Operation,
1437 shapes: *const shape_analysis.ShapeLayoutAnalysis,
1438 claimed: *ClaimMap,
1439 buffer: []*ir.Operation,
1440 ) []*ir.Operation {
1441 buffer[0] = dot;
1442 var count: usize = 1;
1443 var producer = dot;
1444 while (count < buffer.len) {
1445 const consumer = soleFusableConsumer(producer, shapes, claimed) orelse break;
1446 if (!dotEpilogueOperandsSupported(consumer, producer, shapes)) break;
1447 buffer[count] = consumer;
1448 count += 1;
1449 producer = consumer;
1450 }
1451 return buffer[0..count];
1452 }
1453
1454 fn soleFusableConsumer(
1455 producer: *ir.Operation,
1456 shapes: *const shape_analysis.ShapeLayoutAnalysis,
1457 claimed: *ClaimMap,
1458 ) ?*ir.Operation {
1459 const result = producer.getResult(0) orelse return null;
1460 if (!result.hasOneUse()) return null;
1461 const use = result.first_use orelse return null;
1462 const user: *ir.Operation = @ptrCast(@alignCast(use.owner));
1463 if (user.getBlock() != producer.getBlock()) return null;
1464 if (claimed.contains(user)) return null;
1465 if (!isDotEpilogueOp(user)) return null;
1466 if (!compatibleProducerConsumerShapes(producer, user, shapes)) return null;
1467 const user_result = user.getResult(0) orelse return null;
1468 const user_info = shapes.get(user_result) orelse return null;
1469 if (user_info.dtype != .f32) return null;
1470 return user;
1471 }
1472
1473 fn isDotEpilogueOp(op: *ir.Operation) bool {
1474 if (!isFusableElementwiseOp(op)) return false;
1475 const name = op.name.name;
1476 if (std.mem.eql(u8, name, dialect_mod.AccyDialect.CompareOp.operation_name)) return false;
1477 if (std.mem.eql(u8, name, dialect_mod.AccyDialect.ConvertOp.operation_name)) return false;
1478 if (std.mem.eql(u8, name, dialect_mod.AccyDialect.SelectOp.operation_name)) return false;
1479 return true;
1480 }
1481
1482 fn compatibleProducerConsumerShapes(
1483 producer: *ir.Operation,
1484 consumer: *ir.Operation,
1485 shapes: *const shape_analysis.ShapeLayoutAnalysis,
1486 ) bool {
1487 const producer_result = producer.getResult(0) orelse return false;
1488 const consumer_result = consumer.getResult(0) orelse return false;
1489 const producer_info = shapes.get(producer_result) orelse return false;
1490 const consumer_info = shapes.get(consumer_result) orelse return false;
1491 return sameStaticShape(producer_info, consumer_info);
1492 }
1493
1494 fn dotEpilogueOperandsSupported(
1495 consumer: *ir.Operation,
1496 producer: *ir.Operation,
1497 shapes: *const shape_analysis.ShapeLayoutAnalysis,
1498 ) bool {
1499 const producer_result = producer.getResult(0) orelse return false;
1500 const consumer_result = consumer.getResult(0) orelse return false;
1501 const consumer_info = shapes.get(consumer_result) orelse return false;
1502 for (consumer.getOperandValues()) |operand| {
1503 if (operand == producer_result) continue;
1504 if (broadcastLeafSource(operand) != null) continue;
1505 const operand_info = shapes.get(operand) orelse return false;
1506 if (!sameStaticShape(consumer_info, operand_info)) return false;
1507 if (operand.getDefiningOp() != null) return false;
1508 }
1509 return true;
1510 }
1511
1512 pub fn broadcastLeafSource(operand: *ir.Value) ?*ir.Value {
1513 const def_any = operand.getDefiningOp() orelse return null;
1514 const def_op: *ir.Operation = @ptrCast(@alignCast(def_any));
1515 if (!isName(def_op.name.name, dialect_mod.AccyDialect.BroadcastInDimOp.operation_name)) return null;
1516 const operands = def_op.getOperandValues();
1517 if (operands.len != 1) return null;
1518 if (operands[0].getDefiningOp() != null) return null;
1519 return operands[0];
1520 }
1521
1522 fn isName(name: []const u8, expected: []const u8) bool {
1523 return std.mem.eql(u8, name, expected);
1524 }
1525
1526 pub fn isFusableElementwiseOp(op: *ir.Operation) bool {
1527 if (op.regions.items.len > 0) return false;
1528 if (op.getNumResults() != 1) return false;
1529 const name = op.name.name;
1530 return std.mem.eql(u8, name, dialect_mod.AccyDialect.AddOp.operation_name) or
1531 std.mem.eql(u8, name, dialect_mod.AccyDialect.SubOp.operation_name) or
1532 std.mem.eql(u8, name, dialect_mod.AccyDialect.MulOp.operation_name) or
1533 std.mem.eql(u8, name, dialect_mod.AccyDialect.DivOp.operation_name) or
1534 std.mem.eql(u8, name, dialect_mod.AccyDialect.MaxOp.operation_name) or
1535 std.mem.eql(u8, name, dialect_mod.AccyDialect.MinOp.operation_name) or
1536 std.mem.eql(u8, name, dialect_mod.AccyDialect.PowOp.operation_name) or
1537 std.mem.eql(u8, name, dialect_mod.AccyDialect.Atan2Op.operation_name) or
1538 std.mem.eql(u8, name, dialect_mod.AccyDialect.NegOp.operation_name) or
1539 std.mem.eql(u8, name, dialect_mod.AccyDialect.ExpOp.operation_name) or
1540 std.mem.eql(u8, name, dialect_mod.AccyDialect.LogOp.operation_name) or
1541 std.mem.eql(u8, name, dialect_mod.AccyDialect.TanhOp.operation_name) or
1542 std.mem.eql(u8, name, dialect_mod.AccyDialect.SqrtOp.operation_name) or
1543 std.mem.eql(u8, name, dialect_mod.AccyDialect.AbsOp.operation_name) or
1544 std.mem.eql(u8, name, dialect_mod.AccyDialect.SinOp.operation_name) or
1545 std.mem.eql(u8, name, dialect_mod.AccyDialect.CosOp.operation_name) or
1546 std.mem.eql(u8, name, dialect_mod.AccyDialect.TanOp.operation_name) or
1547 std.mem.eql(u8, name, dialect_mod.AccyDialect.FloorOp.operation_name) or
1548 std.mem.eql(u8, name, dialect_mod.AccyDialect.RoundOp.operation_name) or
1549 std.mem.eql(u8, name, dialect_mod.AccyDialect.TruncOp.operation_name) or
1550 std.mem.eql(u8, name, dialect_mod.AccyDialect.CompareOp.operation_name) or
1551 std.mem.eql(u8, name, dialect_mod.AccyDialect.ConvertOp.operation_name) or
1552 std.mem.eql(u8, name, dialect_mod.AccyDialect.SelectOp.operation_name);
1553 }
1554
1555 const testing = std.testing;
1556 const semantic = accy_choir.semantic;
1557
1558 const FusionProbe = struct {
1559 var cluster_count: usize = 0;
1560 var fused_op_count: usize = 0;
1561 var max_cluster_len: usize = 0;
1562 var first_root_name: []const u8 = "";
1563
1564 fn reset() void {
1565 cluster_count = 0;
1566 fused_op_count = 0;
1567 max_cluster_len = 0;
1568 first_root_name = "";
1569 }
1570
1571 fn pass() passes.Pass {
1572 return .{
1573 .name = "accy-choir-fusion-test-probe",
1574 .description = "Inspect Accy Choir fusion plans in tests",
1575 .run_fn = run,
1576 .work_contract = .{
1577 .identity = .{ .name = "accy-choir-fusion-test-probe", .version = 1 },
1578 .estimate = probeWork,
1579 },
1580 };
1581 }
1582
1583 fn probeWork(_: passes.pass.work.Input) !passes.pass.work.Bounds {
1584 return .{ .work = .{ .structural_visits = 16 } };
1585 }
1586
1587 fn run(pass_ctx: *passes.PassContext) passes.PassResult {
1588 const analysis = getFusionPlanAnalysis(pass_ctx, pass_ctx.op) catch return .failure;
1589 cluster_count = analysis.clusterCount();
1590 fused_op_count = analysis.fused_op_count;
1591 max_cluster_len = analysis.max_cluster_len;
1592 if (analysis.clusters.items.len > 0) {
1593 first_root_name = analysis.clusters.items[0].root().?.name.name;
1594 }
1595 pass_ctx.preserveAllAnalyses();
1596 return .success;
1597 }
1598 };
1599
1600 test "fusion planning finds straight-line elementwise chain" {
1601 const allocator = testing.allocator;
1602 FusionProbe.reset();
1603
1604 var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
1605 defer builder.deinit();
1606 const f32_4 = try builder.tensor(.f32, &.{4});
1607 var fb = try builder.beginFunction("fusion_add_mul", &.{ f32_4, f32_4, f32_4 }, &.{f32_4});
1608 const sum = try fb.add(fb.parameter(0), fb.parameter(1));
1609 const product = try fb.mul(sum, fb.parameter(2));
1610 try fb.return_(&.{product});
1611 try fb.finish();
1612 const module = try builder.finish();
1613 defer module.deinit();
1614
1615 const choir_mod = module.choir_module;
1616 const ctx = module.context();
1617 var pm = passes.PassManager.init(allocator);
1618 defer pm.deinit();
1619 try pm.addPass(fusionPlanningPass());
1620 try pm.addPass(FusionProbe.pass());
1621
1622 try runAccountedFusion(&pm, choir_mod, ctx);
1623 try testing.expectEqual(@as(u64, 2), pm.stats.pass_runs);
1624 try testing.expectEqual(@as(usize, 1), FusionProbe.cluster_count);
1625 try testing.expectEqual(@as(usize, 2), FusionProbe.fused_op_count);
1626 try testing.expectEqual(@as(usize, 2), FusionProbe.max_cluster_len);
1627 try testing.expectEqualStrings(dialect_mod.AccyDialect.MulOp.operation_name, FusionProbe.first_root_name);
1628 }
1629
1630 test "fusion planning finds single-use elementwise producer DAG" {
1631 const allocator = testing.allocator;
1632 FusionProbe.reset();
1633
1634 var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
1635 defer builder.deinit();
1636 const f32_4 = try builder.tensor(.f32, &.{4});
1637 var fb = try builder.beginFunction("fusion_branch_sum", &.{ f32_4, f32_4, f32_4, f32_4 }, &.{f32_4});
1638 const left = try fb.mul(fb.parameter(0), fb.parameter(1));
1639 const right = try fb.mul(fb.parameter(2), fb.parameter(3));
1640 const sum = try fb.add(left, right);
1641 try fb.return_(&.{sum});
1642 try fb.finish();
1643 const module = try builder.finish();
1644 defer module.deinit();
1645
1646 const choir_mod = module.choir_module;
1647 const ctx = module.context();
1648 var pm = passes.PassManager.init(allocator);
1649 defer pm.deinit();
1650 try pm.addPass(fusionPlanningPass());
1651 try pm.addPass(FusionProbe.pass());
1652
1653 try runAccountedFusion(&pm, choir_mod, ctx);
1654 try testing.expectEqual(@as(usize, 1), FusionProbe.cluster_count);
1655 try testing.expectEqual(@as(usize, 3), FusionProbe.fused_op_count);
1656 try testing.expectEqual(@as(usize, 3), FusionProbe.max_cluster_len);
1657 try testing.expectEqualStrings(dialect_mod.AccyDialect.AddOp.operation_name, FusionProbe.first_root_name);
1658 }
1659
1660 test "fusion planning claims slices and pads as see-through members" {
1661 const allocator = testing.allocator;
1662 FusionProbe.reset();
1663
1664 var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
1665 defer builder.deinit();
1666 const f32_4x4 = try builder.tensor(.f32, &.{ 4, 4 });
1667 const f32_6x6 = try builder.tensor(.f32, &.{ 6, 6 });
1668 const scalar = try builder.tensor(.f32, &.{});
1669 var fb = try builder.beginFunction("fusion_stencil", &.{f32_4x4}, &.{f32_4x4});
1670 const zero = try fb.constant(scalar, std.mem.asBytes(&@as(f32, 0.0)));
1671 const padded = try fb.pad(fb.parameter(0), zero, f32_6x6, &.{ 1, 1 }, &.{ 1, 1 }, &.{ 0, 0 });
1672 const up = try fb.slice(padded, f32_4x4, &.{ 0, 1 }, &.{ 4, 5 }, &.{ 1, 1 });
1673 const down = try fb.slice(padded, f32_4x4, &.{ 2, 1 }, &.{ 6, 5 }, &.{ 1, 1 });
1674 const sum = try fb.add(up, down);
1675 try fb.return_(&.{sum});
1676 try fb.finish();
1677 const module = try builder.finish();
1678 defer module.deinit();
1679
1680 const choir_mod = module.choir_module;
1681 const ctx = module.context();
1682 var pm = passes.PassManager.init(allocator);
1683 defer pm.deinit();
1684 try pm.addPass(fusionPlanningPass());
1685 try pm.addPass(FusionProbe.pass());
1686
1687 try runAccountedFusion(&pm, choir_mod, ctx);
1688 try testing.expectEqual(@as(usize, 1), FusionProbe.cluster_count);
1689 try testing.expectEqual(@as(usize, 4), FusionProbe.fused_op_count);
1690 try testing.expectEqualStrings(dialect_mod.AccyDialect.AddOp.operation_name, FusionProbe.first_root_name);
1691 }
1692
1693 test "fusion planning claims reduction prologues with shared recompute" {
1694 const allocator = testing.allocator;
1695 FusionProbe.reset();
1696
1697 var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
1698 defer builder.deinit();
1699 const f32_4x8 = try builder.tensor(.f32, &.{ 4, 8 });
1700 const f32_4 = try builder.tensor(.f32, &.{4});
1701 const scalar = try builder.tensor(.f32, &.{});
1702 var fb = try builder.beginFunction("fusion_reduce_pair", &.{ f32_4x8, f32_4x8 }, &.{ f32_4, f32_4 });
1703 const zero = try fb.constant(scalar, std.mem.asBytes(&@as(f32, 0.0)));
1704 const shared = try fb.mul(fb.parameter(0), fb.parameter(1));
1705 const left = try fb.add(shared, fb.parameter(0));
1706 const right = try fb.sub(shared, fb.parameter(1));
1707 const left_sum = try fb.reduce(left, zero, f32_4, "sum", &.{1});
1708 const right_sum = try fb.reduce(right, zero, f32_4, "sum", &.{1});
1709 try fb.return_(&.{ left_sum, right_sum });
1710 try fb.finish();
1711 const module = try builder.finish();
1712 defer module.deinit();
1713
1714 const choir_mod = module.choir_module;
1715 const ctx = module.context();
1716 var pm = passes.PassManager.init(allocator);
1717 defer pm.deinit();
1718 try pm.addPass(fusionPlanningPass());
1719 try pm.addPass(FusionProbe.pass());
1720
1721 try runAccountedFusion(&pm, choir_mod, ctx);
1722 try testing.expectEqual(@as(usize, 2), FusionProbe.cluster_count);
1723 try testing.expectEqual(@as(usize, 5), FusionProbe.fused_op_count);
1724 }
1725
1726 test "fusion planning claims softmax as one row pipeline" {
1727 const allocator = testing.allocator;
1728 FusionProbe.reset();
1729
1730 var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
1731 defer builder.deinit();
1732 const f32_4x1024 = try builder.tensor(.f32, &.{ 4, 1024 });
1733 const f32_4 = try builder.tensor(.f32, &.{4});
1734 const scalar = try builder.tensor(.f32, &.{});
1735 var fb = try builder.beginFunction("fusion_softmax", &.{f32_4x1024}, &.{f32_4x1024});
1736 const neg_inf = try fb.constant(scalar, std.mem.asBytes(&@as(f32, -std.math.inf(f32))));
1737 const zero = try fb.constant(scalar, std.mem.asBytes(&@as(f32, 0.0)));
1738 const row_max = try fb.reduce(fb.parameter(0), neg_inf, f32_4, "max", &.{1});
1739 const max_full = try fb.broadcastInDim(row_max, f32_4x1024, &.{ 4, 1024 }, &.{0});
1740 const shifted = try fb.sub(fb.parameter(0), max_full);
1741 const exps = try fb.exp(shifted);
1742 const row_sum = try fb.reduce(exps, zero, f32_4, "sum", &.{1});
1743 const sum_full = try fb.broadcastInDim(row_sum, f32_4x1024, &.{ 4, 1024 }, &.{0});
1744 const out = try fb.div(exps, sum_full);
1745 try fb.return_(&.{out});
1746 try fb.finish();
1747 const module = try builder.finish();
1748 defer module.deinit();
1749
1750 const choir_mod = module.choir_module;
1751 const ctx = module.context();
1752 var pm = passes.PassManager.init(allocator);
1753 defer pm.deinit();
1754 try pm.addPass(fusionPlanningPass());
1755 try pm.addPass(FusionProbe.pass());
1756
1757 try runAccountedFusion(&pm, choir_mod, ctx);
1758 try testing.expectEqual(@as(usize, 1), FusionProbe.cluster_count);
1759 try testing.expectEqual(@as(usize, 7), FusionProbe.fused_op_count);
1760 try testing.expectEqual(@as(usize, 7), FusionProbe.max_cluster_len);
1761 try testing.expectEqualStrings(dialect_mod.AccyDialect.DivOp.operation_name, FusionProbe.first_root_name);
1762 }
1763
1764 test "fusion planning keeps narrow softmax on sibling recompute" {
1765 const allocator = testing.allocator;
1766 FusionProbe.reset();
1767
1768 var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
1769 defer builder.deinit();
1770 const f32_4x8 = try builder.tensor(.f32, &.{ 4, 8 });
1771 const f32_4 = try builder.tensor(.f32, &.{4});
1772 const scalar = try builder.tensor(.f32, &.{});
1773 var fb = try builder.beginFunction("fusion_softmax_narrow", &.{f32_4x8}, &.{f32_4x8});
1774 const neg_inf = try fb.constant(scalar, std.mem.asBytes(&@as(f32, -std.math.inf(f32))));
1775 const zero = try fb.constant(scalar, std.mem.asBytes(&@as(f32, 0.0)));
1776 const row_max = try fb.reduce(fb.parameter(0), neg_inf, f32_4, "max", &.{1});
1777 const max_full = try fb.broadcastInDim(row_max, f32_4x8, &.{ 4, 8 }, &.{0});
1778 const shifted = try fb.sub(fb.parameter(0), max_full);
1779 const exps = try fb.exp(shifted);
1780 const row_sum = try fb.reduce(exps, zero, f32_4, "sum", &.{1});
1781 const sum_full = try fb.broadcastInDim(row_sum, f32_4x8, &.{ 4, 8 }, &.{0});
1782 const out = try fb.div(exps, sum_full);
1783 try fb.return_(&.{out});
1784 try fb.finish();
1785 const module = try builder.finish();
1786 defer module.deinit();
1787
1788 const choir_mod = module.choir_module;
1789 const ctx = module.context();
1790 var pm = passes.PassManager.init(allocator);
1791 defer pm.deinit();
1792 try pm.addPass(fusionPlanningPass());
1793 try pm.addPass(FusionProbe.pass());
1794
1795 try runAccountedFusion(&pm, choir_mod, ctx);
1796 try testing.expectEqual(@as(usize, 1), FusionProbe.cluster_count);
1797 try testing.expectEqual(@as(usize, 3), FusionProbe.fused_op_count);
1798 try testing.expectEqual(@as(usize, 3), FusionProbe.max_cluster_len);
1799 }
1800
1801 test "fusion planning elides producer shared by sibling roots" {
1802 const allocator = testing.allocator;
1803 FusionProbe.reset();
1804
1805 var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
1806 defer builder.deinit();
1807 const f32_4 = try builder.tensor(.f32, &.{4});
1808 var fb = try builder.beginFunction("fusion_sibling_roots", &.{ f32_4, f32_4 }, &.{ f32_4, f32_4 });
1809 const shared = try fb.mul(fb.parameter(0), fb.parameter(1));
1810 const first = try fb.sin(shared);
1811 const second = try fb.cos(shared);
1812 try fb.return_(&.{ first, second });
1813 try fb.finish();
1814 const module = try builder.finish();
1815 defer module.deinit();
1816
1817 const choir_mod = module.choir_module;
1818 const ctx = module.context();
1819 var pm = passes.PassManager.init(allocator);
1820 defer pm.deinit();
1821 try pm.addPass(fusionPlanningPass());
1822 try pm.addPass(FusionProbe.pass());
1823
1824 try runAccountedFusion(&pm, choir_mod, ctx);
1825 try testing.expectEqual(@as(usize, 1), FusionProbe.cluster_count);
1826 try testing.expectEqual(@as(usize, 2), FusionProbe.fused_op_count);
1827 try testing.expectEqual(@as(usize, 2), FusionProbe.max_cluster_len);
1828 }
1829
1830 test "fusion planning rejects producer with escaping use" {
1831 const allocator = testing.allocator;
1832 FusionProbe.reset();
1833
1834 var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
1835 defer builder.deinit();
1836 const f32_4 = try builder.tensor(.f32, &.{4});
1837 var fb = try builder.beginFunction("fusion_escape", &.{ f32_4, f32_4, f32_4 }, &.{ f32_4, f32_4 });
1838 const sum = try fb.add(fb.parameter(0), fb.parameter(1));
1839 const product = try fb.mul(sum, fb.parameter(2));
1840 try fb.return_(&.{ sum, product });
1841 try fb.finish();
1842 const module = try builder.finish();
1843 defer module.deinit();
1844
1845 const choir_mod = module.choir_module;
1846 const ctx = module.context();
1847 var pm = passes.PassManager.init(allocator);
1848 defer pm.deinit();
1849 try pm.addPass(fusionPlanningPass());
1850 try pm.addPass(FusionProbe.pass());
1851
1852 try runAccountedFusion(&pm, choir_mod, ctx);
1853 try testing.expectEqual(@as(usize, 0), FusionProbe.cluster_count);
1854 try testing.expectEqual(@as(usize, 0), FusionProbe.fused_op_count);
1855 }
1856
1857 fn runAccountedFusion(manager: *passes.PassManager, op: *ir.Operation, ctx: *ir.Context) !void {
1858 const revision = choir.product.revision;
1859 const ledger = try revision.AccountingV1.create(testing.allocator, .{
1860 .allowance = revision.WorkVector.uniform(std.math.maxInt(u64)),
1861 .workspace = std.math.maxInt(u64),
1862 .events = 16,
1863 }, &.{
1864 .{ .name = fusion_planning_pass_name, .version = 1 },
1865 .{ .name = "accy-choir-fusion-test-probe", .version = 1 },
1866 });
1867 defer ledger.destroy();
1868 var cache = try passes.AnalysisCache.initAccounted(testing.allocator, null, ledger, .{}, 2);
1869 defer cache.deinit();
1870 try testing.expectEqual(
1871 passes.PassResult.success,
1872 manager.runWithAnalysisCache(op, ctx, &cache, .{}),
1873 );
1874 try ledger.producersComplete();
1875 try checkFusionStorage(op, ctx, &cache);
1876 }
1877
1878 fn checkFusionStorage(op: *ir.Operation, ctx: *ir.Context, cache: *passes.AnalysisCache) !void {
1879 const bounds = try fusionAnalysisWork(.{ .operation = op });
1880 const bytes = try testing.allocator.alloc(u8, @intCast(bounds.workspace));
1881 defer testing.allocator.free(bytes);
1882 var storage = alloc_fixed.Tracked.init(bytes);
1883 var pass_ctx = passes.PassContext.init(op, ctx, storage.allocator(), cache);
1884 defer pass_ctx.deinit();
1885 const ptr = try computeFusionPlanAnalysis(&pass_ctx, op);
1886 defer cleanupFusionPlanAnalysis(ptr, storage.allocator());
1887 const analysis: *FusionPlanAnalysis = @ptrCast(@alignCast(ptr));
1888 const reference = try getFusionPlanAnalysis(&pass_ctx, op);
1889 try testing.expectEqual(reference.fused_op_count, analysis.fused_op_count);
1890 try testing.expectEqual(reference.clusterCount(), analysis.clusterCount());
1891 for (reference.clusters.items, analysis.clusters.items) |expected, actual| {
1892 try testing.expectEqual(expected.kind, actual.kind);
1893 try testing.expectEqualSlices(*ir.Operation, expected.ops, actual.ops);
1894 }
1895 try testing.expectEqualSlices(*ir.Operation, reference.elided.items, analysis.elided.items);
1896 try testing.expect(!storage.exhausted);
1897 try testing.expect(storage.status().high_water_bytes <= bounds.workspace);
1898 try testing.expect(storage.status().high_water_bytes >= @sizeOf(FusionPlanAnalysis));
1899 }
1900
1901 test "fusion planning contract scales through map growth and rejects overflow" {
1902 for ([_]usize{ 1, 2, 6, 7, 16, 64 }) |count| {
1903 try checkFusionChain(count);
1904 }
1905 try testing.expectError(error.WorkOverflow, work.hashMapCapacity(std.math.maxInt(u64)));
1906 try testing.expectError(
1907 error.WorkOverflow,
1908 work.arrayListGrowth(FusionCluster, std.math.maxInt(u64)),
1909 );
1910 try testing.expectError(
1911 error.WorkOverflow,
1912 (FusionWork{ .rounds = std.math.maxInt(u64) }).bounds(),
1913 );
1914 }
1915
1916 fn checkFusionChain(count: usize) !void {
1917 var builder = try semantic.Builder.init(
1918 testing.allocator,
1919 semantic.Builder.ContextLimits.standard,
1920 );
1921 defer builder.deinit();
1922 const typ = try builder.tensor(.f32, &.{ 2, 3 });
1923 var function = try builder.beginFunction("fusion_chain", &.{ typ, typ }, &.{typ});
1924 var result = function.parameter(0);
1925 for (0..count) |_| result = try function.add(result, function.parameter(1));
1926 try function.return_(&.{result});
1927 try function.finish();
1928 const module = try builder.finish();
1929 defer module.deinit();
1930 var manager = passes.PassManager.init(testing.allocator);
1931 defer manager.deinit();
1932 try manager.addPass(fusionPlanningPass());
1933 try manager.addPass(FusionProbe.pass());
1934 FusionProbe.reset();
1935 try runAccountedFusion(&manager, module.choir_module, module.context());
1936 try testing.expectEqual(@as(usize, if (count > 1) 1 else 0), FusionProbe.cluster_count);
1937 try testing.expectEqual(if (count > 1) count else 0, FusionProbe.fused_op_count);
1938 try checkFusionAdmission(module.choir_module, module.context());
1939 }
1940
1941 fn checkFusionAdmission(op: *ir.Operation, ctx: *ir.Context) !void {
1942 const revision = choir.product.revision;
1943 const bounds = try fusionAnalysisWork(.{ .operation = op });
1944 const shape = shape_analysis.shape_layout_analysis_descriptor.work_contract.?;
1945 const shape_bounds = try shape.estimate(.{ .operation = op });
1946 const analysis_charge = try work.add(
1947 bounds.work.structural_visits,
1948 shape_bounds.work.structural_visits,
1949 );
1950 const charge = try work.add(analysis_charge, 1);
1951 for ([_]i8{ -1, 0, 1 }) |offset| {
1952 var allowance = revision.WorkVector.uniform(std.math.maxInt(u64));
1953 allowance.structural_visits = @intCast(@as(i128, charge) + offset);
1954 const ledger = try revision.AccountingV1.create(testing.allocator, .{
1955 .allowance = allowance,
1956 .workspace = std.math.maxInt(u64),
1957 .events = 8,
1958 }, &.{.{ .name = fusion_planning_pass_name, .version = 1 }});
1959 defer ledger.destroy();
1960 var cache = try passes.AnalysisCache.initAccounted(
1961 testing.allocator,
1962 null,
1963 ledger,
1964 .{},
1965 2,
1966 );
1967 defer cache.deinit();
1968 var manager = passes.PassManager.init(testing.allocator);
1969 defer manager.deinit();
1970 try manager.addPass(fusionPlanningPass());
1971 const result = manager.runWithAnalysisCache(op, ctx, &cache, .{});
1972 if (offset < 0) {
1973 try testing.expectEqual(passes.PassResult.failure, result);
1974 try testing.expectEqual(revision.receipt.Outcome.exhausted, ledger.view().outcome);
1975 try testing.expectEqual(@as(usize, 0), cache.entries.count());
1976 } else {
1977 try testing.expectEqual(passes.PassResult.success, result);
1978 try ledger.producersComplete();
1979 try testing.expectEqual(@as(usize, 2), cache.entries.count());
1980 }
1981 }
1982 }