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 }