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 }