lib/choir/src/profiling/versus/kernels.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const alloc_arena = @import("alloc_arena");
3 const choir = @import("choir");
4
5 const workload_mod = choir.versus.workload;
6
7 const ArithDialect = choir.dialects.ArithDialect;
8 const BuiltinDialect = choir.dialects.BuiltinDialect;
9 const FuncDialect = choir.dialects.FuncDialect;
10 const MemrefDialect = choir.dialects.MemrefDialect;
11 const ScfDialect = choir.dialects.ScfDialect;
12
13 const loc = choir.ir.Location.getUnknown();
14
15 const Types = struct {
16 index: choir.Type,
17 i32_: choir.Type,
18 i64_: choir.Type,
19 f32_: choir.Type,
20 f64_: choir.Type,
21 memref_f32: choir.Type,
22 memref_f64: choir.Type,
23 memref_i64: choir.Type,
24
25 fn init(ctx: *choir.Context) !Types {
26 const f32_type = try ArithDialect.getScalarType(ctx, .f32);
27 const f64_type = try ArithDialect.getScalarType(ctx, .f64);
28 const i32_type = try ArithDialect.getScalarType(ctx, .i32);
29 const i64_type = try ArithDialect.getScalarType(ctx, .i64);
30 return .{
31 .index = try ArithDialect.getIndexType(ctx),
32 .i32_ = i32_type,
33 .i64_ = i64_type,
34 .f32_ = f32_type,
35 .f64_ = f64_type,
36 .memref_f32 = try MemrefDialect.getMemrefTypeDynamic(ctx, f32_type, .host),
37 .memref_f64 = try MemrefDialect.getMemrefTypeDynamic(ctx, f64_type, .host),
38 .memref_i64 = try MemrefDialect.getMemrefTypeDynamic(ctx, i64_type, .host),
39 };
40 }
41 };
42
43 fn appended(block: *choir.Block, op: anytype) !*choir.Value {
44 try block.addOperation(op.op);
45 var mutable = op;
46 return mutable.getResult();
47 }
48
49 fn constIndex(ctx: *choir.Context, block: *choir.Block, types: Types, value: i64) !*choir.Value {
50 return appended(block, try ArithDialect.ConstantOp.createInt(ctx, loc, types.index, value));
51 }
52
53 fn castToIndex(ctx: *choir.Context, block: *choir.Block, types: Types, value: *choir.Value) !*choir.Value {
54 return appended(block, try ArithDialect.CastOp.create(ctx, loc, value, types.index));
55 }
56
57 pub fn build(ctx: *choir.Context, kind: workload_mod.Kind) !BuiltinDialect.ModuleOp {
58 const types = try Types.init(ctx);
59 const module = try BuiltinDialect.ModuleOp.create(ctx, loc);
60 const module_block = module.getBodyBlock();
61
62 switch (kind) {
63 .saxpy => try buildSaxpy(ctx, module_block, types),
64 .dot => try buildDot(ctx, module_block, types),
65 .sum => try buildSum(ctx, module_block, types),
66 .matmul => try buildMatmul(ctx, module_block, types),
67 .polybench_gemm => try buildPolybenchGemm(ctx, module_block, types),
68 .stencil3 => try buildStencil3(ctx, module_block, types),
69 .clampsum => try buildClampsum(ctx, module_block, types),
70 }
71 return module;
72 }
73
74 fn buildSaxpy(ctx: *choir.Context, module_block: *choir.Block, types: Types) !void {
75 var func = try FuncDialect.FuncOp.create(ctx, loc, workload_mod.Kind.saxpy.symbol(), &.{
76 types.i64_, types.memref_f32, types.memref_f32, types.memref_f32, types.memref_f32,
77 }, &.{});
78 try module_block.addOperation(func.op);
79 const entry = func.getEntryBlock();
80
81 const zero = try constIndex(ctx, entry, types, 0);
82 const scale = try appended(entry, try MemrefDialect.LoadOp.create(ctx, loc, func.getArgument(1), zero, types.f32_));
83 const bound = try castToIndex(ctx, entry, types, func.getArgument(0));
84 const step = try constIndex(ctx, entry, types, 1);
85
86 var for_op = try ScfDialect.ForOp.create(ctx, loc, zero, bound, step, &.{}, &.{});
87 try entry.addOperation(for_op.op);
88 const body = for_op.getBodyBlock();
89 const iv = body.arguments.items[0];
90
91 const x = try appended(body, try MemrefDialect.LoadOp.create(ctx, loc, func.getArgument(2), iv, types.f32_));
92 const y = try appended(body, try MemrefDialect.LoadOp.create(ctx, loc, func.getArgument(3), iv, types.f32_));
93 const scaled = try appended(body, try ArithDialect.MulOp.create(ctx, loc, scale, x));
94 const result = try appended(body, try ArithDialect.AddOp.create(ctx, loc, scaled, y));
95 const store = try MemrefDialect.StoreOp.create(ctx, loc, result, func.getArgument(4), iv);
96 try body.addOperation(store.op);
97 const yield_op = try ScfDialect.YieldOp.create(ctx, loc, &.{});
98 try body.addOperation(yield_op.op);
99
100 const ret = try FuncDialect.ReturnOp.create(ctx, loc, &.{});
101 try entry.addOperation(ret.op);
102 }
103
104 fn buildDot(ctx: *choir.Context, module_block: *choir.Block, types: Types) !void {
105 var func = try FuncDialect.FuncOp.create(ctx, loc, workload_mod.Kind.dot.symbol(), &.{
106 types.i64_, types.memref_f64, types.memref_f64, types.memref_f64,
107 }, &.{});
108 try module_block.addOperation(func.op);
109 const entry = func.getEntryBlock();
110
111 const zero = try constIndex(ctx, entry, types, 0);
112 const bound = try castToIndex(ctx, entry, types, func.getArgument(0));
113 const step = try constIndex(ctx, entry, types, 1);
114 const acc_init = try appended(entry, try ArithDialect.ConstantOp.createFloat(ctx, loc, types.f64_, 0.0));
115
116 var for_op = try ScfDialect.ForOp.create(ctx, loc, zero, bound, step, &.{acc_init}, &.{types.f64_});
117 try entry.addOperation(for_op.op);
118 const body = for_op.getBodyBlock();
119 const iv = body.arguments.items[0];
120 const acc = body.arguments.items[1];
121
122 const x = try appended(body, try MemrefDialect.LoadOp.create(ctx, loc, func.getArgument(1), iv, types.f64_));
123 const y = try appended(body, try MemrefDialect.LoadOp.create(ctx, loc, func.getArgument(2), iv, types.f64_));
124 const product = try appended(body, try ArithDialect.MulOp.create(ctx, loc, x, y));
125 const next = try appended(body, try ArithDialect.AddOp.create(ctx, loc, acc, product));
126 const yield_op = try ScfDialect.YieldOp.create(ctx, loc, &.{next});
127 try body.addOperation(yield_op.op);
128
129 const total = for_op.getResult(0) orelse return error.MissingLoopResult;
130 const store = try MemrefDialect.StoreOp.create(ctx, loc, total, func.getArgument(3), zero);
131 try entry.addOperation(store.op);
132 const ret = try FuncDialect.ReturnOp.create(ctx, loc, &.{});
133 try entry.addOperation(ret.op);
134 }
135
136 fn buildSum(ctx: *choir.Context, module_block: *choir.Block, types: Types) !void {
137 var func = try FuncDialect.FuncOp.create(ctx, loc, workload_mod.Kind.sum.symbol(), &.{
138 types.i64_, types.memref_i64, types.memref_i64,
139 }, &.{});
140 try module_block.addOperation(func.op);
141 const entry = func.getEntryBlock();
142
143 const zero = try constIndex(ctx, entry, types, 0);
144 const bound = try castToIndex(ctx, entry, types, func.getArgument(0));
145 const step = try constIndex(ctx, entry, types, 1);
146 const acc_init = try appended(entry, try ArithDialect.ConstantOp.createInt(ctx, loc, types.i64_, 0));
147
148 var for_op = try ScfDialect.ForOp.create(ctx, loc, zero, bound, step, &.{acc_init}, &.{types.i64_});
149 try entry.addOperation(for_op.op);
150 const body = for_op.getBodyBlock();
151 const iv = body.arguments.items[0];
152 const acc = body.arguments.items[1];
153
154 const x = try appended(body, try MemrefDialect.LoadOp.create(ctx, loc, func.getArgument(1), iv, types.i64_));
155 const next = try appended(body, try ArithDialect.AddOp.create(ctx, loc, acc, x));
156 const yield_op = try ScfDialect.YieldOp.create(ctx, loc, &.{next});
157 try body.addOperation(yield_op.op);
158
159 const total = for_op.getResult(0) orelse return error.MissingLoopResult;
160 const store = try MemrefDialect.StoreOp.create(ctx, loc, total, func.getArgument(2), zero);
161 try entry.addOperation(store.op);
162 const ret = try FuncDialect.ReturnOp.create(ctx, loc, &.{});
163 try entry.addOperation(ret.op);
164 }
165
166 fn buildMatmul(ctx: *choir.Context, module_block: *choir.Block, types: Types) !void {
167 var func = try FuncDialect.FuncOp.create(ctx, loc, workload_mod.Kind.matmul.symbol(), &.{
168 types.i64_, types.memref_f64, types.memref_f64, types.memref_f64,
169 }, &.{});
170 try module_block.addOperation(func.op);
171 const entry = func.getEntryBlock();
172
173 const zero = try constIndex(ctx, entry, types, 0);
174 const bound = try castToIndex(ctx, entry, types, func.getArgument(0));
175 const step = try constIndex(ctx, entry, types, 1);
176
177 var i_loop = try ScfDialect.ForOp.create(ctx, loc, zero, bound, step, &.{}, &.{});
178 try entry.addOperation(i_loop.op);
179 const i_body = i_loop.getBodyBlock();
180 const i = i_body.arguments.items[0];
181
182 var j_loop = try ScfDialect.ForOp.create(ctx, loc, zero, bound, step, &.{}, &.{});
183 try i_body.addOperation(j_loop.op);
184 const j_body = j_loop.getBodyBlock();
185 const j = j_body.arguments.items[0];
186
187 const acc_init = try appended(j_body, try ArithDialect.ConstantOp.createFloat(ctx, loc, types.f64_, 0.0));
188 var k_loop = try ScfDialect.ForOp.create(ctx, loc, zero, bound, step, &.{acc_init}, &.{types.f64_});
189 try j_body.addOperation(k_loop.op);
190 const k_body = k_loop.getBodyBlock();
191 const k = k_body.arguments.items[0];
192 const acc = k_body.arguments.items[1];
193
194 const i_row = try appended(k_body, try ArithDialect.MulOp.create(ctx, loc, i, bound));
195 const a_index = try appended(k_body, try ArithDialect.AddOp.create(ctx, loc, i_row, k));
196 const k_row = try appended(k_body, try ArithDialect.MulOp.create(ctx, loc, k, bound));
197 const b_index = try appended(k_body, try ArithDialect.AddOp.create(ctx, loc, k_row, j));
198 const a = try appended(k_body, try MemrefDialect.LoadOp.create(ctx, loc, func.getArgument(1), a_index, types.f64_));
199 const b = try appended(k_body, try MemrefDialect.LoadOp.create(ctx, loc, func.getArgument(2), b_index, types.f64_));
200 const product = try appended(k_body, try ArithDialect.MulOp.create(ctx, loc, a, b));
201 const next = try appended(k_body, try ArithDialect.AddOp.create(ctx, loc, acc, product));
202 const k_yield = try ScfDialect.YieldOp.create(ctx, loc, &.{next});
203 try k_body.addOperation(k_yield.op);
204
205 const total = k_loop.getResult(0) orelse return error.MissingLoopResult;
206 const c_row = try appended(j_body, try ArithDialect.MulOp.create(ctx, loc, i, bound));
207 const c_index = try appended(j_body, try ArithDialect.AddOp.create(ctx, loc, c_row, j));
208 const store = try MemrefDialect.StoreOp.create(ctx, loc, total, func.getArgument(3), c_index);
209 try j_body.addOperation(store.op);
210 const j_yield = try ScfDialect.YieldOp.create(ctx, loc, &.{});
211 try j_body.addOperation(j_yield.op);
212 const i_yield = try ScfDialect.YieldOp.create(ctx, loc, &.{});
213 try i_body.addOperation(i_yield.op);
214
215 const ret = try FuncDialect.ReturnOp.create(ctx, loc, &.{});
216 try entry.addOperation(ret.op);
217 }
218
219 fn buildPolybenchGemm(ctx: *choir.Context, module_block: *choir.Block, types: Types) !void {
220 var func = try FuncDialect.FuncOp.create(ctx, loc, workload_mod.Kind.polybench_gemm.symbol(), &.{
221 types.i32_, types.i32_, types.i32_, types.f64_, types.f64_, types.memref_f64, types.memref_f64, types.memref_f64,
222 }, &.{});
223 try module_block.addOperation(func.op);
224 const entry = func.getEntryBlock();
225
226 const zero = try constIndex(ctx, entry, types, 0);
227 const ni = try castToIndex(ctx, entry, types, func.getArgument(0));
228 const nj = try castToIndex(ctx, entry, types, func.getArgument(1));
229 const nk = try castToIndex(ctx, entry, types, func.getArgument(2));
230 const step = try constIndex(ctx, entry, types, 1);
231
232 var i_loop = try ScfDialect.ForOp.create(ctx, loc, zero, ni, step, &.{}, &.{});
233 try entry.addOperation(i_loop.op);
234 const i_body = i_loop.getBodyBlock();
235 const i = i_body.arguments.items[0];
236
237 var scale_j_loop = try ScfDialect.ForOp.create(ctx, loc, zero, nj, step, &.{}, &.{});
238 try i_body.addOperation(scale_j_loop.op);
239 const scale_j_body = scale_j_loop.getBodyBlock();
240 const scale_j = scale_j_body.arguments.items[0];
241
242 const scale_c_row = try appended(scale_j_body, try ArithDialect.MulOp.create(ctx, loc, i, nj));
243 const scale_c_index = try appended(scale_j_body, try ArithDialect.AddOp.create(ctx, loc, scale_c_row, scale_j));
244 const scale_c = try appended(scale_j_body, try MemrefDialect.LoadOp.create(ctx, loc, func.getArgument(5), scale_c_index, types.f64_));
245 const scaled_c = try appended(scale_j_body, try ArithDialect.MulOp.create(ctx, loc, scale_c, func.getArgument(4)));
246 const scale_store = try MemrefDialect.StoreOp.create(ctx, loc, scaled_c, func.getArgument(5), scale_c_index);
247 try scale_j_body.addOperation(scale_store.op);
248 const scale_j_yield = try ScfDialect.YieldOp.create(ctx, loc, &.{});
249 try scale_j_body.addOperation(scale_j_yield.op);
250
251 var k_loop = try ScfDialect.ForOp.create(ctx, loc, zero, nk, step, &.{}, &.{});
252 try i_body.addOperation(k_loop.op);
253 const k_body = k_loop.getBodyBlock();
254 const k = k_body.arguments.items[0];
255
256 var update_j_loop = try ScfDialect.ForOp.create(ctx, loc, zero, nj, step, &.{}, &.{});
257 try k_body.addOperation(update_j_loop.op);
258 const update_j_body = update_j_loop.getBodyBlock();
259 const j = update_j_body.arguments.items[0];
260
261 const c_row = try appended(update_j_body, try ArithDialect.MulOp.create(ctx, loc, i, nj));
262 const c_index = try appended(update_j_body, try ArithDialect.AddOp.create(ctx, loc, c_row, j));
263 const a_row = try appended(update_j_body, try ArithDialect.MulOp.create(ctx, loc, i, nk));
264 const a_index = try appended(update_j_body, try ArithDialect.AddOp.create(ctx, loc, a_row, k));
265 const b_row = try appended(update_j_body, try ArithDialect.MulOp.create(ctx, loc, k, nj));
266 const b_index = try appended(update_j_body, try ArithDialect.AddOp.create(ctx, loc, b_row, j));
267 const c_value = try appended(update_j_body, try MemrefDialect.LoadOp.create(ctx, loc, func.getArgument(5), c_index, types.f64_));
268 const a = try appended(update_j_body, try MemrefDialect.LoadOp.create(ctx, loc, func.getArgument(6), a_index, types.f64_));
269 const b = try appended(update_j_body, try MemrefDialect.LoadOp.create(ctx, loc, func.getArgument(7), b_index, types.f64_));
270 const scaled_a = try appended(update_j_body, try ArithDialect.MulOp.create(ctx, loc, func.getArgument(3), a));
271 const product = try appended(update_j_body, try ArithDialect.MulOp.create(ctx, loc, scaled_a, b));
272 const next = try appended(update_j_body, try ArithDialect.AddOp.create(ctx, loc, c_value, product));
273 const update_store = try MemrefDialect.StoreOp.create(ctx, loc, next, func.getArgument(5), c_index);
274 try update_j_body.addOperation(update_store.op);
275 const update_j_yield = try ScfDialect.YieldOp.create(ctx, loc, &.{});
276 try update_j_body.addOperation(update_j_yield.op);
277 const k_yield = try ScfDialect.YieldOp.create(ctx, loc, &.{});
278 try k_body.addOperation(k_yield.op);
279 const i_yield = try ScfDialect.YieldOp.create(ctx, loc, &.{});
280 try i_body.addOperation(i_yield.op);
281
282 const ret = try FuncDialect.ReturnOp.create(ctx, loc, &.{});
283 try entry.addOperation(ret.op);
284 }
285
286 fn buildStencil3(ctx: *choir.Context, module_block: *choir.Block, types: Types) !void {
287 var func = try FuncDialect.FuncOp.create(ctx, loc, workload_mod.Kind.stencil3.symbol(), &.{
288 types.i64_, types.memref_f64, types.memref_f64,
289 }, &.{});
290 try module_block.addOperation(func.op);
291 const entry = func.getEntryBlock();
292
293 const one = try constIndex(ctx, entry, types, 1);
294 const bound = try castToIndex(ctx, entry, types, func.getArgument(0));
295 const upper = try appended(entry, try ArithDialect.SubOp.create(ctx, loc, bound, one));
296 const quarter = try appended(entry, try ArithDialect.ConstantOp.createFloat(ctx, loc, types.f64_, 0.25));
297 const half = try appended(entry, try ArithDialect.ConstantOp.createFloat(ctx, loc, types.f64_, 0.5));
298
299 var for_op = try ScfDialect.ForOp.create(ctx, loc, one, upper, one, &.{}, &.{});
300 try entry.addOperation(for_op.op);
301 const body = for_op.getBodyBlock();
302 const iv = body.arguments.items[0];
303
304 const left_index = try appended(body, try ArithDialect.SubOp.create(ctx, loc, iv, one));
305 const right_index = try appended(body, try ArithDialect.AddOp.create(ctx, loc, iv, one));
306 const left = try appended(body, try MemrefDialect.LoadOp.create(ctx, loc, func.getArgument(1), left_index, types.f64_));
307 const center = try appended(body, try MemrefDialect.LoadOp.create(ctx, loc, func.getArgument(1), iv, types.f64_));
308 const right = try appended(body, try MemrefDialect.LoadOp.create(ctx, loc, func.getArgument(1), right_index, types.f64_));
309 const left_term = try appended(body, try ArithDialect.MulOp.create(ctx, loc, quarter, left));
310 const center_term = try appended(body, try ArithDialect.MulOp.create(ctx, loc, half, center));
311 const right_term = try appended(body, try ArithDialect.MulOp.create(ctx, loc, quarter, right));
312 const partial = try appended(body, try ArithDialect.AddOp.create(ctx, loc, left_term, center_term));
313 const result = try appended(body, try ArithDialect.AddOp.create(ctx, loc, partial, right_term));
314 const store = try MemrefDialect.StoreOp.create(ctx, loc, result, func.getArgument(2), iv);
315 try body.addOperation(store.op);
316 const yield_op = try ScfDialect.YieldOp.create(ctx, loc, &.{});
317 try body.addOperation(yield_op.op);
318
319 const ret = try FuncDialect.ReturnOp.create(ctx, loc, &.{});
320 try entry.addOperation(ret.op);
321 }
322
323 fn buildClampsum(ctx: *choir.Context, module_block: *choir.Block, types: Types) !void {
324 var func = try FuncDialect.FuncOp.create(ctx, loc, workload_mod.Kind.clampsum.symbol(), &.{
325 types.i64_, types.memref_i64, types.i64_, types.i64_, types.memref_i64,
326 }, &.{});
327 try module_block.addOperation(func.op);
328 const entry = func.getEntryBlock();
329
330 const zero = try constIndex(ctx, entry, types, 0);
331 const bound = try castToIndex(ctx, entry, types, func.getArgument(0));
332 const step = try constIndex(ctx, entry, types, 1);
333 const acc_init = try appended(entry, try ArithDialect.ConstantOp.createInt(ctx, loc, types.i64_, 0));
334
335 var for_op = try ScfDialect.ForOp.create(ctx, loc, zero, bound, step, &.{acc_init}, &.{types.i64_});
336 try entry.addOperation(for_op.op);
337 const body = for_op.getBodyBlock();
338 const iv = body.arguments.items[0];
339 const acc = body.arguments.items[1];
340
341 const x = try appended(body, try MemrefDialect.LoadOp.create(ctx, loc, func.getArgument(1), iv, types.i64_));
342 const below = try appended(body, try ArithDialect.CmpOp.create(ctx, loc, .lt, x, func.getArgument(2)));
343 const low_clamped = try appended(body, try ArithDialect.SelectOp.create(ctx, loc, below, func.getArgument(2), x));
344 const above = try appended(body, try ArithDialect.CmpOp.create(ctx, loc, .gt, low_clamped, func.getArgument(3)));
345 const clamped = try appended(body, try ArithDialect.SelectOp.create(ctx, loc, above, func.getArgument(3), low_clamped));
346 const next = try appended(body, try ArithDialect.AddOp.create(ctx, loc, acc, clamped));
347 const yield_op = try ScfDialect.YieldOp.create(ctx, loc, &.{next});
348 try body.addOperation(yield_op.op);
349
350 const total = for_op.getResult(0) orelse return error.MissingLoopResult;
351 const store = try MemrefDialect.StoreOp.create(ctx, loc, total, func.getArgument(4), zero);
352 try entry.addOperation(store.op);
353 const ret = try FuncDialect.ReturnOp.create(ctx, loc, &.{});
354 try entry.addOperation(ret.op);
355 }
356
357 test "every battery kernel builds a module" {
358 var arena = alloc_arena.Arena.init(std.testing.allocator);
359 defer arena.deinit();
360
361 inline for (std.meta.tags(workload_mod.Kind)) |kind| {
362 var ctx = try choir.Context.init(arena.allocator(), choir.Context.Limits.testing);
363 defer ctx.deinit(arena.allocator());
364 try choir.dialects.registerAllDialects(&ctx);
365 const module = try build(&ctx, kind);
366 try std.testing.expect(!module.getBodyBlock().operations.isEmpty());
367 }
368 }