lib/accy/src/preparation/shape/pass.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 
  8 const ir = choir.ir;
  9 const passes = choir.passes;
 10 const work = passes.pass.work;
 11 
 12 pub const shape_layout_analysis_name = "accy-choir-shape-layout";
 13 pub const shape_layout_pass_name = "accy-choir-shape-layout-propagate";
 14 pub const shape_layout_pass_description =
 15     "Decode Accy Choir tensor shape and row-major layout facts";
 16 
 17 pub const TensorInfo = struct {
 18     dtype: choir_abi.DType,
 19     dims: []i64,
 20     element_count: ?u64,
 21     row_major_strides: ?[]u64,
 22 
 23     pub fn rank(self: TensorInfo) usize {
 24         return self.dims.len;
 25     }
 26 
 27     pub fn hasStaticLayout(self: TensorInfo) bool {
 28         return self.element_count != null and self.row_major_strides != null;
 29     }
 30 
 31     fn deinit(self: *TensorInfo, allocator: std.mem.Allocator) void {
 32         allocator.free(self.dims);
 33         if (self.row_major_strides) |strides| allocator.free(strides);
 34         self.* = undefined;
 35     }
 36 };
 37 
 38 pub const ShapeLayoutAnalysis = struct {
 39     allocator: std.mem.Allocator,
 40     tensors: std.AutoHashMap(*ir.Value, TensorInfo),
 41     tensor_value_count: usize = 0,
 42     accy_op_count: usize = 0,
 43 
 44     pub fn init(allocator: std.mem.Allocator) ShapeLayoutAnalysis {
 45         return .{
 46             .allocator = allocator,
 47             .tensors = std.AutoHashMap(*ir.Value, TensorInfo).init(allocator),
 48         };
 49     }
 50 
 51     pub fn deinit(self: *ShapeLayoutAnalysis) void {
 52         var iter = self.tensors.valueIterator();
 53         while (iter.next()) |info| {
 54             info.deinit(self.allocator);
 55         }
 56         self.tensors.deinit();
 57         self.* = undefined;
 58     }
 59 
 60     pub fn get(self: *const ShapeLayoutAnalysis, value: *ir.Value) ?TensorInfo {
 61         return self.tensors.get(value);
 62     }
 63 
 64     fn recordValue(self: *ShapeLayoutAnalysis, value: *ir.Value) !void {
 65         const info = try tensorInfoFromType(self.allocator, value.type) orelse return;
 66         errdefer {
 67             var mutable = info;
 68             mutable.deinit(self.allocator);
 69         }
 70         self.tensors.putAssumeCapacityNoClobber(value, info);
 71         self.tensor_value_count += 1;
 72     }
 73 };
 74 
 75 const ShapeWork = struct {
 76     visits: u64 = 0,
 77     values: u64 = 0,
 78     dimensions: u64 = 0,
 79     key_bytes: u64 = 0,
 80 
 81     fn inspect(op: *ir.Operation) !ShapeWork {
 82         var result: ShapeWork = .{};
 83         _ = try op.walk(.{ .order = .pre_order }, &result, visit);
 84         return result;
 85     }
 86 
 87     fn visit(self: *ShapeWork, op: *ir.Operation) !ir.Operation.WalkResult {
 88         self.visits = try work.add(self.visits, 1);
 89         for (op.results.items) |*result| try self.include(result.type);
 90         for (op.regions.items) |*region| {
 91             self.visits = try work.add(self.visits, 1);
 92             var blocks = region.getBlocks();
 93             while (blocks.next()) |block| {
 94                 self.visits = try work.add(self.visits, 1);
 95                 for (block.arguments.items) |arg| try self.include(arg.type);
 96             }
 97         }
 98         return .advance;
 99     }
100 
101     fn include(self: *ShapeWork, typ: ir.Type) !void {
102         self.visits = try work.add(self.visits, 1);
103         const name = typ.getDialectTypeName() orelse return;
104         self.key_bytes = try work.add(self.key_bytes, name.len);
105         if (!std.mem.eql(u8, name, dialect_mod.tensor_type_name)) return;
106         const key = typ.getDialectParamKey() orelse return;
107         self.key_bytes = try work.add(self.key_bytes, key.len);
108         const comma = std.mem.indexOfScalar(u8, key, ',') orelse return;
109         _ = choir_abi.DType.fromName(key[0..comma]) orelse return;
110         const dims = key[comma + 1 ..];
111         var rank: u64 = @intFromBool(dims.len != 0);
112         for (dims) |byte| if (byte == 'x') {
113             rank = try work.add(rank, 1);
114         };
115         self.values = try work.add(self.values, 1);
116         self.dimensions = try work.add(self.dimensions, rank);
117     }
118 
119     fn storage(self: ShapeWork) !u64 {
120         const alignment = @max(@alignOf(usize), @alignOf(TensorInfo), @alignOf(ShapeLayoutAnalysis));
121         var bytes: u64 = @sizeOf(ShapeLayoutAnalysis) + alignment;
122         const capacity = try work.hashMapCapacity(self.values);
123         if (capacity != 0) {
124             const table = try work.multiply(capacity, 1 + @sizeOf(*ir.Value) + @sizeOf(TensorInfo));
125             bytes = try work.add(bytes, try work.add(table, 4 * @sizeOf(usize) + 3 * alignment));
126         }
127         bytes = try work.add(bytes, try work.multiply(self.dimensions, 16));
128         bytes = try work.add(bytes, try work.multiply(self.values, 2 * @alignOf(u64)));
129         if (bytes > std.math.maxInt(usize)) return error.WorkOverflow;
130         return bytes;
131     }
132 };
133 
134 fn analysisWork(input: work.Input) !work.Bounds {
135     const counts = try ShapeWork.inspect(input.operation);
136     const bytes = try counts.storage();
137     const traversal = try work.add(try work.add(counts.visits, counts.key_bytes), counts.dimensions);
138     const probes = try work.multiply(counts.values, try work.hashMapCapacity(counts.values));
139     return .{
140         .work = .{
141             .input_bytes = counts.key_bytes,
142             .structural_visits = try work.add(try work.multiply(traversal, 8), probes),
143             .analysis_computations = 1,
144             .allocation_capacity = bytes,
145         },
146         .workspace = bytes,
147         .retained_storage = bytes,
148     };
149 }
150 
151 fn propagationWork(_: work.Input) !work.Bounds {
152     return .{ .work = .{ .structural_visits = 1 } };
153 }
154 
155 const ShapeLayoutAnalysisRegistration = passes.Analysis(
156     ShapeLayoutAnalysis,
157     shape_layout_analysis_name,
158     &.{},
159     computeShapeLayoutAnalysis,
160     cleanupShapeLayoutAnalysis,
161     .{ .identity = .{ .name = shape_layout_analysis_name, .version = 1 }, .estimate = analysisWork },
162 );
163 
164 pub const shape_layout_analysis_descriptor = ShapeLayoutAnalysisRegistration.descriptor;
165 
166 pub fn getShapeLayoutAnalysis(
167     pass_ctx: *passes.PassContext,
168     op: *ir.Operation,
169 ) !*ShapeLayoutAnalysis {
170     return ShapeLayoutAnalysisRegistration.get(pass_ctx, op);
171 }
172 
173 pub fn shapeLayoutPropagationPass() passes.Pass {
174     return .{
175         .name = shape_layout_pass_name,
176         .description = shape_layout_pass_description,
177         .run_fn = runShapeLayoutPropagationPass,
178         .work_contract = .{
179             .identity = .{ .name = shape_layout_pass_name, .version = 1 },
180             .estimate = propagationWork,
181         },
182     };
183 }
184 
185 fn runShapeLayoutPropagationPass(pass_ctx: *passes.PassContext) passes.PassResult {
186     _ = getShapeLayoutAnalysis(pass_ctx, pass_ctx.op) catch return .failure;
187     pass_ctx.preserveAllAnalyses();
188     return .success;
189 }
190 
191 fn computeShapeLayoutAnalysis(
192     pass_ctx: *passes.PassContext,
193     op: *ir.Operation,
194 ) anyerror!*ShapeLayoutAnalysis {
195     const analysis = try pass_ctx.allocator.create(ShapeLayoutAnalysis);
196     analysis.* = ShapeLayoutAnalysis.init(pass_ctx.allocator);
197     errdefer {
198         analysis.deinit();
199         pass_ctx.allocator.destroy(analysis);
200     }
201 
202     const counts = try ShapeWork.inspect(op);
203     _ = try work.hashMapCapacity(counts.values);
204     try analysis.tensors.ensureTotalCapacity(@intCast(counts.values));
205     try collectShapeLayoutFacts(analysis, op);
206     std.debug.assert(analysis.tensor_value_count == counts.values);
207     std.debug.assert(analysis.tensors.count() == analysis.tensor_value_count);
208     return analysis;
209 }
210 
211 fn cleanupShapeLayoutAnalysis(analysis: *ShapeLayoutAnalysis, allocator: std.mem.Allocator) void {
212     analysis.deinit();
213     allocator.destroy(analysis);
214 }
215 
216 fn collectShapeLayoutFacts(analysis: *ShapeLayoutAnalysis, op: *ir.Operation) !void {
217     if (std.mem.startsWith(u8, op.name.name, "accy.")) {
218         analysis.accy_op_count += 1;
219     }
220 
221     for (op.results.items) |*result| {
222         try analysis.recordValue(result);
223     }
224 
225     for (op.regions.items) |*region| {
226         var block_iter = region.getBlocks();
227         while (block_iter.next()) |block| {
228             for (block.arguments.items) |arg| {
229                 try analysis.recordValue(arg);
230             }
231             var current: ?*ir.Operation = @ptrCast(@alignCast(block.operations.head));
232             while (current) |current_op| {
233                 try collectShapeLayoutFacts(analysis, current_op);
234                 current = current_op.next_op;
235             }
236         }
237     }
238 }
239 
240 const StaticLayout = struct {
241     element_count: u64,
242     strides: []u64,
243 };
244 
245 fn tensorInfoFromType(allocator: std.mem.Allocator, typ: ir.Type) !?TensorInfo {
246     const type_name = typ.getDialectTypeName() orelse return null;
247     if (!std.mem.eql(u8, type_name, dialect_mod.tensor_type_name)) return null;
248     const key = typ.getDialectParamKey() orelse return null;
249     const comma = std.mem.indexOfScalar(u8, key, ',') orelse return null;
250     const dtype_name = key[0..comma];
251     const dtype = choir_abi.DType.fromName(dtype_name) orelse return null;
252     const dims = try parseDimsAlloc(allocator, key[comma + 1 ..]);
253     errdefer allocator.free(dims);
254 
255     const static_layout = try computeStaticLayoutAlloc(allocator, dims);
256     return .{
257         .dtype = dtype,
258         .dims = dims,
259         .element_count = if (static_layout) |layout| layout.element_count else null,
260         .row_major_strides = if (static_layout) |layout| layout.strides else null,
261     };
262 }
263 
264 fn parseDimsAlloc(allocator: std.mem.Allocator, dims_text: []const u8) ![]i64 {
265     if (dims_text.len == 0) return try allocator.alloc(i64, 0);
266 
267     var dim_count: usize = 1;
268     for (dims_text) |ch| {
269         if (ch == 'x') dim_count += 1;
270     }
271 
272     const dims = try allocator.alloc(i64, dim_count);
273     errdefer allocator.free(dims);
274     var iter = std.mem.splitScalar(u8, dims_text, 'x');
275     var index: usize = 0;
276     while (iter.next()) |part| {
277         if (part.len == 0 or index >= dim_count) return error.InvalidTensorShape;
278         dims[index] = std.fmt.parseInt(i64, part, 10) catch return error.InvalidTensorShape;
279         index += 1;
280     }
281     if (index != dim_count) return error.InvalidTensorShape;
282     return dims;
283 }
284 
285 fn computeStaticLayoutAlloc(allocator: std.mem.Allocator, dims: []const i64) !?StaticLayout {
286     const strides = try allocator.alloc(u64, dims.len);
287     errdefer allocator.free(strides);
288 
289     var stride: u64 = 1;
290     var i = dims.len;
291     while (i > 0) {
292         i -= 1;
293         const dim = dims[i];
294         if (dim < 0) {
295             allocator.free(strides);
296             return null;
297         }
298         strides[i] = stride;
299         const dim_u64: u64 = @intCast(dim);
300         stride = std.math.mul(u64, stride, dim_u64) catch {
301             allocator.free(strides);
302             return null;
303         };
304     }
305 
306     return .{ .element_count = stride, .strides = strides };
307 }
308 
309 const testing = std.testing;
310 const semantic = accy_choir.semantic;
311 
312 fn readSymbolName(func: *ir.Operation) ?[]const u8 {
313     return ir.SymbolTable.getSymbolName(func);
314 }
315 
316 fn findOpNamedInBlock(block: *ir.Block, name: []const u8) ?*ir.Operation {
317     var iter = block.operations.head;
318     while (iter) |op_ptr| {
319         const op: *ir.Operation = @ptrCast(@alignCast(op_ptr));
320         if (std.mem.eql(u8, op.name.name, name)) return op;
321         iter = op.next_op;
322     }
323     return null;
324 }
325 
326 test "shape layout analysis records accy tensor values" {
327     const allocator = testing.allocator;
328 
329     var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
330     defer builder.deinit();
331     const f32_2x3 = try builder.tensor(.f32, &.{ 2, 3 });
332     const f32_3x2 = try builder.tensor(.f32, &.{ 3, 2 });
333     var fb = try builder.beginFunction("shape_layout", &.{f32_2x3}, &.{f32_3x2});
334     const reshaped = try fb.reshape(fb.parameter(0), f32_3x2, &.{ 3, 2 });
335     try fb.return_(&.{reshaped});
336     try fb.finish();
337     const module = try builder.finish();
338     defer module.deinit();
339 
340     const choir_mod = module.choir_module;
341     const ctx = module.context();
342     var cache = passes.AnalysisCache.init(allocator, null);
343     defer cache.deinit();
344     var pass_ctx = passes.PassContext.init(choir_mod, ctx, allocator, &cache);
345     defer pass_ctx.deinit();
346 
347     const analysis = try getShapeLayoutAnalysis(&pass_ctx, choir_mod);
348     try testing.expectEqual(@as(usize, 3), analysis.tensor_value_count);
349     try testing.expectEqual(@as(usize, 1), analysis.accy_op_count);
350 
351     const module_body = choir_mod.getRegion(0).?.getEntryBlock().?;
352     const func = ir.inspection.functionByNameInBlock(module_body, "shape_layout") orelse
353         return error.TestExpectedFunc;
354     const entry = func.getRegion(0).?.getEntryBlock().?;
355     const arg_info = analysis.get(entry.getArgument(0).?) orelse
356         return error.TestExpectedTensorInfo;
357     try testing.expectEqual(choir_abi.DType.f32, arg_info.dtype);
358     try testing.expectEqualSlices(i64, &.{ 2, 3 }, arg_info.dims);
359     try testing.expectEqual(@as(?u64, 6), arg_info.element_count);
360     try testing.expectEqualSlices(u64, &.{ 3, 1 }, arg_info.row_major_strides.?);
361 
362     const reshape_op = findOpNamedInBlock(
363         entry,
364         dialect_mod.AccyDialect.ReshapeOp.operation_name,
365     ) orelse return error.TestExpectedReshape;
366     const result_info = analysis.get(reshape_op.getResult(0).?) orelse
367         return error.TestExpectedTensorInfo;
368     try testing.expectEqualSlices(i64, &.{ 3, 2 }, result_info.dims);
369     try testing.expectEqual(@as(?u64, 6), result_info.element_count);
370     try testing.expectEqualSlices(u64, &.{ 2, 1 }, result_info.row_major_strides.?);
371 }
372 
373 test "shape layout propagation pass populates analysis without modification" {
374     const allocator = testing.allocator;
375 
376     var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
377     defer builder.deinit();
378     const f32_4 = try builder.tensor(.f32, &.{4});
379     var fb = try builder.beginFunction("shape_pass_add4", &.{ f32_4, f32_4 }, &.{f32_4});
380     const sum = try fb.add(fb.parameter(0), fb.parameter(1));
381     try fb.return_(&.{sum});
382     try fb.finish();
383     const module = try builder.finish();
384     defer module.deinit();
385 
386     const choir_mod = module.choir_module;
387     const ctx = module.context();
388     var pm = passes.PassManager.init(allocator);
389     defer pm.deinit();
390     try pm.addPass(shapeLayoutPropagationPass());
391 
392     const revision = choir.product.revision;
393     const ledger = try revision.AccountingV1.create(allocator, .{
394         .allowance = revision.WorkVector.uniform(1 << 30),
395         .workspace = 1 << 24,
396         .events = 8,
397     }, &.{.{ .name = shape_layout_pass_name, .version = 1 }});
398     defer ledger.destroy();
399     var cache = try passes.AnalysisCache.initAccounted(allocator, null, ledger, .{}, 1);
400     defer cache.deinit();
401     try testing.expectEqual(
402         passes.PassResult.success,
403         pm.runWithAnalysisCache(choir_mod, ctx, &cache, .{}),
404     );
405     try ledger.producersComplete();
406     try testing.expectEqual(@as(u64, 1), pm.stats.pass_runs);
407     try testing.expectEqual(@as(u64, 0), pm.stats.passes_modified);
408     try testing.expectEqual(@as(u64, 1), pm.stats.analysis_misses);
409 }
410 
411 test "shape layout storage bound covers scalar dynamic overflow and table growth" {
412     const shapes = [_][]const i64{ &.{}, &.{ 2, 3 }, &.{ -1, 3 }, &.{ std.math.maxInt(i64), 3 } };
413     for (shapes) |dims| {
414         for ([_]usize{ 0, 1, 6, 7, 13, 64 }) |count| {
415             try checkShapeStorage(dims, count);
416         }
417     }
418     try testing.expectError(
419         error.WorkOverflow,
420         (ShapeWork{ .values = std.math.maxInt(u64) }).storage(),
421     );
422     try testing.expectError(
423         error.WorkOverflow,
424         (ShapeWork{ .dimensions = std.math.maxInt(u64) }).storage(),
425     );
426 }
427 
428 fn checkShapeStorage(dims: []const i64, count: usize) !void {
429     var builder = try semantic.Builder.init(
430         testing.allocator,
431         semantic.Builder.ContextLimits.standard,
432     );
433     defer builder.deinit();
434     const typ = try builder.tensor(.f32, dims);
435     var parameters: [64]ir.Type = undefined;
436     @memset(parameters[0..count], typ);
437     var function = try builder.beginFunction("storage", parameters[0..count], &.{});
438     try function.return_(&.{});
439     try function.finish();
440     const module = try builder.finish();
441     defer module.deinit();
442     const bounds = try analysisWork(.{ .operation = module.choir_module });
443     const bytes = try testing.allocator.alloc(u8, @intCast(bounds.workspace));
444     defer testing.allocator.free(bytes);
445     var storage = @import("alloc_fixed").Tracked.init(bytes);
446     var cache = passes.AnalysisCache.init(testing.allocator, null);
447     defer cache.deinit();
448     var pass_ctx = passes.PassContext.init(
449         module.choir_module,
450         module.context(),
451         storage.allocator(),
452         &cache,
453     );
454     defer pass_ctx.deinit();
455     const analysis = try computeShapeLayoutAnalysis(&pass_ctx, module.choir_module);
456     defer cleanupShapeLayoutAnalysis(analysis, storage.allocator());
457     try testing.expectEqual(count, analysis.tensor_value_count);
458     try testing.expect(!storage.exhausted);
459     try testing.expect(storage.status().high_water_bytes <= bounds.workspace);
460     try testing.expect(storage.status().high_water_bytes >= @sizeOf(ShapeLayoutAnalysis));
461     var entries = analysis.tensors.valueIterator();
462     while (entries.next()) |info| try testing.expectEqualSlices(i64, dims, info.dims);
463 }
464 
465 test "shape layout production accounting rejects below the exact charge" {
466     var builder = try semantic.Builder.init(
467         testing.allocator,
468         semantic.Builder.ContextLimits.standard,
469     );
470     defer builder.deinit();
471     const typ = try builder.tensor(.f32, &.{ 2, 3 });
472     var function = try builder.beginFunction("budget", &.{typ}, &.{typ});
473     try function.return_(&.{function.parameter(0)});
474     try function.finish();
475     const module = try builder.finish();
476     defer module.deinit();
477     const analysis = try analysisWork(.{ .operation = module.choir_module });
478     const pass = try propagationWork(.{ .operation = module.choir_module });
479     const cache_bytes = try passes.AnalysisCache.storageBound(1);
480     const producer_charge = try analysis.work.add(pass.work);
481     const charge = try producer_charge.add(.{ .allocation_capacity = cache_bytes });
482     for ([_]i8{ -1, 0, 1 }) |offset| {
483         var allowance = charge;
484         allowance.structural_visits = @intCast(@as(i128, charge.structural_visits) + offset);
485         try checkShapeBudget(module.choir_module, module.context(), allowance, offset >= 0);
486     }
487 }
488 
489 fn checkShapeBudget(
490     operation: *ir.Operation,
491     context: *ir.Context,
492     allowance: choir.product.revision.WorkVector,
493     admitted: bool,
494 ) !void {
495     const revision = choir.product.revision;
496     const ledger = try revision.AccountingV1.create(testing.allocator, .{
497         .allowance = allowance,
498         .workspace = 1 << 24,
499         .events = 8,
500     }, &.{.{ .name = shape_layout_pass_name, .version = 1 }});
501     defer ledger.destroy();
502     var cache = try passes.AnalysisCache.initAccounted(testing.allocator, null, ledger, .{}, 1);
503     defer cache.deinit();
504     var manager = passes.PassManager.init(testing.allocator);
505     defer manager.deinit();
506     try manager.addPass(shapeLayoutPropagationPass());
507     const result = manager.runWithAnalysisCache(operation, context, &cache, .{});
508     try testing.expectEqual(if (admitted) passes.PassResult.success else .failure, result);
509     const events = ledger.view().events;
510     try testing.expectEqual(@as(usize, 3), events.len);
511     try testing.expectEqual(revision.receipt.Phase.analysis, events[2].phase);
512     try testing.expectEqual(admitted, events[2].admitted);
513     if (admitted) {
514         try ledger.producersComplete();
515         try testing.expectEqual(revision.receipt.Outcome.success, events[2].outcome);
516     } else {
517         try testing.expectEqual(revision.receipt.Outcome.exhausted, ledger.view().outcome);
518         try testing.expectEqual(
519             revision.WorkVector.Component.structural_visits,
520             ledger.view().exceeded.?,
521         );
522         try testing.expectEqual(@as(usize, 0), cache.entries.count());
523     }
524 }