lib/accy/src/preparation/dtype.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 dialect_mod = accy_choir.dialect;
7 const shape_analysis = @import("shape/root.zig");
8
9 const ir = choir.ir;
10 const passes = choir.passes;
11 const work = passes.pass.work;
12
13 pub const dtype_legalization_pass_name = "accy-choir-legalize-dtypes";
14 pub const dtype_legalization_pass_description =
15 "Check Accy Choir dtypes supported by backend lowering";
16
17 pub fn dtypeLegalizationPass() passes.Pass {
18 return .{
19 .name = dtype_legalization_pass_name,
20 .description = dtype_legalization_pass_description,
21 .run_fn = runDTypeLegalizationPass,
22 .work_contract = .{
23 .identity = .{ .name = dtype_legalization_pass_name, .version = 1 },
24 .estimate = dtypePassWork,
25 },
26 };
27 }
28
29 fn dtypePassWork(input: work.Input) !work.Bounds {
30 const counts = try work.Census.inspect(input.operation);
31 const units = try work.add(try work.add(counts.atoms, counts.input_bytes), 1);
32 const values = try work.add(try work.add(counts.values, counts.operands), 1);
33 return .{ .work = .{
34 .input_bytes = counts.input_bytes,
35 .structural_visits = try work.multiply(256, try work.multiply(units, values)),
36 } };
37 }
38
39 fn runDTypeLegalizationPass(pass_ctx: *passes.PassContext) passes.PassResult {
40 const analysis = shape_analysis.getShapeLayoutAnalysis(pass_ctx, pass_ctx.op) catch return .failure;
41 if (!legalizeOnOp(pass_ctx.op, analysis)) return .failure;
42 pass_ctx.preserveAllAnalyses();
43 return .success;
44 }
45
46 fn legalizeOnOp(op: *ir.Operation, analysis: *const shape_analysis.ShapeLayoutAnalysis) bool {
47 if (!checkResults(op, analysis, .signature)) return false;
48
49 if (std.mem.startsWith(u8, op.name.name, "accy.")) {
50 if (!legalizeAccyOp(op, analysis)) return false;
51 } else if (std.mem.eql(u8, op.name.name, "func.return")) {
52 if (!checkOperands(op, analysis, .signature)) return false;
53 }
54
55 for (op.regions.items) |*region| {
56 var block_iter = region.getBlocks();
57 while (block_iter.next()) |block| {
58 for (block.arguments.items) |arg| {
59 if (!checkValue(arg, analysis, .signature)) return false;
60 }
61 var current: ?*ir.Operation = @ptrCast(@alignCast(block.operations.head));
62 while (current) |current_op| {
63 if (!legalizeOnOp(current_op, analysis)) return false;
64 current = current_op.next_op;
65 }
66 }
67 }
68 return true;
69 }
70
71 fn legalizeAccyOp(op: *ir.Operation, analysis: *const shape_analysis.ShapeLayoutAnalysis) bool {
72 if (!checkOperands(op, analysis, .signature)) return false;
73 if (!checkResults(op, analysis, .signature)) return false;
74
75 const name = op.name.name;
76 if (isName(name, dialect_mod.AccyDialect.AddOp.operation_name) or
77 isName(name, dialect_mod.AccyDialect.SubOp.operation_name) or
78 isName(name, dialect_mod.AccyDialect.MulOp.operation_name) or
79 isName(name, dialect_mod.AccyDialect.DivOp.operation_name) or
80 isName(name, dialect_mod.AccyDialect.MaxOp.operation_name) or
81 isName(name, dialect_mod.AccyDialect.MinOp.operation_name) or
82 isName(name, dialect_mod.AccyDialect.NegOp.operation_name) or
83 isName(name, dialect_mod.AccyDialect.AbsOp.operation_name) or
84 isName(name, dialect_mod.AccyDialect.ReduceOp.operation_name) or
85 isName(name, dialect_mod.AccyDialect.DotGeneralOp.operation_name))
86 {
87 return checkOperands(op, analysis, .numeric) and checkResults(op, analysis, .numeric);
88 }
89
90 if (isName(name, dialect_mod.AccyDialect.PowOp.operation_name) or
91 isName(name, dialect_mod.AccyDialect.Atan2Op.operation_name) or
92 isName(name, dialect_mod.AccyDialect.ExpOp.operation_name) or
93 isName(name, dialect_mod.AccyDialect.LogOp.operation_name) or
94 isName(name, dialect_mod.AccyDialect.TanhOp.operation_name) or
95 isName(name, dialect_mod.AccyDialect.SqrtOp.operation_name) or
96 isName(name, dialect_mod.AccyDialect.SinOp.operation_name) or
97 isName(name, dialect_mod.AccyDialect.CosOp.operation_name) or
98 isName(name, dialect_mod.AccyDialect.TanOp.operation_name) or
99 isName(name, dialect_mod.AccyDialect.FloorOp.operation_name) or
100 isName(name, dialect_mod.AccyDialect.RoundOp.operation_name) or
101 isName(name, dialect_mod.AccyDialect.TruncOp.operation_name))
102 {
103 return checkOperands(op, analysis, .float) and checkResults(op, analysis, .float);
104 }
105
106 if (isName(name, dialect_mod.AccyDialect.CompareOp.operation_name)) {
107 return op.getNumOperands() == 2 and
108 op.getNumResults() == 1 and
109 checkOperands(op, analysis, .numeric) and
110 checkResults(op, analysis, .bool_only);
111 }
112
113 if (isName(name, dialect_mod.AccyDialect.ConvertOp.operation_name)) {
114 return op.getNumOperands() == 1 and
115 op.getNumResults() == 1 and
116 checkOperands(op, analysis, .numeric) and
117 checkResults(op, analysis, .numeric);
118 }
119
120 if (isName(name, dialect_mod.AccyDialect.SelectOp.operation_name)) {
121 return op.getNumOperands() == 3 and
122 op.getNumResults() == 1 and
123 checkOperandAt(op, 0, analysis, .bool_only) and
124 checkOperandAt(op, 1, analysis, .selectable) and
125 checkOperandAt(op, 2, analysis, .selectable) and
126 checkResults(op, analysis, .selectable);
127 }
128
129 if (isName(name, dialect_mod.AccyDialect.IotaOp.operation_name)) {
130 return op.getNumOperands() == 0 and checkResults(op, analysis, .numeric);
131 }
132
133 if (isName(name, dialect_mod.AccyDialect.GatherOp.operation_name)) {
134 return op.getNumOperands() == 2 and
135 op.getNumResults() == 1 and
136 checkOperandAt(op, 0, analysis, .signature) and
137 checkOperandAt(op, 1, analysis, .index_integer) and
138 checkResults(op, analysis, .signature);
139 }
140
141 if (isName(name, dialect_mod.AccyDialect.ScatterOp.operation_name) or
142 isName(name, dialect_mod.AccyDialect.ScatterAddOp.operation_name))
143 {
144 return op.getNumOperands() == 3 and
145 op.getNumResults() == 1 and
146 checkOperandAt(op, 0, analysis, .signature) and
147 checkOperandAt(op, 1, analysis, .index_integer) and
148 checkOperandAt(op, 2, analysis, .signature) and
149 checkResults(op, analysis, .signature);
150 }
151
152 if (isName(name, dialect_mod.AccyDialect.PadOp.operation_name)) {
153 return op.getNumOperands() == 2 and
154 op.getNumResults() == 1 and
155 checkOperands(op, analysis, .signature) and
156 checkResults(op, analysis, .signature);
157 }
158
159 if (isName(name, dialect_mod.AccyDialect.ConstantOp.operation_name) or
160 isName(name, dialect_mod.AccyDialect.ReshapeOp.operation_name) or
161 isName(name, dialect_mod.AccyDialect.BroadcastOp.operation_name) or
162 isName(name, dialect_mod.AccyDialect.BroadcastInDimOp.operation_name) or
163 isName(name, dialect_mod.AccyDialect.TransposeOp.operation_name) or
164 isName(name, dialect_mod.AccyDialect.SliceOp.operation_name) or
165 isName(name, dialect_mod.AccyDialect.KernelCallOp.operation_name) or
166 isName(name, dialect_mod.AccyDialect.ConcatenateOp.operation_name))
167 {
168 return true;
169 }
170
171 if (isName(name, dialect_mod.AccyDialect.IterateOp.operation_name)) {
172 return checkOperands(op, analysis, .iterable) and checkResults(op, analysis, .iterable);
173 }
174
175 if (isName(name, dialect_mod.AccyDialect.CumsumOp.operation_name)) {
176 return checkOperandAt(op, 0, analysis, .numeric) and checkResults(op, analysis, .numeric);
177 }
178
179 if (isName(name, dialect_mod.AccyDialect.ScratchOp.operation_name)) {
180 return true;
181 }
182
183 if (isName(name, dialect_mod.AccyDialect.IterateYieldOp.operation_name)) {
184 return op.getNumOperands() >= 2 and
185 checkOperandAt(op, 0, analysis, .bool_only);
186 }
187
188 return false;
189 }
190
191 fn isName(actual: []const u8, expected: []const u8) bool {
192 return std.mem.eql(u8, actual, expected);
193 }
194
195 const DTypeSet = enum {
196 signature,
197 numeric,
198 float,
199 bool_only,
200 index_integer,
201 selectable,
202 iterable,
203 };
204
205 fn checkOperands(
206 op: *ir.Operation,
207 analysis: *const shape_analysis.ShapeLayoutAnalysis,
208 set: DTypeSet,
209 ) bool {
210 for (op.operands.items) |operand| {
211 if (!checkValue(operand.value, analysis, set)) return false;
212 }
213 return true;
214 }
215
216 fn checkResults(
217 op: *ir.Operation,
218 analysis: *const shape_analysis.ShapeLayoutAnalysis,
219 set: DTypeSet,
220 ) bool {
221 for (op.results.items) |*result| {
222 if (!checkValue(result, analysis, set)) return false;
223 }
224 return true;
225 }
226
227 fn checkOperandAt(
228 op: *ir.Operation,
229 index: usize,
230 analysis: *const shape_analysis.ShapeLayoutAnalysis,
231 set: DTypeSet,
232 ) bool {
233 const value = op.getOperand(index) orelse return false;
234 return checkValue(value, analysis, set);
235 }
236
237 fn checkValue(
238 value: *ir.Value,
239 analysis: *const shape_analysis.ShapeLayoutAnalysis,
240 set: DTypeSet,
241 ) bool {
242 const info = analysis.get(value) orelse return true;
243 return dtypeAllowed(info.dtype, set);
244 }
245
246 fn dtypeAllowed(dtype: choir_abi.DType, set: DTypeSet) bool {
247 return switch (set) {
248 .signature => isSignatureDType(dtype),
249 .numeric => isNumericDType(dtype),
250 .float => dtype == .f32 or dtype == .f64 or dtype == .f16 or dtype == .bf16,
251 .bool_only => dtype == .i1,
252 .index_integer => isIntegerDType(dtype),
253 .selectable => isNumericDType(dtype) or dtype == .i1,
254 .iterable => isNumericDType(dtype) or dtype == .i1,
255 };
256 }
257
258 fn isSignatureDType(dtype: choir_abi.DType) bool {
259 return isNumericDType(dtype) or dtype == .i1 or dtype == .key;
260 }
261
262 fn isNumericDType(dtype: choir_abi.DType) bool {
263 return switch (dtype) {
264 .f32, .f64, .f16, .bf16 => true,
265 else => isIntegerDType(dtype),
266 };
267 }
268
269 fn isIntegerDType(dtype: choir_abi.DType) bool {
270 return dtype.isSignedInt() or dtype.isUnsignedInt();
271 }
272
273 const testing = std.testing;
274 const semantic = accy_choir.semantic;
275
276 test "dtype legalization accepts static numeric lowering dtypes" {
277 const allocator = testing.allocator;
278
279 var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
280 defer builder.deinit();
281 const f32_4 = try builder.tensor(.f32, &.{4});
282 var fb = try builder.beginFunction("legalize_add4", &.{ f32_4, f32_4 }, &.{f32_4});
283 const sum = try fb.add(fb.parameter(0), fb.parameter(1));
284 try fb.return_(&.{sum});
285 try fb.finish();
286 const module = try builder.finish();
287 defer module.deinit();
288
289 const choir_mod = module.choir_module;
290 const ctx = module.context();
291 var pm = passes.PassManager.init(allocator);
292 defer pm.deinit();
293 try pm.addPass(dtypeLegalizationPass());
294
295 try testing.expectEqual(passes.PassResult.success, pm.run(choir_mod, ctx));
296 try testing.expectEqual(@as(u64, 1), pm.stats.pass_runs);
297 try testing.expectEqual(@as(u64, 0), pm.stats.passes_modified);
298 try testing.expectEqual(@as(u64, 1), pm.stats.analysis_misses);
299 }
300
301 test "dtype legalization accepts unsigned arithmetic before backend capability checks" {
302 const allocator = testing.allocator;
303
304 var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
305 defer builder.deinit();
306 const u32_4 = try builder.tensor(.u32, &.{4});
307 var fb = try builder.beginFunction("legalize_unsigned_add4", &.{ u32_4, u32_4 }, &.{u32_4});
308 const sum = try fb.add(fb.parameter(0), fb.parameter(1));
309 try fb.return_(&.{sum});
310 try fb.finish();
311 const module = try builder.finish();
312 defer module.deinit();
313
314 const choir_mod = module.choir_module;
315 const ctx = module.context();
316 var pm = passes.PassManager.init(allocator);
317 defer pm.deinit();
318 try pm.addPass(dtypeLegalizationPass());
319
320 try testing.expectEqual(passes.PassResult.success, pm.run(choir_mod, ctx));
321 try testing.expectEqual(@as(u64, 1), pm.stats.pass_runs);
322 try testing.expectEqual(@as(u64, 0), pm.stats.passes_modified);
323 }
324
325 test "dtype legalization accepts indexing and padding ops" {
326 const allocator = testing.allocator;
327
328 var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
329 defer builder.deinit();
330 const f32_scalar = try builder.tensor(.f32, &.{});
331 const f32_4 = try builder.tensor(.f32, &.{4});
332 const i32_4 = try builder.tensor(.i32, &.{4});
333 var fb = try builder.beginFunction("legalize_indexing_pad", &.{ f32_4, i32_4 }, &.{f32_4});
334 const zero: f32 = 0;
335 const padding = try fb.constant(f32_scalar, std.mem.asBytes(&zero));
336 const gathered = try fb.gather(fb.parameter(0), fb.parameter(1), f32_4, 0);
337 const scattered = try fb.scatter(fb.parameter(0), fb.parameter(1), gathered, f32_4, 0);
338 const accumulated = try fb.scatterAdd(scattered, fb.parameter(1), gathered, f32_4, 0);
339 const padded = try fb.pad(accumulated, padding, f32_4, &.{1}, &.{-1}, &.{0});
340 try fb.return_(&.{padded});
341 try fb.finish();
342 const module = try builder.finish();
343 defer module.deinit();
344
345 const choir_mod = module.choir_module;
346 const ctx = module.context();
347 var pm = passes.PassManager.init(allocator);
348 defer pm.deinit();
349 try pm.addPass(dtypeLegalizationPass());
350
351 try testing.expectEqual(passes.PassResult.success, pm.run(choir_mod, ctx));
352 try testing.expectEqual(@as(u64, 1), pm.stats.pass_runs);
353 try testing.expectEqual(@as(u64, 0), pm.stats.passes_modified);
354 }
355
356 test "dtype legalization accepts unsigned indexing dtypes" {
357 const allocator = testing.allocator;
358
359 var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
360 defer builder.deinit();
361 const f32_4 = try builder.tensor(.f32, &.{4});
362 const u32_4 = try builder.tensor(.u32, &.{4});
363 var fb = try builder.beginFunction("legalize_unsigned_indexing", &.{ f32_4, u32_4 }, &.{f32_4});
364 const gathered = try fb.gather(fb.parameter(0), fb.parameter(1), f32_4, 0);
365 try fb.return_(&.{gathered});
366 try fb.finish();
367 const module = try builder.finish();
368 defer module.deinit();
369
370 const choir_mod = module.choir_module;
371 const ctx = module.context();
372 var pm = passes.PassManager.init(allocator);
373 defer pm.deinit();
374 try pm.addPass(dtypeLegalizationPass());
375
376 try testing.expectEqual(passes.PassResult.success, pm.run(choir_mod, ctx));
377 try testing.expectEqual(@as(u64, 1), pm.stats.pass_runs);
378 try testing.expectEqual(@as(u64, 0), pm.stats.passes_modified);
379 }
380
381 test "dtype legalization accepts kernel_call contracts" {
382 const allocator = testing.allocator;
383
384 var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
385 defer builder.deinit();
386 const f32_4 = try builder.tensor(.f32, &.{4});
387 var fb = try builder.beginFunction("legalize_kernel_call", &.{f32_4}, &.{f32_4});
388 const call = try fb.kernelCall(
389 &.{fb.parameter(0)},
390 &.{f32_4},
391 .{
392 .target = "accy.custom.scale",
393 .operand_effects = &.{.read},
394 .result_aliases = &.{null},
395 },
396 );
397 try fb.return_(&.{call.getFirstResult()});
398 try fb.finish();
399 const module = try builder.finish();
400 defer module.deinit();
401
402 const choir_mod = module.choir_module;
403 const ctx = module.context();
404 var pm = passes.PassManager.init(allocator);
405 defer pm.deinit();
406 try pm.addPass(dtypeLegalizationPass());
407
408 try testing.expectEqual(passes.PassResult.success, pm.run(choir_mod, ctx));
409 try testing.expectEqual(@as(u64, 1), pm.stats.pass_runs);
410 try testing.expectEqual(@as(u64, 0), pm.stats.passes_modified);
411 }
412
413 test "dtype legalization accepts bf16 arithmetic before backend capability checks" {
414 const allocator = testing.allocator;
415
416 var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
417 defer builder.deinit();
418 const bf16_4 = try builder.tensor(.bf16, &.{4});
419 var fb = try builder.beginFunction("legalize_bf16_add", &.{ bf16_4, bf16_4 }, &.{bf16_4});
420 const sum = try fb.add(fb.parameter(0), fb.parameter(1));
421 try fb.return_(&.{sum});
422 try fb.finish();
423 const module = try builder.finish();
424 defer module.deinit();
425
426 const choir_mod = module.choir_module;
427 const ctx = module.context();
428 var pm = passes.PassManager.init(allocator);
429 defer pm.deinit();
430 try pm.addPass(dtypeLegalizationPass());
431
432 try testing.expectEqual(passes.PassResult.success, pm.run(choir_mod, ctx));
433 try testing.expectEqual(@as(u64, 1), pm.stats.pass_runs);
434 try testing.expectEqual(@as(u64, 0), pm.stats.passes_modified);
435 }
436
437 test "dtype legalization accepts key kernel_call boundaries" {
438 const allocator = testing.allocator;
439
440 var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
441 defer builder.deinit();
442 const key_8 = try builder.tensor(.key, &.{8});
443 const f32_8 = try builder.tensor(.f32, &.{8});
444 var fb = try builder.beginFunction("legalize_key_kernel_call", &.{key_8}, &.{f32_8});
445 const call = try fb.kernelCall(
446 &.{fb.parameter(0)},
447 &.{f32_8},
448 .{
449 .target = "accy.kernel.random.philox_key_uniform_family_10r_64_f32",
450 .operand_effects = &.{.read},
451 .result_aliases = &.{null},
452 },
453 );
454 try fb.return_(&.{call.getFirstResult()});
455 try fb.finish();
456 const module = try builder.finish();
457 defer module.deinit();
458
459 const choir_mod = module.choir_module;
460 const ctx = module.context();
461 var pm = passes.PassManager.init(allocator);
462 defer pm.deinit();
463 try pm.addPass(dtypeLegalizationPass());
464
465 try testing.expectEqual(passes.PassResult.success, pm.run(choir_mod, ctx));
466 try testing.expectEqual(@as(u64, 1), pm.stats.pass_runs);
467 try testing.expectEqual(@as(u64, 0), pm.stats.passes_modified);
468 }
469
470 test "dtype legalization accepts bool iterate carries" {
471 const allocator = testing.allocator;
472
473 var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
474 defer builder.deinit();
475 const bool_4 = try builder.tensor(.i1, &.{4});
476 var fb = try builder.beginFunction("legalize_bool_iterate", &.{bool_4}, &.{bool_4});
477 var iterate = try fb.beginIterate(&.{fb.parameter(0)}, 4);
478 try iterate.yield_(iterate.carry(0), &.{iterate.carry(0)});
479 try fb.return_(&.{iterate.result(0)});
480 try fb.finish();
481 const module = try builder.finish();
482 defer module.deinit();
483
484 const choir_mod = module.choir_module;
485 const ctx = module.context();
486 var pm = passes.PassManager.init(allocator);
487 defer pm.deinit();
488 try pm.addPass(dtypeLegalizationPass());
489
490 try testing.expectEqual(passes.PassResult.success, pm.run(choir_mod, ctx));
491 try testing.expectEqual(@as(u64, 1), pm.stats.pass_runs);
492 try testing.expectEqual(@as(u64, 0), pm.stats.passes_modified);
493 }
494
495 test "key arithmetic is rejected upstream by the semantic dialect" {
496 const allocator = testing.allocator;
497
498 var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
499 defer builder.deinit();
500 const key_4 = try builder.tensor(.key, &.{4});
501 var fb = try builder.beginFunction("illegal_key_add", &.{ key_4, key_4 }, &.{key_4});
502 try testing.expectError(error.UnsupportedDType, fb.add(fb.parameter(0), fb.parameter(1)));
503 }
504
505 test "dtype legalization rejects bool iota before lowering" {
506 const allocator = testing.allocator;
507
508 var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
509 defer builder.deinit();
510 const bool_4 = try builder.tensor(.i1, &.{4});
511 var fb = try builder.beginFunction("illegal_bool_iota", &.{}, &.{bool_4});
512 const ramp = try fb.iota(bool_4, 0);
513 try fb.return_(&.{ramp});
514 try fb.finish();
515 const module = try builder.finish();
516 defer module.deinit();
517
518 const choir_mod = module.choir_module;
519 const ctx = module.context();
520 var pm = passes.PassManager.init(allocator);
521 defer pm.deinit();
522 try pm.addPass(dtypeLegalizationPass());
523
524 try testing.expectEqual(passes.PassResult.failure, pm.run(choir_mod, ctx));
525 try testing.expectEqual(@as(u64, 1), pm.stats.pass_failures);
526 }