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 }