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 }