lib/choir/src/profiling/versus/core/compile.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const alloc_arena = @import("alloc_arena");
3 const dialects = @import("../../../dialects/root.zig");
4 const ir = @import("../../../core/root.zig");
5
6 const ArithDialect = dialects.ArithDialect;
7 const BuiltinDialect = dialects.BuiltinDialect;
8 const FuncDialect = dialects.FuncDialect;
9 const ScfDialect = dialects.ScfDialect;
10
11 pub const Shape = enum {
12 arith_chain,
13 fold_cse_grid,
14 control_loop,
15 symbol_call_forest,
16
17 pub fn name(self: Shape) []const u8 {
18 return switch (self) {
19 .arith_chain => "arith_chain",
20 .fold_cse_grid => "fold_cse_grid",
21 .control_loop => "control_loop",
22 .symbol_call_forest => "symbol_call_forest",
23 };
24 }
25 };
26
27 pub const Workload = struct {
28 name: []const u8,
29 shape: Shape,
30 functions: u32,
31 ops_per_function: u32,
32
33 pub fn sourceOps(self: Workload) u64 {
34 return @as(u64, self.functions) * @as(u64, self.ops_per_function);
35 }
36 };
37
38 pub const matrix = [_]Workload{
39 .{ .name = "compile_arith_1x32", .shape = .arith_chain, .functions = 1, .ops_per_function = 32 },
40 .{ .name = "compile_arith_8x128", .shape = .arith_chain, .functions = 8, .ops_per_function = 128 },
41 .{ .name = "compile_foldcse_4x64", .shape = .fold_cse_grid, .functions = 4, .ops_per_function = 64 },
42 .{ .name = "compile_control_4x32", .shape = .control_loop, .functions = 4, .ops_per_function = 32 },
43 .{ .name = "compile_calls_8x32", .shape = .symbol_call_forest, .functions = 8, .ops_per_function = 32 },
44 };
45
46 pub fn byName(name: []const u8) ?Workload {
47 for (matrix) |workload| {
48 if (std.mem.eql(u8, workload.name, name)) return workload;
49 }
50 return null;
51 }
52
53 pub fn build(ctx: *ir.Context, workload: Workload) !BuiltinDialect.ModuleOp {
54 const loc = ir.Location.getUnknown();
55 const module = try BuiltinDialect.ModuleOp.create(ctx, loc);
56 const i32_type = try ArithDialect.getScalarType(ctx, .i32);
57 const module_block = module.getBodyBlock();
58
59 for (0..workload.functions) |function_index_usize| {
60 const function_index: u32 = @intCast(function_index_usize);
61 const name = try std.fmt.allocPrint(ir.context.transientAllocator(ctx), "compile_f_{d}", .{function_index});
62 const func = try FuncDialect.FuncOp.create(
63 ctx,
64 loc,
65 name,
66 &.{ i32_type, i32_type },
67 &.{i32_type},
68 );
69 try module_block.addOperation(func.op);
70 try populateFunction(ctx, loc, func, i32_type, workload, function_index);
71 }
72
73 return module;
74 }
75
76 fn populateFunction(
77 ctx: *ir.Context,
78 loc: ir.Location,
79 func: FuncDialect.FuncOp,
80 i32_type: ir.Type,
81 workload: Workload,
82 function_index: u32,
83 ) !void {
84 const entry = func.getEntryBlock();
85 var current = func.getArgument(0);
86 const rhs = func.getArgument(1);
87
88 for (0..workload.ops_per_function) |op_index_usize| {
89 const op_index: u32 = @intCast(op_index_usize);
90 current = switch (workload.shape) {
91 .arith_chain => blk: {
92 const constant = try appendConstant(ctx, loc, entry, i32_type, @intCast((op_index % 17) + 1));
93 break :blk try appendArithChainOp(ctx, loc, entry, current, rhs, constant.getResult(), op_index);
94 },
95 .fold_cse_grid => blk: {
96 const constant = try appendConstant(ctx, loc, entry, i32_type, @intCast((op_index % 17) + 1));
97 break :blk try appendFoldCseGridOp(ctx, loc, entry, current, rhs, constant.getResult(), i32_type, function_index, op_index);
98 },
99 .control_loop => try appendControlLoopOp(ctx, loc, entry, current, rhs, i32_type, op_index),
100 .symbol_call_forest => blk: {
101 const constant = try appendConstant(ctx, loc, entry, i32_type, @intCast((op_index % 23) + 1));
102 break :blk try appendSymbolCallForestOp(ctx, loc, entry, current, rhs, constant.getResult(), i32_type, function_index, op_index);
103 },
104 };
105 }
106
107 const ret = try FuncDialect.ReturnOp.create(ctx, loc, &.{current});
108 try entry.addOperation(ret.op);
109 }
110
111 fn appendConstant(
112 ctx: *ir.Context,
113 loc: ir.Location,
114 entry: *ir.Block,
115 i32_type: ir.Type,
116 value: i64,
117 ) !ArithDialect.ConstantOp {
118 const constant = try ArithDialect.ConstantOp.createInt(ctx, loc, i32_type, value);
119 try entry.addOperation(constant.op);
120 return constant;
121 }
122
123 fn appendArithChainOp(
124 ctx: *ir.Context,
125 loc: ir.Location,
126 entry: *ir.Block,
127 current: *ir.Value,
128 rhs: *ir.Value,
129 constant: *ir.Value,
130 op_index: u32,
131 ) !*ir.Value {
132 switch (op_index % 4) {
133 0 => {
134 var op = try ArithDialect.AddOp.create(ctx, loc, current, constant);
135 try entry.addOperation(op.op);
136 return op.getResult();
137 },
138 1 => {
139 var op = try ArithDialect.MulOp.create(ctx, loc, current, rhs);
140 try entry.addOperation(op.op);
141 return op.getResult();
142 },
143 2 => {
144 var op = try ArithDialect.SubOp.create(ctx, loc, current, constant);
145 try entry.addOperation(op.op);
146 return op.getResult();
147 },
148 else => {
149 var op = try ArithDialect.AddOp.create(ctx, loc, current, rhs);
150 try entry.addOperation(op.op);
151 return op.getResult();
152 },
153 }
154 }
155
156 fn appendFoldCseGridOp(
157 ctx: *ir.Context,
158 loc: ir.Location,
159 entry: *ir.Block,
160 current: *ir.Value,
161 rhs: *ir.Value,
162 constant: *ir.Value,
163 i32_type: ir.Type,
164 function_index: u32,
165 op_index: u32,
166 ) !*ir.Value {
167 var lhs_dead = try ArithDialect.ConstantOp.createInt(ctx, loc, i32_type, @intCast(function_index % 11));
168 var rhs_dead = try ArithDialect.ConstantOp.createInt(ctx, loc, i32_type, @intCast(op_index % 13));
169 const folded_dead = try ArithDialect.AddOp.create(ctx, loc, lhs_dead.getResult(), rhs_dead.getResult());
170 try entry.addOperation(lhs_dead.op);
171 try entry.addOperation(rhs_dead.op);
172 try entry.addOperation(folded_dead.op);
173
174 var next = try ArithDialect.AddOp.create(ctx, loc, current, constant);
175 try entry.addOperation(next.op);
176 const duplicate = try ArithDialect.AddOp.create(ctx, loc, current, constant);
177 try entry.addOperation(duplicate.op);
178 var mixed = try ArithDialect.MulOp.create(ctx, loc, next.getResult(), rhs);
179 try entry.addOperation(mixed.op);
180 return mixed.getResult();
181 }
182
183 fn appendControlLoopOp(
184 ctx: *ir.Context,
185 loc: ir.Location,
186 entry: *ir.Block,
187 current: *ir.Value,
188 rhs: *ir.Value,
189 i32_type: ir.Type,
190 op_index: u32,
191 ) !*ir.Value {
192 var lower = try ArithDialect.ConstantOp.createInt(ctx, loc, i32_type, 0);
193 var upper = try ArithDialect.ConstantOp.createInt(ctx, loc, i32_type, @intCast((op_index % 31) + 2));
194 var step = try ArithDialect.ConstantOp.createInt(ctx, loc, i32_type, 1);
195 try entry.addOperation(lower.op);
196 try entry.addOperation(upper.op);
197 try entry.addOperation(step.op);
198
199 var for_op = try ScfDialect.ForOp.create(ctx, loc, lower.getResult(), upper.getResult(), step.getResult(), &.{current}, &.{i32_type});
200 try entry.addOperation(for_op.op);
201 const body = for_op.getBodyBlock();
202 const iter_args = for_op.getIterArgs();
203 const iter_arg = iter_args[0];
204 var invariant_const = try ArithDialect.ConstantOp.createInt(ctx, loc, i32_type, @intCast((op_index % 17) + 1));
205 try body.addOperation(invariant_const.op);
206 var invariant = try ArithDialect.AddOp.create(ctx, loc, rhs, invariant_const.getResult());
207 try body.addOperation(invariant.op);
208 var carried = try ArithDialect.AddOp.create(ctx, loc, iter_arg, invariant.getResult());
209 try body.addOperation(carried.op);
210 var stepped = try ArithDialect.AddOp.create(ctx, loc, carried.getResult(), for_op.getInductionVar());
211 try body.addOperation(stepped.op);
212 const yield = try ScfDialect.YieldOp.create(ctx, loc, &.{stepped.getResult()});
213 try body.addOperation(yield.op);
214 return for_op.getResult(0).?;
215 }
216
217 fn appendSymbolCallForestOp(
218 ctx: *ir.Context,
219 loc: ir.Location,
220 entry: *ir.Block,
221 current: *ir.Value,
222 rhs: *ir.Value,
223 constant: *ir.Value,
224 i32_type: ir.Type,
225 function_index: u32,
226 op_index: u32,
227 ) !*ir.Value {
228 if (function_index == 0) {
229 return appendArithChainOp(ctx, loc, entry, current, rhs, constant, op_index);
230 }
231
232 var lhs_dead = try ArithDialect.ConstantOp.createInt(ctx, loc, i32_type, @intCast((function_index % 17) + 3));
233 var rhs_dead = try ArithDialect.ConstantOp.createInt(ctx, loc, i32_type, @intCast((op_index % 19) + 5));
234 const folded_dead = try ArithDialect.AddOp.create(ctx, loc, lhs_dead.getResult(), rhs_dead.getResult());
235 try entry.addOperation(lhs_dead.op);
236 try entry.addOperation(rhs_dead.op);
237 try entry.addOperation(folded_dead.op);
238
239 const callee_name = try symbolCallCalleeName(ctx, function_index, op_index);
240 var call = try FuncDialect.CallOp.create(ctx, loc, callee_name, &.{ current, rhs }, &.{i32_type});
241 try entry.addOperation(call.op);
242
243 var combined = try ArithDialect.AddOp.create(ctx, loc, call.getResult(0).?, constant);
244 try entry.addOperation(combined.op);
245 return combined.getResult();
246 }
247
248 fn symbolCallCalleeName(ctx: *ir.Context, function_index: u32, op_index: u32) ![]const u8 {
249 const window = @min(function_index, 4);
250 const callee_index = function_index - 1 - (op_index % window);
251 return try std.fmt.allocPrint(ir.context.transientAllocator(ctx), "compile_f_{d}", .{callee_index});
252 }
253
254 test "compile matrix names resolve" {
255 for (matrix) |workload| {
256 try std.testing.expectEqual(workload.shape, byName(workload.name).?.shape);
257 try std.testing.expect(workload.sourceOps() > 0);
258 }
259 try std.testing.expect(byName("missing") == null);
260 }
261
262 test "compile matrix workloads build modules" {
263 var arena = alloc_arena.Arena.init(std.testing.allocator);
264 defer arena.deinit();
265
266 for (matrix) |workload| {
267 var ctx = try ir.Context.init(arena.allocator(), ir.Context.Limits.testing);
268 defer ctx.deinit(arena.allocator());
269 try dialects.registerAllDialects(&ctx);
270 const module = try build(&ctx, workload);
271 try ir.verifyOperation(module.op, ir.verify.default_options);
272 try std.testing.expect(ir.inspection.countOperationsNamed(module.op, FuncDialect.FuncOp.operation_name) > 0);
273 }
274 }