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 }