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 }