lib/accy/src/preparation/folding.zig

daab053ee43316e1809a84551d573ddd1e5bf3d2

  1 const std = @import("std");
  2 const choir_abi = @import("choir_abi");
  3 const choir = @import("choir");
  4 const accy_root = @import("../root.zig");
  5 const accy_choir = @import("../choir/root.zig");
  6 const canonicalization = @import("canonicalization.zig");
  7 const dialect_mod = accy_choir.dialect;
  8 const shape_analysis = @import("shape/root.zig");
  9 
 10 const ir = choir.ir;
 11 const rewrite = ir.rewrite;
 12 const passes = choir.passes;
 13 const work = passes.pass.work;
 14 const DType = choir_abi.DType;
 15 
 16 pub const constant_folding_pass_name = "accy-choir-constant-fold";
 17 pub const constant_folding_pass_description =
 18     "Fold Accy Choir operations with constant operands";
 19 
 20 pub fn constantFoldingPass() passes.Pass {
 21     return .{
 22         .name = constant_folding_pass_name,
 23         .description = constant_folding_pass_description,
 24         .run_fn = runConstantFoldingPass,
 25         .work_contract = .{
 26             .identity = .{ .name = constant_folding_pass_name, .version = 1 },
 27             .estimate = foldingWork,
 28         },
 29     };
 30 }
 31 
 32 const FoldWork = struct {
 33     candidates: u64 = 0,
 34     payload: u64 = 0,
 35     shape_key: u64 = 0,
 36     list_bytes: u64 = 0,
 37     broadcasts: bool = false,
 38 
 39     fn visit(self: *FoldWork, op: *ir.Operation) !ir.WalkResult {
 40         const name = op.name.name;
 41         const broadcast = std.mem.eql(
 42             u8,
 43             name,
 44             dialect_mod.AccyDialect.BroadcastInDimOp.operation_name,
 45         );
 46         const reshape = std.mem.eql(u8, name, dialect_mod.AccyDialect.ReshapeOp.operation_name);
 47         if (broadcast or reshape or binaryFoldKind(name) != null or unaryFoldKind(name) != null) {
 48             self.candidates = try work.add(self.candidates, 1);
 49         }
 50         self.broadcasts = self.broadcasts or broadcast;
 51         if (std.mem.eql(u8, name, dialect_mod.AccyDialect.ConstantOp.operation_name)) {
 52             if ((dialect_mod.AccyDialect.ConstantOp{ .op = op }).getPayload()) |payload| {
 53                 self.payload = @max(self.payload, payload.len);
 54             }
 55         }
 56         for (op.results.items) |*result| {
 57             if (result.type.getDialectParamKey()) |key| {
 58                 self.shape_key = @max(self.shape_key, key.len);
 59             }
 60         }
 61         for (op.getOperandValues()) |operand| {
 62             if (operand.type.getDialectParamKey()) |key| {
 63                 self.shape_key = @max(self.shape_key, key.len);
 64             }
 65         }
 66         if (broadcast) {
 67             if (op.getAttr("broadcast_dims")) |attr| {
 68                 if (attr.cast(ir.Attribute.DialectAttr)) |value| {
 69                     self.list_bytes = try work.add(self.list_bytes, value.payload.len);
 70                 }
 71             }
 72         }
 73         return .advance;
 74     }
 75 };
 76 
 77 fn foldingWork(input: work.Input) !work.Bounds {
 78     const counts = try work.Census.inspect(input.operation);
 79     var facts: FoldWork = .{};
 80     _ = try input.operation.walk(.{ .order = .pre_order }, &facts, FoldWork.visit);
 81     const payload = @max(facts.payload, if (facts.broadcasts) broadcast_fold_payload_limit else 0);
 82     const payload_bytes = try work.multiply(facts.candidates, payload);
 83     const dimensions = try work.multiply(facts.shape_key, @sizeOf(usize));
 84     const per_fold = try work.add(try work.add(payload, dimensions), 2 * @alignOf(usize));
 85     const queues = try work.multiply(2, try work.arrayListGrowth(*ir.Operation, facts.candidates));
 86     const temporary = try work.add(facts.list_bytes, try work.multiply(facts.candidates, per_fold));
 87     const bytes = try work.add(queues, temporary);
 88     const visits = try work.add(counts.atoms, counts.input_bytes);
 89     const uses = try work.add(try work.add(counts.values, counts.operands), 1);
 90     const traversal = try work.multiply(64, try work.multiply(try work.add(visits, 1), uses));
 91     const processing = try work.multiply(16, try work.multiply(
 92         payload_bytes,
 93         try work.add(facts.shape_key, 1),
 94     ));
 95     const nodes = try work.multiply(
 96         facts.candidates,
 97         @sizeOf(ir.Operation) + @sizeOf(ir.Value) + 64,
 98     );
 99     return .{
100         .work = .{
101             .input_bytes = counts.input_bytes,
102             .output_bytes = try work.add(payload_bytes, nodes),
103             .structural_visits = try work.add(traversal, processing),
104             .allocation_capacity = bytes,
105         },
106         .workspace = bytes,
107     };
108 }
109 
110 fn runConstantFoldingPass(pass_ctx: *passes.PassContext) passes.PassResult {
111     const analysis = shape_analysis.getShapeLayoutAnalysis(pass_ctx, pass_ctx.op) catch return .failure;
112 
113     var rewriter = rewrite.PatternRewriter.init(pass_ctx.allocator, pass_ctx.ir_ctx);
114     defer rewriter.deinit();
115 
116     var folded_count: usize = 0;
117     foldOnOp(pass_ctx.ir_ctx, pass_ctx.op, analysis, &rewriter, &folded_count) catch return .failure;
118     if (folded_count == 0) {
119         pass_ctx.preserveAllAnalyses();
120     } else {
121         rewriter.finalize(pass_ctx.op);
122         pass_ctx.markModified();
123     }
124     return .success;
125 }
126 
127 fn foldOnOp(
128     ctx: *ir.Context,
129     op: *ir.Operation,
130     analysis: *const shape_analysis.ShapeLayoutAnalysis,
131     rewriter: *rewrite.PatternRewriter,
132     folded_count: *usize,
133 ) !void {
134     for (op.regions.items) |*region| {
135         var block_iter = region.getBlocks();
136         while (block_iter.next()) |block| {
137             var current: ?*ir.Operation = @ptrCast(@alignCast(block.operations.head));
138             while (current) |current_op| {
139                 const next = current_op.next_op;
140                 if (current_op.regions.items.len > 0) {
141                     try foldOnOp(ctx, current_op, analysis, rewriter, folded_count);
142                 }
143                 if (try foldOp(ctx, current_op, analysis, rewriter)) {
144                     folded_count.* += 1;
145                 }
146                 current = next;
147             }
148         }
149     }
150 }
151 
152 fn foldOp(
153     ctx: *ir.Context,
154     op: *ir.Operation,
155     analysis: *const shape_analysis.ShapeLayoutAnalysis,
156     rewriter: *rewrite.PatternRewriter,
157 ) !bool {
158     if (try foldReshapeConstant(ctx, op, analysis, rewriter)) return true;
159     if (try foldBroadcastInDimConstant(ctx, op, analysis, rewriter)) return true;
160     if (try foldElementwiseConstant(ctx, op, analysis, rewriter)) return true;
161     return false;
162 }
163 
164 const broadcast_fold_payload_limit: usize = 64 * 1024;
165 
166 fn foldBroadcastInDimConstant(
167     ctx: *ir.Context,
168     op: *ir.Operation,
169     analysis: *const shape_analysis.ShapeLayoutAnalysis,
170     rewriter: *rewrite.PatternRewriter,
171 ) !bool {
172     if (!std.mem.eql(u8, op.name.name, dialect_mod.AccyDialect.BroadcastInDimOp.operation_name)) {
173         return false;
174     }
175     if (op.getNumOperands() != 1 or op.getNumResults() != 1) return false;
176 
177     const input = op.getOperand(0) orelse return false;
178     const result = op.getResult(0) orelse return false;
179     const input_const = constantDefiningOp(input) orelse return false;
180     const input_info = analysis.get(input) orelse return false;
181     const result_info = analysis.get(result) orelse return false;
182 
183     if (input_info.dtype != result_info.dtype) return false;
184     const input_bytes = expectedPayloadBytes(input_info) orelse return false;
185     const result_bytes = expectedPayloadBytes(result_info) orelse return false;
186     if (result_bytes > broadcast_fold_payload_limit) return false;
187 
188     const payload = (dialect_mod.AccyDialect.ConstantOp{ .op = input_const }).getPayload() orelse return false;
189     if (payload.len != input_bytes) return false;
190 
191     const broadcast_dims = (try canonicalization.readI64ListAttrAlloc(
192         rewriter.allocator,
193         op,
194         "broadcast_dims",
195         dialect_mod.AccyDialect.BroadcastInDimOp.dialectAttrName("broadcast_dims"),
196     )) orelse return false;
197     defer rewriter.allocator.free(broadcast_dims);
198     if (broadcast_dims.len != input_info.dims.len) return false;
199     for (broadcast_dims) |mapped| {
200         if (mapped < 0 or @as(usize, @intCast(mapped)) >= result_info.dims.len) return false;
201     }
202 
203     const folded_payload = (try foldBroadcastPayload(
204         rewriter.allocator,
205         result_info.dtype,
206         payload,
207         input_info.dims,
208         result_info.dims,
209         broadcast_dims,
210     )) orelse return false;
211     defer rewriter.allocator.free(folded_payload);
212 
213     const folded = try createConstantBefore(ctx, rewriter, op, folded_payload, result.type);
214     try rewriter.replaceOpWithValue(op, folded.getResult());
215     return true;
216 }
217 
218 fn foldBroadcastPayload(
219     allocator: std.mem.Allocator,
220     dtype: DType,
221     input_payload: []const u8,
222     input_dims: []const i64,
223     result_dims: []const i64,
224     broadcast_dims: []const i64,
225 ) !?[]u8 {
226     const element_size: usize = dtype.sizeOf();
227     var result_elements: usize = 1;
228     for (result_dims) |dim| {
229         if (dim < 0) return null;
230         result_elements = std.math.mul(usize, result_elements, @intCast(dim)) catch return null;
231     }
232 
233     const out = try allocator.alloc(u8, result_elements * element_size);
234     errdefer allocator.free(out);
235 
236     const coords = try allocator.alloc(usize, result_dims.len);
237     defer allocator.free(coords);
238     @memset(coords, 0);
239 
240     var out_index: usize = 0;
241     while (out_index < result_elements) : (out_index += 1) {
242         var in_index: usize = 0;
243         for (input_dims, broadcast_dims) |in_dim, mapped| {
244             if (in_dim < 0) return null;
245             const extent: usize = @intCast(in_dim);
246             const coord = if (extent <= 1) 0 else coords[@intCast(mapped)];
247             in_index = in_index * extent + coord;
248         }
249         @memcpy(
250             out[out_index * element_size ..][0..element_size],
251             input_payload[in_index * element_size ..][0..element_size],
252         );
253 
254         var axis = result_dims.len;
255         while (axis > 0) {
256             axis -= 1;
257             coords[axis] += 1;
258             if (coords[axis] < @as(usize, @intCast(result_dims[axis]))) break;
259             coords[axis] = 0;
260         }
261     }
262     return out;
263 }
264 
265 fn foldReshapeConstant(
266     ctx: *ir.Context,
267     op: *ir.Operation,
268     analysis: *const shape_analysis.ShapeLayoutAnalysis,
269     rewriter: *rewrite.PatternRewriter,
270 ) !bool {
271     if (!std.mem.eql(u8, op.name.name, dialect_mod.AccyDialect.ReshapeOp.operation_name)) {
272         return false;
273     }
274     if (op.getNumOperands() != 1 or op.getNumResults() != 1) return false;
275 
276     const input = op.getOperand(0) orelse return false;
277     const result = op.getResult(0) orelse return false;
278     const input_const = constantDefiningOp(input) orelse return false;
279     const input_info = analysis.get(input) orelse return false;
280     const result_info = analysis.get(result) orelse return false;
281 
282     if (!input_info.hasStaticLayout() or !result_info.hasStaticLayout()) return false;
283     if (input_info.dtype != result_info.dtype) return false;
284     if (input_info.element_count.? != result_info.element_count.?) return false;
285 
286     const payload = (dialect_mod.AccyDialect.ConstantOp{ .op = input_const }).getPayload() orelse return false;
287     const expected_bytes = std.math.mul(
288         u64,
289         result_info.element_count.?,
290         @as(u64, result_info.dtype.sizeOf()),
291     ) catch return false;
292     if (expected_bytes != payload.len) return false;
293 
294     const folded = try createConstantBefore(ctx, rewriter, op, payload, result.type);
295     try rewriter.replaceOpWithValue(op, folded.getResult());
296     return true;
297 }
298 
299 const BinaryFoldKind = enum {
300     add,
301     sub,
302     mul,
303 };
304 
305 const UnaryFoldKind = enum {
306     neg,
307     floor,
308     round,
309     trunc,
310 };
311 
312 fn foldElementwiseConstant(
313     ctx: *ir.Context,
314     op: *ir.Operation,
315     analysis: *const shape_analysis.ShapeLayoutAnalysis,
316     rewriter: *rewrite.PatternRewriter,
317 ) !bool {
318     if (binaryFoldKind(op.name.name)) |kind| {
319         return foldBinaryElementwiseConstant(ctx, op, analysis, rewriter, kind);
320     }
321     if (unaryFoldKind(op.name.name)) |kind| {
322         return foldUnaryElementwiseConstant(ctx, op, analysis, rewriter, kind);
323     }
324     return false;
325 }
326 
327 fn binaryFoldKind(name: []const u8) ?BinaryFoldKind {
328     if (std.mem.eql(u8, name, dialect_mod.AccyDialect.AddOp.operation_name)) return .add;
329     if (std.mem.eql(u8, name, dialect_mod.AccyDialect.SubOp.operation_name)) return .sub;
330     if (std.mem.eql(u8, name, dialect_mod.AccyDialect.MulOp.operation_name)) return .mul;
331     return null;
332 }
333 
334 fn unaryFoldKind(name: []const u8) ?UnaryFoldKind {
335     if (std.mem.eql(u8, name, dialect_mod.AccyDialect.NegOp.operation_name)) return .neg;
336     if (std.mem.eql(u8, name, dialect_mod.AccyDialect.FloorOp.operation_name)) return .floor;
337     if (std.mem.eql(u8, name, dialect_mod.AccyDialect.RoundOp.operation_name)) return .round;
338     if (std.mem.eql(u8, name, dialect_mod.AccyDialect.TruncOp.operation_name)) return .trunc;
339     return null;
340 }
341 
342 fn foldBinaryElementwiseConstant(
343     ctx: *ir.Context,
344     op: *ir.Operation,
345     analysis: *const shape_analysis.ShapeLayoutAnalysis,
346     rewriter: *rewrite.PatternRewriter,
347     kind: BinaryFoldKind,
348 ) !bool {
349     if (op.getNumOperands() != 2 or op.getNumResults() != 1) return false;
350     const lhs = op.getOperand(0) orelse return false;
351     const rhs = op.getOperand(1) orelse return false;
352     const result = op.getResult(0) orelse return false;
353     const lhs_const = constantDefiningOp(lhs) orelse return false;
354     const rhs_const = constantDefiningOp(rhs) orelse return false;
355     const result_info = analysis.get(result) orelse return false;
356     const expected_bytes = expectedPayloadBytes(result_info) orelse return false;
357 
358     const lhs_payload = (dialect_mod.AccyDialect.ConstantOp{ .op = lhs_const }).getPayload() orelse return false;
359     const rhs_payload = (dialect_mod.AccyDialect.ConstantOp{ .op = rhs_const }).getPayload() orelse return false;
360     if (lhs_payload.len != expected_bytes or rhs_payload.len != expected_bytes) return false;
361 
362     const folded_payload = (try foldBinaryPayload(
363         rewriter.allocator,
364         kind,
365         result_info.dtype,
366         lhs_payload,
367         rhs_payload,
368         result_info.element_count.?,
369     )) orelse return false;
370     defer rewriter.allocator.free(folded_payload);
371 
372     const folded = try createConstantBefore(ctx, rewriter, op, folded_payload, result.type);
373     try rewriter.replaceOpWithValue(op, folded.getResult());
374     return true;
375 }
376 
377 fn foldUnaryElementwiseConstant(
378     ctx: *ir.Context,
379     op: *ir.Operation,
380     analysis: *const shape_analysis.ShapeLayoutAnalysis,
381     rewriter: *rewrite.PatternRewriter,
382     kind: UnaryFoldKind,
383 ) !bool {
384     if (op.getNumOperands() != 1 or op.getNumResults() != 1) return false;
385     const input = op.getOperand(0) orelse return false;
386     const result = op.getResult(0) orelse return false;
387     const input_const = constantDefiningOp(input) orelse return false;
388     const result_info = analysis.get(result) orelse return false;
389     const expected_bytes = expectedPayloadBytes(result_info) orelse return false;
390 
391     const input_payload = (dialect_mod.AccyDialect.ConstantOp{ .op = input_const }).getPayload() orelse return false;
392     if (input_payload.len != expected_bytes) return false;
393 
394     const folded_payload = (try foldUnaryPayload(
395         rewriter.allocator,
396         kind,
397         result_info.dtype,
398         input_payload,
399         result_info.element_count.?,
400     )) orelse return false;
401     defer rewriter.allocator.free(folded_payload);
402 
403     const folded = try createConstantBefore(ctx, rewriter, op, folded_payload, result.type);
404     try rewriter.replaceOpWithValue(op, folded.getResult());
405     return true;
406 }
407 
408 fn expectedPayloadBytes(info: shape_analysis.TensorInfo) ?usize {
409     if (!info.hasStaticLayout()) return null;
410     const bytes = std.math.mul(
411         u64,
412         info.element_count.?,
413         @as(u64, info.dtype.sizeOf()),
414     ) catch return null;
415     if (bytes > std.math.maxInt(usize)) return null;
416     return @intCast(bytes);
417 }
418 
419 fn foldBinaryPayload(
420     allocator: std.mem.Allocator,
421     kind: BinaryFoldKind,
422     dtype: DType,
423     lhs_payload: []const u8,
424     rhs_payload: []const u8,
425     element_count: u64,
426 ) !?[]u8 {
427     if (element_count > std.math.maxInt(usize)) return null;
428     const count: usize = @intCast(element_count);
429     return switch (dtype) {
430         .f32 => try foldBinaryPayloadTyped(f32, allocator, kind, lhs_payload, rhs_payload, count),
431         .f64 => try foldBinaryPayloadTyped(f64, allocator, kind, lhs_payload, rhs_payload, count),
432         .i32 => try foldBinaryPayloadTyped(i32, allocator, kind, lhs_payload, rhs_payload, count),
433         .i64 => try foldBinaryPayloadTyped(i64, allocator, kind, lhs_payload, rhs_payload, count),
434         else => null,
435     };
436 }
437 
438 fn foldUnaryPayload(
439     allocator: std.mem.Allocator,
440     kind: UnaryFoldKind,
441     dtype: DType,
442     input_payload: []const u8,
443     element_count: u64,
444 ) !?[]u8 {
445     if (element_count > std.math.maxInt(usize)) return null;
446     const count: usize = @intCast(element_count);
447     return switch (dtype) {
448         .f32 => try foldUnaryPayloadTyped(f32, allocator, kind, input_payload, count),
449         .f64 => try foldUnaryPayloadTyped(f64, allocator, kind, input_payload, count),
450         .i32 => try foldUnaryPayloadTyped(i32, allocator, kind, input_payload, count),
451         .i64 => try foldUnaryPayloadTyped(i64, allocator, kind, input_payload, count),
452         else => null,
453     };
454 }
455 
456 fn foldBinaryPayloadTyped(
457     comptime T: type,
458     allocator: std.mem.Allocator,
459     kind: BinaryFoldKind,
460     lhs_payload: []const u8,
461     rhs_payload: []const u8,
462     element_count: usize,
463 ) !?[]u8 {
464     const byte_count = std.math.mul(usize, element_count, @sizeOf(T)) catch return null;
465     const payload = try allocator.alloc(u8, byte_count);
466     var keep_payload = false;
467     defer if (!keep_payload) allocator.free(payload);
468     for (0..element_count) |i| {
469         const lhs = readPayloadScalar(T, lhs_payload, i);
470         const rhs = readPayloadScalar(T, rhs_payload, i);
471         const folded = foldBinaryScalar(T, kind, lhs, rhs) orelse return null;
472         writePayloadScalar(T, payload, i, folded);
473     }
474     keep_payload = true;
475     return payload;
476 }
477 
478 fn foldUnaryPayloadTyped(
479     comptime T: type,
480     allocator: std.mem.Allocator,
481     kind: UnaryFoldKind,
482     input_payload: []const u8,
483     element_count: usize,
484 ) !?[]u8 {
485     const byte_count = std.math.mul(usize, element_count, @sizeOf(T)) catch return null;
486     const payload = try allocator.alloc(u8, byte_count);
487     var keep_payload = false;
488     defer if (!keep_payload) allocator.free(payload);
489     for (0..element_count) |i| {
490         const input = readPayloadScalar(T, input_payload, i);
491         const folded = foldUnaryScalar(T, kind, input) orelse return null;
492         writePayloadScalar(T, payload, i, folded);
493     }
494     keep_payload = true;
495     return payload;
496 }
497 
498 fn foldBinaryScalar(comptime T: type, kind: BinaryFoldKind, lhs: T, rhs: T) ?T {
499     return switch (@typeInfo(T)) {
500         .float => switch (kind) {
501             .add => lhs + rhs,
502             .sub => lhs - rhs,
503             .mul => lhs * rhs,
504         },
505         .int => checkedBinaryInt(T, kind, lhs, rhs),
506         else => null,
507     };
508 }
509 
510 fn foldUnaryScalar(comptime T: type, kind: UnaryFoldKind, value: T) ?T {
511     return switch (@typeInfo(T)) {
512         .float => switch (kind) {
513             .neg => -value,
514             .floor => @floor(value),
515             .round => @round(value),
516             .trunc => @trunc(value),
517         },
518         .int => switch (kind) {
519             .neg => checkedNegInt(T, value),
520             .floor, .round, .trunc => null,
521         },
522         else => null,
523     };
524 }
525 
526 fn checkedBinaryInt(comptime T: type, kind: BinaryFoldKind, lhs: T, rhs: T) ?T {
527     const folded: i128 = switch (kind) {
528         .add => @as(i128, lhs) + @as(i128, rhs),
529         .sub => @as(i128, lhs) - @as(i128, rhs),
530         .mul => @as(i128, lhs) * @as(i128, rhs),
531     };
532     if (folded < @as(i128, std.math.minInt(T))) return null;
533     if (folded > @as(i128, std.math.maxInt(T))) return null;
534     return @intCast(folded);
535 }
536 
537 fn checkedNegInt(comptime T: type, value: T) ?T {
538     if (value == std.math.minInt(T)) return null;
539     return -value;
540 }
541 
542 fn readPayloadScalar(comptime T: type, payload: []const u8, index: usize) T {
543     const start = index * @sizeOf(T);
544     var value: T = undefined;
545     @memcpy(std.mem.asBytes(&value), payload[start..][0..@sizeOf(T)]);
546     return value;
547 }
548 
549 fn writePayloadScalar(comptime T: type, payload: []u8, index: usize, value: T) void {
550     const start = index * @sizeOf(T);
551     @memcpy(payload[start..][0..@sizeOf(T)], std.mem.asBytes(&value));
552 }
553 
554 fn constantDefiningOp(value: *ir.Value) ?*ir.Operation {
555     const def_any = value.getDefiningOp() orelse return null;
556     const def_op: *ir.Operation = @ptrCast(@alignCast(def_any));
557     if (!std.mem.eql(u8, def_op.name.name, dialect_mod.AccyDialect.ConstantOp.operation_name)) {
558         return null;
559     }
560     return def_op;
561 }
562 
563 fn createConstantBefore(
564     ctx: *ir.Context,
565     rewriter: *rewrite.PatternRewriter,
566     before: *ir.Operation,
567     payload: []const u8,
568     result_type: ir.Type,
569 ) !dialect_mod.AccyDialect.ConstantOp {
570     rewriter.setInsertionPointBefore(before);
571     var state = ir.Operation.State.init(
572         dialect_mod.AccyDialect.ConstantOp.operation_name,
573         before.location,
574     );
575     state.addTypes(&.{result_type});
576     const op = try rewriter.create(state);
577     const payload_attr = try ctx.getDialectAttr(dialect_mod.AccyDialect.ConstantOp.payload_attr_name, payload);
578     try op.setAttr("payload", payload_attr);
579     return .{ .op = op };
580 }
581 
582 const testing = std.testing;
583 const semantic = accy_choir.semantic;
584 
585 fn readSymbolName(func: *ir.Operation) ?[]const u8 {
586     return ir.SymbolTable.getSymbolName(func);
587 }
588 
589 fn findOpNamedInBlock(block: *ir.Block, name: []const u8) ?*ir.Operation {
590     var iter = block.operations.head;
591     while (iter) |op_ptr| {
592         const op: *ir.Operation = @ptrCast(@alignCast(op_ptr));
593         if (std.mem.eql(u8, op.name.name, name)) return op;
594         iter = op.next_op;
595     }
596     return null;
597 }
598 
599 fn foldingFixture(shape: []const i64, copies: u32) !*semantic.SemanticModule {
600     var builder = try semantic.Builder.init(testing.allocator, .standard);
601     defer builder.deinit();
602     const scalar = try builder.tensor(.f32, &.{1});
603     const vector = try builder.tensor(.f32, shape);
604     var function = try builder.beginFunction("fold_accounting", &.{}, &.{vector});
605     const value = [_]f32{2};
606     var result = try function.constant(scalar, std.mem.sliceAsBytes(&value));
607     result = try function.broadcastInDim(result, vector, shape, &.{0});
608     for (0..copies) |_| result = try function.neg(result);
609     try function.return_(&.{result});
610     try function.finish();
611     return builder.finish();
612 }
613 
614 fn expectFoldingPayload(module: *semantic.SemanticModule, span: u32, copies: u32) !void {
615     const body = module.choir_module.getRegion(0).?.getEntryBlock().?;
616     const function = ir.inspection.functionByNameInBlock(body, "fold_accounting").?;
617     const block = function.getRegion(0).?.getEntryBlock().?;
618     const ret = findOpNamedInBlock(block, "func.return").?;
619     const constant = constantDefiningOp(ret.getOperand(0).?) orelse
620         return error.TestExpectedConstant;
621     const payload = (dialect_mod.AccyDialect.ConstantOp{ .op = constant }).getPayload().?;
622     try testing.expectEqual(@as(usize, span) * @sizeOf(f32), payload.len);
623     const expected: f32 = if (copies % 2 == 0) 2 else -2;
624     for (0..span) |index| try testing.expectEqual(expected, readPayloadScalar(f32, payload, index));
625 }
626 
627 fn checkFoldingAccounting(admitted: bool) !void {
628     const allocator = testing.allocator;
629     const revision = choir.product.revision;
630     const module = try foldingFixture(&.{17}, 3);
631     defer module.deinit();
632     const root = module.choir_module;
633     const before = try choir.bytecode.encodeModule(allocator, root);
634     defer allocator.free(before);
635     const bounds = try foldingWork(.{ .operation = root });
636     var allowance = revision.WorkVector.uniform(1 << 40);
637     if (!admitted) allowance.structural_visits = bounds.work.structural_visits - 1;
638     const ledger = try revision.AccountingV1.create(allocator, .{
639         .allowance = allowance,
640         .workspace = 1 << 24,
641         .events = 8,
642     }, &.{.{ .name = constant_folding_pass_name, .version = 1 }});
643     defer ledger.destroy();
644     var cache = try passes.AnalysisCache.initAccounted(
645         allocator,
646         null,
647         ledger,
648         .{ .context = module.context() },
649         1,
650     );
651     defer cache.deinit();
652     var manager = passes.PassManager.init(allocator);
653     defer manager.deinit();
654     try manager.addPass(constantFoldingPass());
655     const result = manager.runWithAnalysisCache(root, module.context(), &cache, .{});
656     try testing.expectEqual(if (admitted) passes.PassResult.success else .failure, result);
657     if (admitted) {
658         try ledger.producersComplete();
659         try testing.expect(!ledger.view().missing_work_contract);
660         try testing.expectEqual(@as(u64, 1), ledger.view().executed.counters.pass_runs);
661         try expectFoldingPayload(module, 17, 3);
662     } else {
663         try testing.expectEqual(.exhausted, ledger.view().outcome);
664         try testing.expectEqual(@as(u64, 0), manager.stats.pass_runs);
665         try testing.expectEqual(@as(usize, 0), cache.entries.count());
666         const after = try choir.bytecode.encodeModule(allocator, root);
667         defer allocator.free(after);
668         try testing.expectEqualSlices(u8, before, after);
669     }
670 }
671 
672 test "constant folding accounts real output and refuses below its charge before mutation" {
673     try checkFoldingAccounting(false);
674     try checkFoldingAccounting(true);
675 }
676 
677 fn checkFoldingStorage(module: *semantic.SemanticModule, span: u32, copies: u32) !void {
678     const allocator = testing.allocator;
679     var cache = passes.AnalysisCache.init(allocator, null);
680     defer cache.deinit();
681     var context = passes.PassContext.init(module.choir_module, module.context(), allocator, &cache);
682     defer context.deinit();
683     _ = try shape_analysis.getShapeLayoutAnalysis(&context, module.choir_module);
684     const bounds = try foldingWork(.{ .operation = module.choir_module });
685     const bytes = try allocator.alloc(u8, @intCast(bounds.workspace));
686     defer allocator.free(bytes);
687     var storage = @import("alloc_fixed").Tracked.init(bytes);
688     context.allocator = storage.allocator();
689     defer context.allocator = allocator;
690     try testing.expectEqual(.success, runConstantFoldingPass(&context));
691     try testing.expect(!storage.exhausted);
692     try testing.expect(storage.status().high_water_bytes <= bounds.workspace);
693     try testing.expect(storage.status().high_water_bytes > 0);
694     try expectFoldingPayload(module, span, copies);
695 }
696 
697 test "constant folding scratch bound covers generated payloads and rewrite queue growth" {
698     const shapes = [_][]const i64{
699         &.{1},
700         &.{ 2, 3 },
701         &.{ 128, 128 },
702         &.{ 1, 1, 1, 1, 1, 1, 1, 1 },
703     };
704     for (shapes) |shape| {
705         var elements: u32 = 1;
706         for (shape) |dim| elements = try std.math.mul(u32, elements, @intCast(dim));
707         for ([_]u32{ 1, 17 }) |copies| {
708             const module = try foldingFixture(shape, copies);
709             defer module.deinit();
710             try checkFoldingStorage(module, elements, copies);
711         }
712     }
713 }
714 
715 test "constant folding scratch bound includes original payloads above the broadcast limit" {
716     var builder = try semantic.Builder.init(testing.allocator, .standard);
717     defer builder.deinit();
718     const tensor = try builder.tensor(.f32, &.{20000});
719     const values: [20000]f32 = @splat(2);
720     var function = try builder.beginFunction("fold_accounting", &.{}, &.{tensor});
721     const source = try function.constant(tensor, std.mem.sliceAsBytes(&values));
722     const negated = try function.neg(source);
723     try function.return_(&.{negated});
724     try function.finish();
725     const module = try builder.finish();
726     defer module.deinit();
727     try checkFoldingStorage(module, values.len, 1);
728 }
729 
730 test "constant folding folds reshape of accy constant" {
731     const allocator = testing.allocator;
732 
733     var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
734     defer builder.deinit();
735     const f32_2x2 = try builder.tensor(.f32, &.{ 2, 2 });
736     const f32_4 = try builder.tensor(.f32, &.{4});
737     var payload: [16]u8 = undefined;
738     @as(*align(1) f32, @ptrCast(&payload[0])).* = 1.0;
739     @as(*align(1) f32, @ptrCast(&payload[4])).* = 2.0;
740     @as(*align(1) f32, @ptrCast(&payload[8])).* = 3.0;
741     @as(*align(1) f32, @ptrCast(&payload[12])).* = 4.0;
742 
743     var fb = try builder.beginFunction("fold_reshape_constant", &.{}, &.{f32_4});
744     const c = try fb.constant(f32_2x2, &payload);
745     const r = try fb.reshape(c, f32_4, &.{4});
746     try fb.return_(&.{r});
747     try fb.finish();
748     const module = try builder.finish();
749     defer module.deinit();
750 
751     const choir_mod = module.choir_module;
752     const ctx = module.context();
753     var pm = passes.PassManager.init(allocator);
754     defer pm.deinit();
755     try pm.addPass(shape_analysis.shapeLayoutPropagationPass());
756     try pm.addPass(constantFoldingPass());
757 
758     try testing.expectEqual(passes.PassResult.success, pm.run(choir_mod, ctx));
759     try testing.expectEqual(@as(u64, 2), pm.stats.pass_runs);
760     try testing.expectEqual(@as(u64, 1), pm.stats.passes_modified);
761     try testing.expectEqual(@as(u64, 1), pm.stats.analysis_hits);
762     try testing.expectEqual(@as(usize, 0), ir.inspection.countOperationsNamed(choir_mod, dialect_mod.AccyDialect.ReshapeOp.operation_name));
763     try testing.expectEqual(@as(usize, 2), ir.inspection.countOperationsNamed(choir_mod, dialect_mod.AccyDialect.ConstantOp.operation_name));
764 
765     const module_body = choir_mod.getRegion(0).?.getEntryBlock().?;
766     const func = ir.inspection.functionByNameInBlock(module_body, "fold_reshape_constant") orelse return error.TestExpectedFunc;
767     const entry = func.getRegion(0).?.getEntryBlock().?;
768     const ret = findOpNamedInBlock(entry, "func.return") orelse return error.TestExpectedReturn;
769     const folded_const = constantDefiningOp(ret.getOperand(0).?) orelse return error.TestExpectedConstant;
770     const folded_payload = (dialect_mod.AccyDialect.ConstantOp{ .op = folded_const }).getPayload() orelse return error.TestExpectedPayload;
771     try testing.expectEqualSlices(u8, &payload, folded_payload);
772     try testing.expectEqualStrings("f32,4", ret.getOperand(0).?.type.getDialectParamKey().?);
773 }
774 
775 test "constant folding folds broadcast_in_dim of accy constant" {
776     const allocator = testing.allocator;
777 
778     var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
779     defer builder.deinit();
780     const f32_2 = try builder.tensor(.f32, &.{2});
781     const f32_2x3 = try builder.tensor(.f32, &.{ 2, 3 });
782     const values = [_]f32{ 1.5, -2.0 };
783     const expected = [_]f32{ 1.5, 1.5, 1.5, -2.0, -2.0, -2.0 };
784 
785     var fb = try builder.beginFunction("fold_broadcast_constant", &.{}, &.{f32_2x3});
786     const c = try fb.constant(f32_2, std.mem.sliceAsBytes(values[0..]));
787     const b = try fb.broadcastInDim(c, f32_2x3, &.{ 2, 3 }, &.{0});
788     try fb.return_(&.{b});
789     try fb.finish();
790     const module = try builder.finish();
791     defer module.deinit();
792 
793     const choir_mod = module.choir_module;
794     const ctx = module.context();
795     var pm = passes.PassManager.init(allocator);
796     defer pm.deinit();
797     try pm.addPass(shape_analysis.shapeLayoutPropagationPass());
798     try pm.addPass(constantFoldingPass());
799 
800     try testing.expectEqual(passes.PassResult.success, pm.run(choir_mod, ctx));
801     try testing.expectEqual(@as(usize, 0), ir.inspection.countOperationsNamed(choir_mod, dialect_mod.AccyDialect.BroadcastInDimOp.operation_name));
802 
803     const module_body = choir_mod.getRegion(0).?.getEntryBlock().?;
804     const func = ir.inspection.functionByNameInBlock(module_body, "fold_broadcast_constant") orelse return error.TestExpectedFunc;
805     const entry = func.getRegion(0).?.getEntryBlock().?;
806     const ret = findOpNamedInBlock(entry, "func.return") orelse return error.TestExpectedReturn;
807     const folded_const = constantDefiningOp(ret.getOperand(0).?) orelse return error.TestExpectedConstant;
808     const folded_payload = (dialect_mod.AccyDialect.ConstantOp{ .op = folded_const }).getPayload() orelse return error.TestExpectedPayload;
809     try testing.expectEqualSlices(u8, std.mem.sliceAsBytes(expected[0..]), folded_payload);
810 }
811 
812 test "constant folding folds f32 add constants" {
813     const allocator = testing.allocator;
814 
815     var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
816     defer builder.deinit();
817     const f32_3 = try builder.tensor(.f32, &.{3});
818     const lhs = [_]f32{ 1.25, -2.5, 4.0 };
819     const rhs = [_]f32{ 3.75, 10.0, -1.5 };
820     const expected = [_]f32{ 5.0, 7.5, 2.5 };
821 
822     var fb = try builder.beginFunction("fold_add_constants", &.{}, &.{f32_3});
823     const l = try fb.constant(f32_3, std.mem.sliceAsBytes(lhs[0..]));
824     const r = try fb.constant(f32_3, std.mem.sliceAsBytes(rhs[0..]));
825     const sum = try fb.add(l, r);
826     try fb.return_(&.{sum});
827     try fb.finish();
828     const module = try builder.finish();
829     defer module.deinit();
830 
831     const choir_mod = module.choir_module;
832     const ctx = module.context();
833     var pm = passes.PassManager.init(allocator);
834     defer pm.deinit();
835     try pm.addPass(shape_analysis.shapeLayoutPropagationPass());
836     try pm.addPass(constantFoldingPass());
837 
838     try testing.expectEqual(passes.PassResult.success, pm.run(choir_mod, ctx));
839     try testing.expectEqual(@as(u64, 2), pm.stats.pass_runs);
840     try testing.expectEqual(@as(u64, 1), pm.stats.passes_modified);
841     try testing.expectEqual(@as(usize, 0), ir.inspection.countOperationsNamed(choir_mod, dialect_mod.AccyDialect.AddOp.operation_name));
842     try testing.expectEqual(@as(usize, 3), ir.inspection.countOperationsNamed(choir_mod, dialect_mod.AccyDialect.ConstantOp.operation_name));
843 
844     const module_body = choir_mod.getRegion(0).?.getEntryBlock().?;
845     const func = ir.inspection.functionByNameInBlock(module_body, "fold_add_constants") orelse return error.TestExpectedFunc;
846     const entry = func.getRegion(0).?.getEntryBlock().?;
847     const ret = findOpNamedInBlock(entry, "func.return") orelse return error.TestExpectedReturn;
848     const folded_const = constantDefiningOp(ret.getOperand(0).?) orelse return error.TestExpectedConstant;
849     const folded_payload = (dialect_mod.AccyDialect.ConstantOp{ .op = folded_const }).getPayload() orelse return error.TestExpectedPayload;
850     try testing.expectEqualSlices(u8, std.mem.sliceAsBytes(expected[0..]), folded_payload);
851 }
852 
853 test "constant folding folds i32 neg constants" {
854     const allocator = testing.allocator;
855 
856     var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
857     defer builder.deinit();
858     const i32_3 = try builder.tensor(.i32, &.{3});
859     const input = [_]i32{ 7, -11, 0 };
860     const expected = [_]i32{ -7, 11, 0 };
861 
862     var fb = try builder.beginFunction("fold_neg_constants", &.{}, &.{i32_3});
863     const c = try fb.constant(i32_3, std.mem.sliceAsBytes(input[0..]));
864     const n = try fb.neg(c);
865     try fb.return_(&.{n});
866     try fb.finish();
867     const module = try builder.finish();
868     defer module.deinit();
869 
870     const choir_mod = module.choir_module;
871     const ctx = module.context();
872     var pm = passes.PassManager.init(allocator);
873     defer pm.deinit();
874     try pm.addPass(shape_analysis.shapeLayoutPropagationPass());
875     try pm.addPass(constantFoldingPass());
876 
877     try testing.expectEqual(passes.PassResult.success, pm.run(choir_mod, ctx));
878     try testing.expectEqual(@as(u64, 2), pm.stats.pass_runs);
879     try testing.expectEqual(@as(u64, 1), pm.stats.passes_modified);
880     try testing.expectEqual(@as(usize, 0), ir.inspection.countOperationsNamed(choir_mod, dialect_mod.AccyDialect.NegOp.operation_name));
881     try testing.expectEqual(@as(usize, 2), ir.inspection.countOperationsNamed(choir_mod, dialect_mod.AccyDialect.ConstantOp.operation_name));
882 
883     const module_body = choir_mod.getRegion(0).?.getEntryBlock().?;
884     const func = ir.inspection.functionByNameInBlock(module_body, "fold_neg_constants") orelse return error.TestExpectedFunc;
885     const entry = func.getRegion(0).?.getEntryBlock().?;
886     const ret = findOpNamedInBlock(entry, "func.return") orelse return error.TestExpectedReturn;
887     const folded_const = constantDefiningOp(ret.getOperand(0).?) orelse return error.TestExpectedConstant;
888     const folded_payload = (dialect_mod.AccyDialect.ConstantOp{ .op = folded_const }).getPayload() orelse return error.TestExpectedPayload;
889     try testing.expectEqualSlices(u8, std.mem.sliceAsBytes(expected[0..]), folded_payload);
890 }
891 
892 test "constant folding folds f32 floor, round, and trunc constants" {
893     const allocator = testing.allocator;
894 
895     var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
896     defer builder.deinit();
897     const f32_4 = try builder.tensor(.f32, &.{4});
898     const input = [_]f32{ -1.75, -0.25, 0.25, 1.75 };
899     const expected_floor = [_]f32{ -2.0, -1.0, 0.0, 1.0 };
900     const expected_round = [_]f32{ -2.0, -0.0, 0.0, 2.0 };
901     const expected_trunc = [_]f32{ -1.0, -0.0, 0.0, 1.0 };
902 
903     var fb = try builder.beginFunction("fold_rounding_constants", &.{}, &.{ f32_4, f32_4, f32_4 });
904     const c = try fb.constant(f32_4, std.mem.sliceAsBytes(input[0..]));
905     const floored = try fb.floor(c);
906     const rounded = try fb.round(c);
907     const truncated = try fb.trunc(c);
908     try fb.return_(&.{ floored, rounded, truncated });
909     try fb.finish();
910     const module = try builder.finish();
911     defer module.deinit();
912 
913     const choir_mod = module.choir_module;
914     const ctx = module.context();
915     var pm = passes.PassManager.init(allocator);
916     defer pm.deinit();
917     try pm.addPass(shape_analysis.shapeLayoutPropagationPass());
918     try pm.addPass(constantFoldingPass());
919 
920     try testing.expectEqual(passes.PassResult.success, pm.run(choir_mod, ctx));
921     try testing.expectEqual(@as(u64, 2), pm.stats.pass_runs);
922     try testing.expectEqual(@as(u64, 1), pm.stats.passes_modified);
923     try testing.expectEqual(@as(usize, 0), ir.inspection.countOperationsNamed(choir_mod, dialect_mod.AccyDialect.FloorOp.operation_name));
924     try testing.expectEqual(@as(usize, 0), ir.inspection.countOperationsNamed(choir_mod, dialect_mod.AccyDialect.RoundOp.operation_name));
925     try testing.expectEqual(@as(usize, 0), ir.inspection.countOperationsNamed(choir_mod, dialect_mod.AccyDialect.TruncOp.operation_name));
926     try testing.expectEqual(@as(usize, 4), ir.inspection.countOperationsNamed(choir_mod, dialect_mod.AccyDialect.ConstantOp.operation_name));
927 
928     const module_body = choir_mod.getRegion(0).?.getEntryBlock().?;
929     const func = ir.inspection.functionByNameInBlock(module_body, "fold_rounding_constants") orelse return error.TestExpectedFunc;
930     const entry = func.getRegion(0).?.getEntryBlock().?;
931     const ret = findOpNamedInBlock(entry, "func.return") orelse return error.TestExpectedReturn;
932     const floor_const = constantDefiningOp(ret.getOperand(0).?) orelse return error.TestExpectedConstant;
933     const round_const = constantDefiningOp(ret.getOperand(1).?) orelse return error.TestExpectedConstant;
934     const trunc_const = constantDefiningOp(ret.getOperand(2).?) orelse return error.TestExpectedConstant;
935     const floor_payload = (dialect_mod.AccyDialect.ConstantOp{ .op = floor_const }).getPayload() orelse return error.TestExpectedPayload;
936     const round_payload = (dialect_mod.AccyDialect.ConstantOp{ .op = round_const }).getPayload() orelse return error.TestExpectedPayload;
937     const trunc_payload = (dialect_mod.AccyDialect.ConstantOp{ .op = trunc_const }).getPayload() orelse return error.TestExpectedPayload;
938     try testing.expectEqualSlices(u8, std.mem.sliceAsBytes(expected_floor[0..]), floor_payload);
939     try testing.expectEqualSlices(u8, std.mem.sliceAsBytes(expected_round[0..]), round_payload);
940     try testing.expectEqualSlices(u8, std.mem.sliceAsBytes(expected_trunc[0..]), trunc_payload);
941 }
942 
943 test "constant folding skips overflowing i32 add constants" {
944     const allocator = testing.allocator;
945 
946     var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
947     defer builder.deinit();
948     const i32_1 = try builder.tensor(.i32, &.{1});
949     const lhs = [_]i32{std.math.maxInt(i32)};
950     const rhs = [_]i32{1};
951 
952     var fb = try builder.beginFunction("fold_skip_overflow", &.{}, &.{i32_1});
953     const l = try fb.constant(i32_1, std.mem.sliceAsBytes(lhs[0..]));
954     const r = try fb.constant(i32_1, std.mem.sliceAsBytes(rhs[0..]));
955     const sum = try fb.add(l, r);
956     try fb.return_(&.{sum});
957     try fb.finish();
958     const module = try builder.finish();
959     defer module.deinit();
960 
961     const choir_mod = module.choir_module;
962     const ctx = module.context();
963     var pm = passes.PassManager.init(allocator);
964     defer pm.deinit();
965     try pm.addPass(shape_analysis.shapeLayoutPropagationPass());
966     try pm.addPass(constantFoldingPass());
967 
968     try testing.expectEqual(passes.PassResult.success, pm.run(choir_mod, ctx));
969     try testing.expectEqual(@as(u64, 2), pm.stats.pass_runs);
970     try testing.expectEqual(@as(u64, 0), pm.stats.passes_modified);
971     try testing.expectEqual(@as(usize, 1), ir.inspection.countOperationsNamed(choir_mod, dialect_mod.AccyDialect.AddOp.operation_name));
972     try testing.expectEqual(@as(usize, 2), ir.inspection.countOperationsNamed(choir_mod, dialect_mod.AccyDialect.ConstantOp.operation_name));
973 }
974 
975 test "constant folding pass preserves analyses when no constants fold" {
976     const allocator = testing.allocator;
977 
978     var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
979     defer builder.deinit();
980     const f32_4 = try builder.tensor(.f32, &.{4});
981     var fb = try builder.beginFunction("fold_noop_add4", &.{ f32_4, f32_4 }, &.{f32_4});
982     const sum = try fb.add(fb.parameter(0), fb.parameter(1));
983     try fb.return_(&.{sum});
984     try fb.finish();
985     const module = try builder.finish();
986     defer module.deinit();
987 
988     const choir_mod = module.choir_module;
989     const ctx = module.context();
990     var pm = passes.PassManager.init(allocator);
991     defer pm.deinit();
992     try pm.addPass(constantFoldingPass());
993 
994     try testing.expectEqual(passes.PassResult.success, pm.run(choir_mod, ctx));
995     try testing.expectEqual(@as(u64, 1), pm.stats.pass_runs);
996     try testing.expectEqual(@as(u64, 0), pm.stats.passes_modified);
997     try testing.expectEqual(@as(u64, 1), pm.stats.analysis_misses);
998 }