lib/accy/src/preparation/layout.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 bufferization = @import("bufferization/root.zig");
  6 const accy_choir = @import("../choir/root.zig");
  7 const dialect_mod = accy_choir.dialect;
  8 const memory_space = @import("root.zig").memory;
  9 
 10 const ir = choir.ir;
 11 const passes = choir.passes;
 12 const accounting = passes.pass.work;
 13 
 14 pub const layout_plan_analysis_name = "accy-choir-layout-plan";
 15 pub const layout_planning_pass_name = "accy-choir-plan-layouts";
 16 pub const layout_planning_pass_description =
 17     "Plan Accy Choir buffer layouts before scheduling";
 18 
 19 pub const LayoutKind = accy_choir.record.memory.LayoutKind;
 20 
 21 pub const LayoutAssignment = struct {
 22     slot_id: usize,
 23     value: *ir.Value,
 24     producer: ?*ir.Operation,
 25     role: bufferization.BufferRole,
 26     dtype: choir_abi.DType,
 27     memory_space: memory_space.MemorySpace,
 28     kind: LayoutKind,
 29     rank: usize,
 30     dims: []i64,
 31     element_strides: ?[]u64,
 32     minor_to_major: []usize,
 33     element_count: ?u64,
 34     byte_size: ?u64,
 35     element_size: u64,
 36     alignment: u64,
 37     contiguous: bool,
 38     static_layout: bool,
 39 
 40     fn init(
 41         allocator: std.mem.Allocator,
 42         slot: bufferization.BufferSlot,
 43         memory_assignment: memory_space.MemorySpaceAssignment,
 44     ) !LayoutAssignment {
 45         const dims = try allocator.dupe(i64, slot.dims);
 46         errdefer allocator.free(dims);
 47 
 48         var strides: ?[]u64 = null;
 49         if (slot.row_major_strides) |existing| {
 50             strides = try allocator.dupe(u64, existing);
 51             errdefer if (strides) |owned| allocator.free(owned);
 52         }
 53 
 54         const minor_to_major = try rowMajorMinorToMajorAlloc(allocator, slot.dims.len);
 55         errdefer allocator.free(minor_to_major);
 56 
 57         const kind = layoutKindForSlot(slot);
 58         return .{
 59             .slot_id = slot.id,
 60             .value = slot.value,
 61             .producer = slot.producer,
 62             .role = slot.role,
 63             .dtype = slot.dtype,
 64             .memory_space = memory_assignment.space,
 65             .kind = kind,
 66             .rank = slot.dims.len,
 67             .dims = dims,
 68             .element_strides = strides,
 69             .minor_to_major = minor_to_major,
 70             .element_count = slot.element_count,
 71             .byte_size = slot.byte_size,
 72             .element_size = @as(u64, slot.dtype.sizeOf()),
 73             .alignment = @as(u64, slot.dtype.alignOf()),
 74             .contiguous = true,
 75             .static_layout = kind.hasStaticStrides(),
 76         };
 77     }
 78 
 79     fn deinit(self: *LayoutAssignment, allocator: std.mem.Allocator) void {
 80         allocator.free(self.dims);
 81         if (self.element_strides) |strides| allocator.free(strides);
 82         allocator.free(self.minor_to_major);
 83         self.* = undefined;
 84     }
 85 
 86     pub fn hasStaticByteSize(self: LayoutAssignment) bool {
 87         return self.byte_size != null;
 88     }
 89 };
 90 
 91 pub const LayoutPlanAnalysis = struct {
 92     allocator: std.mem.Allocator,
 93     assignments: std.ArrayListUnmanaged(LayoutAssignment),
 94     slot_to_assignment: std.AutoHashMap(usize, usize),
 95     scalar_layout_count: usize = 0,
 96     row_major_layout_count: usize = 0,
 97     dynamic_row_major_layout_count: usize = 0,
 98     host_slot_count: usize = 0,
 99     device_global_slot_count: usize = 0,
100     device_constant_slot_count: usize = 0,
101     device_shared_slot_count: usize = 0,
102     unified_slot_count: usize = 0,
103     dynamic_slot_count: usize = 0,
104     elided_value_count: usize = 0,
105     total_static_bytes: u64 = 0,
106 
107     pub fn init(allocator: std.mem.Allocator) LayoutPlanAnalysis {
108         return .{
109             .allocator = allocator,
110             .assignments = .empty,
111             .slot_to_assignment = std.AutoHashMap(usize, usize).init(allocator),
112         };
113     }
114 
115     pub fn deinit(self: *LayoutPlanAnalysis) void {
116         for (self.assignments.items) |*assignment| {
117             assignment.deinit(self.allocator);
118         }
119         self.assignments.deinit(self.allocator);
120         self.slot_to_assignment.deinit();
121         self.* = undefined;
122     }
123 
124     pub fn assignmentCount(self: LayoutPlanAnalysis) usize {
125         return self.assignments.items.len;
126     }
127 
128     pub fn getAssignmentForSlot(
129         self: *const LayoutPlanAnalysis,
130         slot_id: usize,
131     ) ?*const LayoutAssignment {
132         const index = self.slot_to_assignment.get(slot_id) orelse return null;
133         return &self.assignments.items[index];
134     }
135 
136     pub fn getAssignmentForValue(
137         self: *const LayoutPlanAnalysis,
138         buffers: *const bufferization.BufferPlanAnalysis,
139         value: *ir.Value,
140     ) ?*const LayoutAssignment {
141         const slot = buffers.getSlot(value) orelse return null;
142         return self.getAssignmentForSlot(slot.id);
143     }
144 
145     fn addAssignment(self: *LayoutPlanAnalysis, assignment: LayoutAssignment) !void {
146         if (self.slot_to_assignment.contains(assignment.slot_id)) {
147             return error.DuplicateLayoutAssignment;
148         }
149         const index = self.assignments.items.len;
150         try self.slot_to_assignment.put(assignment.slot_id, index);
151         errdefer _ = self.slot_to_assignment.remove(assignment.slot_id);
152         try self.assignments.append(self.allocator, assignment);
153 
154         switch (assignment.kind) {
155             .scalar => self.scalar_layout_count += 1,
156             .row_major => self.row_major_layout_count += 1,
157             .dynamic_row_major => self.dynamic_row_major_layout_count += 1,
158         }
159         switch (assignment.memory_space) {
160             .host => self.host_slot_count += 1,
161             .device_global => self.device_global_slot_count += 1,
162             .device_constant => self.device_constant_slot_count += 1,
163             .device_shared => self.device_shared_slot_count += 1,
164             .unified => self.unified_slot_count += 1,
165         }
166         if (assignment.byte_size) |bytes| {
167             self.total_static_bytes += bytes;
168         } else {
169             self.dynamic_slot_count += 1;
170         }
171     }
172 };
173 
174 const LayoutWork = struct {
175     input: accounting.Census,
176 
177     fn storage(self: LayoutWork) !u64 {
178         var bytes: u64 = @sizeOf(LayoutPlanAnalysis) + @alignOf(LayoutPlanAnalysis);
179         bytes = try accounting.add(
180             bytes,
181             try accounting.arrayListGrowth(LayoutAssignment, self.input.values),
182         );
183         bytes = try accounting.add(
184             bytes,
185             try accounting.hashMapGrowth(usize, usize, self.input.values),
186         );
187         const element = @sizeOf(i64) + @sizeOf(u64) + @sizeOf(usize);
188         bytes = try accounting.add(bytes, try accounting.multiply(self.input.input_bytes, element));
189         const alignment = @alignOf(i64) + @alignOf(u64) + @alignOf(usize);
190         bytes = try accounting.add(bytes, try accounting.multiply(self.input.values, alignment));
191         if (bytes > std.math.maxInt(usize)) return error.WorkOverflow;
192         return bytes;
193     }
194 
195     fn bounds(self: LayoutWork) !accounting.Bounds {
196         const bytes = try self.storage();
197         const input_units = try accounting.add(self.input.atoms, self.input.input_bytes);
198         const units = try accounting.add(input_units, 1);
199         const probes = try accounting.add(try accounting.hashMapCapacity(self.input.values), 1);
200         const visits = try accounting.multiply(32, try accounting.multiply(units, probes));
201         return .{
202             .work = .{
203                 .input_bytes = self.input.input_bytes,
204                 .structural_visits = visits,
205                 .analysis_computations = 1,
206                 .allocation_capacity = bytes,
207             },
208             .workspace = bytes,
209             .retained_storage = bytes,
210         };
211     }
212 };
213 
214 fn layoutAnalysisWork(input: accounting.Input) !accounting.Bounds {
215     return (LayoutWork{ .input = try accounting.Census.inspect(input.operation) }).bounds();
216 }
217 
218 fn layoutPassWork(_: accounting.Input) !accounting.Bounds {
219     return .{ .work = .{ .structural_visits = 1 } };
220 }
221 
222 pub const layout_plan_analysis_descriptor = passes.AnalysisDescriptor{
223     .id = passes.analysisId(layout_plan_analysis_name),
224     .name = layout_plan_analysis_name,
225     .work_contract = .{
226         .identity = .{ .name = layout_plan_analysis_name, .version = 1 },
227         .estimate = layoutAnalysisWork,
228     },
229 };
230 
231 pub fn getLayoutPlanAnalysis(
232     pass_ctx: *passes.PassContext,
233     op: *ir.Operation,
234 ) !*LayoutPlanAnalysis {
235     const ptr = try pass_ctx.getAnalysis(
236         op,
237         &layout_plan_analysis_descriptor,
238         computeLayoutPlanAnalysis,
239         cleanupLayoutPlanAnalysis,
240     );
241     return @ptrCast(@alignCast(ptr));
242 }
243 
244 pub fn layoutPlanningPass() passes.Pass {
245     return .{
246         .name = layout_planning_pass_name,
247         .description = layout_planning_pass_description,
248         .run_fn = runLayoutPlanningPass,
249         .work_contract = .{
250             .identity = .{ .name = layout_planning_pass_name, .version = 1 },
251             .estimate = layoutPassWork,
252         },
253     };
254 }
255 
256 fn runLayoutPlanningPass(pass_ctx: *passes.PassContext) passes.PassResult {
257     _ = getLayoutPlanAnalysis(pass_ctx, pass_ctx.op) catch return .failure;
258     pass_ctx.preserveAllAnalyses();
259     return .success;
260 }
261 
262 fn computeLayoutPlanAnalysis(
263     pass_ctx: *passes.PassContext,
264     op: *ir.Operation,
265 ) anyerror!*anyopaque {
266     const buffers = try bufferization.getBufferPlanAnalysis(pass_ctx, op);
267     const memory_plan = try memory_space.getMemorySpacePlanAnalysis(pass_ctx, op);
268 
269     const analysis = try pass_ctx.allocator.create(LayoutPlanAnalysis);
270     analysis.* = LayoutPlanAnalysis.init(pass_ctx.allocator);
271     errdefer {
272         analysis.deinit();
273         pass_ctx.allocator.destroy(analysis);
274     }
275 
276     analysis.elided_value_count = buffers.elisionCount();
277     for (buffers.slots.items) |slot| {
278         const memory_assignment = memory_plan.getAssignmentForSlot(slot.id) orelse {
279             return error.MissingMemorySpaceAssignment;
280         };
281         {
282             var assignment = try LayoutAssignment.init(
283                 pass_ctx.allocator,
284                 slot,
285                 memory_assignment.*,
286             );
287             errdefer assignment.deinit(pass_ctx.allocator);
288             try analysis.addAssignment(assignment);
289         }
290     }
291     std.debug.assert(analysis.assignmentCount() == buffers.slotCount());
292     std.debug.assert(analysis.total_static_bytes == buffers.total_static_bytes);
293 
294     return @ptrCast(analysis);
295 }
296 
297 fn cleanupLayoutPlanAnalysis(ptr: *anyopaque, allocator: std.mem.Allocator) void {
298     const analysis: *LayoutPlanAnalysis = @ptrCast(@alignCast(ptr));
299     analysis.deinit();
300     allocator.destroy(analysis);
301 }
302 
303 fn layoutKindForSlot(slot: bufferization.BufferSlot) LayoutKind {
304     if (slot.dims.len == 0) return .scalar;
305     if (slot.element_count != null and slot.row_major_strides != null) return .row_major;
306     return .dynamic_row_major;
307 }
308 
309 fn rowMajorMinorToMajorAlloc(allocator: std.mem.Allocator, rank: usize) ![]usize {
310     const order = try allocator.alloc(usize, rank);
311     for (order, 0..) |*axis, index| {
312         axis.* = rank - 1 - index;
313     }
314     return order;
315 }
316 
317 const testing = std.testing;
318 const semantic = accy_choir.semantic;
319 
320 fn findOpNamedInBlock(block: *ir.Block, name: []const u8) ?*ir.Operation {
321     var iter = block.operations.head;
322     while (iter) |op_ptr| {
323         const op: *ir.Operation = @ptrCast(@alignCast(op_ptr));
324         if (isName(op.name.name, name)) return op;
325         iter = op.next_op;
326     }
327     return null;
328 }
329 
330 fn isName(actual: []const u8, expected: []const u8) bool {
331     return std.mem.eql(u8, actual, expected);
332 }
333 
334 test "layout planning records row-major slot layouts" {
335     const allocator = testing.allocator;
336 
337     var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
338     defer builder.deinit();
339     const f32_2x3 = try builder.tensor(.f32, &.{ 2, 3 });
340     var fb = try builder.beginFunction("layout_add2x3", &.{ f32_2x3, f32_2x3 }, &.{f32_2x3});
341     const sum = try fb.add(fb.parameter(0), fb.parameter(1));
342     try fb.return_(&.{sum});
343     try fb.finish();
344     const module = try builder.finish();
345     defer module.deinit();
346 
347     const choir_mod = module.choir_module;
348     const ctx = module.context();
349     const ledger = try layoutTestLedger();
350     defer ledger.destroy();
351     var cache = try passes.AnalysisCache.initAccounted(allocator, null, ledger, .{}, 6);
352     defer cache.deinit();
353     var pass_ctx = passes.PassContext.init(choir_mod, ctx, allocator, &cache);
354     defer pass_ctx.deinit();
355 
356     const buffers = try bufferization.getBufferPlanAnalysis(&pass_ctx, choir_mod);
357     const analysis = try getLayoutPlanAnalysis(&pass_ctx, choir_mod);
358     try ledger.producersComplete();
359     try checkLayoutStorage(choir_mod, ctx, &cache, analysis);
360     try testing.expectEqual(@as(usize, 3), analysis.assignmentCount());
361     try testing.expectEqual(@as(usize, 3), analysis.row_major_layout_count);
362     try testing.expectEqual(@as(usize, 3), analysis.device_global_slot_count);
363     try testing.expectEqual(@as(usize, 0), analysis.dynamic_slot_count);
364     try testing.expectEqual(@as(u64, 72), analysis.total_static_bytes);
365 
366     const body = choir_mod.getRegion(0).?.getEntryBlock().?;
367     const func = ir.inspection.functionByNameInBlock(body, "layout_add2x3") orelse return error.TestExpectedFunc;
368     const entry = func.getRegion(0).?.getEntryBlock().?;
369     const add = findOpNamedInBlock(entry, dialect_mod.AccyDialect.AddOp.operation_name) orelse return error.TestExpectedAdd;
370     const output = analysis.getAssignmentForValue(buffers, add.getResult(0).?) orelse return error.TestExpectedAssignment;
371     try testing.expectEqual(LayoutKind.row_major, output.kind);
372     try testing.expectEqual(memory_space.MemorySpace.device_global, output.memory_space);
373     try testing.expectEqual(@as(usize, 2), output.rank);
374     try testing.expectEqualSlices(i64, &.{ 2, 3 }, output.dims);
375     try testing.expectEqualSlices(u64, &.{ 3, 1 }, output.element_strides.?);
376     try testing.expectEqualSlices(usize, &.{ 1, 0 }, output.minor_to_major);
377     try testing.expectEqual(@as(?u64, 6), output.element_count);
378     try testing.expectEqual(@as(?u64, 24), output.byte_size);
379     try testing.expectEqual(@as(u64, 4), output.element_size);
380     try testing.expectEqual(@as(u64, 4), output.alignment);
381     try testing.expect(output.contiguous);
382     try testing.expect(output.static_layout);
383 }
384 
385 test "layout planning records constant memory layouts" {
386     const allocator = testing.allocator;
387 
388     var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
389     defer builder.deinit();
390     const i32_4 = try builder.tensor(.i32, &.{4});
391     var fb = try builder.beginFunction("layout_const_add", &.{i32_4}, &.{i32_4});
392     const values = [_]i32{ 1, 2, 3, 4 };
393     const c = try fb.constant(i32_4, std.mem.sliceAsBytes(values[0..]));
394     const sum = try fb.add(fb.parameter(0), c);
395     try fb.return_(&.{sum});
396     try fb.finish();
397     const module = try builder.finish();
398     defer module.deinit();
399 
400     const choir_mod = module.choir_module;
401     const ctx = module.context();
402     const ledger = try layoutTestLedger();
403     defer ledger.destroy();
404     var cache = try passes.AnalysisCache.initAccounted(allocator, null, ledger, .{}, 6);
405     defer cache.deinit();
406     var pass_ctx = passes.PassContext.init(choir_mod, ctx, allocator, &cache);
407     defer pass_ctx.deinit();
408 
409     const buffers = try bufferization.getBufferPlanAnalysis(&pass_ctx, choir_mod);
410     const analysis = try getLayoutPlanAnalysis(&pass_ctx, choir_mod);
411     try ledger.producersComplete();
412     try checkLayoutStorage(choir_mod, ctx, &cache, analysis);
413     try testing.expectEqual(@as(usize, 3), analysis.assignmentCount());
414     try testing.expectEqual(@as(usize, 2), analysis.device_global_slot_count);
415     try testing.expectEqual(@as(usize, 1), analysis.device_constant_slot_count);
416     try testing.expectEqual(@as(u64, 48), analysis.total_static_bytes);
417 
418     const body = choir_mod.getRegion(0).?.getEntryBlock().?;
419     const func = ir.inspection.functionByNameInBlock(body, "layout_const_add") orelse return error.TestExpectedFunc;
420     const entry = func.getRegion(0).?.getEntryBlock().?;
421     const constant = findOpNamedInBlock(entry, dialect_mod.AccyDialect.ConstantOp.operation_name) orelse return error.TestExpectedConstant;
422     const assignment = analysis.getAssignmentForValue(buffers, constant.getResult(0).?) orelse return error.TestExpectedAssignment;
423     try testing.expectEqual(LayoutKind.row_major, assignment.kind);
424     try testing.expectEqual(memory_space.MemorySpace.device_constant, assignment.memory_space);
425     try testing.expectEqualSlices(i64, &.{4}, assignment.dims);
426     try testing.expectEqualSlices(u64, &.{1}, assignment.element_strides.?);
427     try testing.expectEqualSlices(usize, &.{0}, assignment.minor_to_major);
428     try testing.expectEqual(@as(?u64, 16), assignment.byte_size);
429 }
430 
431 test "layout planning records fusion elisions without layouts" {
432     const allocator = testing.allocator;
433 
434     var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
435     defer builder.deinit();
436     const f32_4 = try builder.tensor(.f32, &.{4});
437     var fb = try builder.beginFunction("layout_fused_add_mul", &.{ f32_4, f32_4, f32_4 }, &.{f32_4});
438     const sum = try fb.add(fb.parameter(0), fb.parameter(1));
439     const product = try fb.mul(sum, fb.parameter(2));
440     try fb.return_(&.{product});
441     try fb.finish();
442     const module = try builder.finish();
443     defer module.deinit();
444 
445     const choir_mod = module.choir_module;
446     const ctx = module.context();
447     const ledger = try layoutTestLedger();
448     defer ledger.destroy();
449     var cache = try passes.AnalysisCache.initAccounted(allocator, null, ledger, .{}, 6);
450     defer cache.deinit();
451     var pass_ctx = passes.PassContext.init(choir_mod, ctx, allocator, &cache);
452     defer pass_ctx.deinit();
453 
454     const buffers = try bufferization.getBufferPlanAnalysis(&pass_ctx, choir_mod);
455     const analysis = try getLayoutPlanAnalysis(&pass_ctx, choir_mod);
456     try ledger.producersComplete();
457     try checkLayoutStorage(choir_mod, ctx, &cache, analysis);
458     try testing.expectEqual(@as(usize, 4), analysis.assignmentCount());
459     try testing.expectEqual(@as(usize, 1), analysis.elided_value_count);
460     try testing.expectEqual(@as(usize, 4), analysis.device_global_slot_count);
461 
462     const body = choir_mod.getRegion(0).?.getEntryBlock().?;
463     const func = ir.inspection.functionByNameInBlock(body, "layout_fused_add_mul") orelse return error.TestExpectedFunc;
464     const entry = func.getRegion(0).?.getEntryBlock().?;
465     const add = findOpNamedInBlock(entry, dialect_mod.AccyDialect.AddOp.operation_name) orelse return error.TestExpectedAdd;
466     try testing.expect(buffers.getSlot(add.getResult(0).?) == null);
467     try testing.expect(analysis.getAssignmentForValue(buffers, add.getResult(0).?) == null);
468 }
469 
470 test "layout planning pass preserves IR" {
471     const allocator = testing.allocator;
472 
473     var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
474     defer builder.deinit();
475     const f32_4 = try builder.tensor(.f32, &.{4});
476     var fb = try builder.beginFunction("layout_pass_add4", &.{ f32_4, f32_4 }, &.{f32_4});
477     const sum = try fb.add(fb.parameter(0), fb.parameter(1));
478     try fb.return_(&.{sum});
479     try fb.finish();
480     const module = try builder.finish();
481     defer module.deinit();
482 
483     const choir_mod = module.choir_module;
484     const ctx = module.context();
485     var pm = passes.PassManager.init(allocator);
486     defer pm.deinit();
487     try pm.addPass(layoutPlanningPass());
488 
489     const ledger = try choir.product.revision.AccountingV1.create(allocator, .{
490         .allowance = choir.product.revision.WorkVector.uniform(std.math.maxInt(u64)),
491         .workspace = std.math.maxInt(u64),
492         .events = 14,
493     }, &.{.{ .name = layout_planning_pass_name, .version = 1 }});
494     defer ledger.destroy();
495     var cache = try passes.AnalysisCache.initAccounted(allocator, null, ledger, .{}, 6);
496     defer cache.deinit();
497     try testing.expectEqual(
498         passes.PassResult.success,
499         pm.runWithAnalysisCache(choir_mod, ctx, &cache, .{}),
500     );
501     try ledger.producersComplete();
502     try testing.expectEqual(@as(u64, 1), pm.stats.pass_runs);
503     try testing.expectEqual(@as(u64, 0), pm.stats.passes_modified);
504 }
505 
506 fn layoutTestLedger() !*choir.product.revision.AccountingV1 {
507     const revision = choir.product.revision;
508     return revision.AccountingV1.create(testing.allocator, .{
509         .allowance = revision.WorkVector.uniform(std.math.maxInt(u64)),
510         .workspace = std.math.maxInt(u64),
511         .events = 14,
512     }, &.{});
513 }
514 
515 fn checkLayoutStorage(
516     op: *ir.Operation,
517     ctx: *ir.Context,
518     cache: *passes.AnalysisCache,
519     expected: *const LayoutPlanAnalysis,
520 ) !void {
521     const bounds = try layoutAnalysisWork(.{ .operation = op });
522     const fixed = @import("alloc_fixed");
523     const bytes = try testing.allocator.alignedAlloc(u8, .@"64", @intCast(bounds.workspace));
524     defer testing.allocator.free(bytes);
525     var storage = fixed.Tracked.init(bytes);
526     var retained = fixed.Monotonic.init(storage.allocator(), @max(1, bytes.len));
527     const allocator = retained.allocator();
528     var pass_ctx = passes.PassContext.init(op, ctx, allocator, cache);
529     defer pass_ctx.deinit();
530     const ptr = try computeLayoutPlanAnalysis(&pass_ctx, op);
531     defer cleanupLayoutPlanAnalysis(ptr, allocator);
532     const actual: *LayoutPlanAnalysis = @ptrCast(@alignCast(ptr));
533     inline for (.{
534         "scalar_layout_count",      "row_major_layout_count",   "dynamic_row_major_layout_count",
535         "host_slot_count",          "device_global_slot_count", "device_constant_slot_count",
536         "device_shared_slot_count", "unified_slot_count",       "dynamic_slot_count",
537         "elided_value_count",       "total_static_bytes",
538     }) |field| try testing.expectEqual(@field(expected, field), @field(actual, field));
539     try testing.expectEqual(expected.assignmentCount(), actual.assignmentCount());
540     for (expected.assignments.items, actual.assignments.items) |left, right| {
541         inline for (.{
542             "slot_id",       "value",        "producer",  "dtype",
543             "memory_space",  "kind",         "rank",      "element_count",
544             "byte_size",     "element_size", "alignment", "contiguous",
545             "static_layout",
546         }) |field| try testing.expectEqual(@field(left, field), @field(right, field));
547         try testing.expectEqualDeep(left.role, right.role);
548         try testing.expectEqualSlices(i64, left.dims, right.dims);
549         try testing.expectEqualSlices(usize, left.minor_to_major, right.minor_to_major);
550         if (left.element_strides) |strides| {
551             try testing.expectEqualSlices(u64, strides, right.element_strides.?);
552         } else try testing.expect(right.element_strides == null);
553         try testing.expectEqual(right.value, actual.getAssignmentForSlot(right.slot_id).?.value);
554     }
555     try testing.expectEqual(expected.slot_to_assignment.count(), actual.slot_to_assignment.count());
556     try testing.expect(!storage.exhausted);
557     const used = if (retained.current) |*current| fixed.used(current) else 0;
558     try testing.expect(used <= bounds.workspace);
559     try testing.expect(used >= @sizeOf(LayoutPlanAnalysis));
560 }
561 
562 test "layout planning work contract covers slot growth and layout dimensions" {
563     const high_rank: [512]i64 = @splat(1);
564     for ([_]usize{ 0, 1, 6, 7, 16, 64 }) |count| {
565         try checkLayoutBoundary(count, &.{}, .scalar, false);
566         try checkLayoutBoundary(count, &.{ 2, 3 }, .row_major, false);
567         try checkLayoutBoundary(count, &high_rank, .row_major, false);
568         try checkLayoutBoundary(count, &.{ -1, 3 }, .dynamic_row_major, true);
569         try checkLayoutBoundary(count, &.{ std.math.maxInt(i64), 3 }, .dynamic_row_major, true);
570         try checkLayoutBoundary(count, &.{std.math.maxInt(i64)}, .row_major, true);
571     }
572     try testing.expectError(error.WorkOverflow, (LayoutWork{
573         .input = .{ .values = std.math.maxInt(u64) },
574     }).bounds());
575     try testing.expectError(error.WorkOverflow, (LayoutWork{
576         .input = .{ .input_bytes = std.math.maxInt(u64) },
577     }).bounds());
578     try testing.expectError(error.WorkOverflow, (LayoutWork{
579         .input = .{ .atoms = std.math.maxInt(u64) },
580     }).bounds());
581 }
582 
583 fn checkLayoutBoundary(
584     count: usize,
585     dims: []const i64,
586     kind: LayoutKind,
587     unknown_bytes: bool,
588 ) !void {
589     std.debug.assert(count <= 64);
590     var context_limits = semantic.Builder.ContextLimits.standard;
591     context_limits.transient_bytes = 2 * 1024 * 1024;
592     var builder = try semantic.Builder.init(testing.allocator, context_limits);
593     defer builder.deinit();
594     const typ = try builder.tensor(.f32, dims);
595     var types: [64]@TypeOf(typ) = @splat(typ);
596     var function = try builder.beginFunction("layout_boundary", types[0..count], types[0..count]);
597     var values: [64]@TypeOf(function.parameter(0)) = undefined;
598     for (values[0..count], 0..) |*value, index| value.* = function.parameter(index);
599     try function.return_(values[0..count]);
600     try function.finish();
601     const module = try builder.finish();
602     defer module.deinit();
603     const ledger = try layoutTestLedger();
604     defer ledger.destroy();
605     var cache = try passes.AnalysisCache.initAccounted(testing.allocator, null, ledger, .{}, 6);
606     defer cache.deinit();
607     var pass_ctx = passes.PassContext.init(
608         module.choir_module,
609         module.context(),
610         testing.allocator,
611         &cache,
612     );
613     defer pass_ctx.deinit();
614     const analysis = try getLayoutPlanAnalysis(&pass_ctx, module.choir_module);
615     try ledger.producersComplete();
616     try testing.expectEqual(count, analysis.assignmentCount());
617     try testing.expectEqual(if (unknown_bytes) count else 0, analysis.dynamic_slot_count);
618     for (analysis.assignments.items, 0..) |assignment, index| {
619         try testing.expectEqual(index, assignment.slot_id);
620         try testing.expectEqual(values[index], assignment.value);
621         try testing.expectEqual(kind, assignment.kind);
622         try testing.expectEqual(dims.len, assignment.rank);
623         try testing.expectEqualSlices(i64, dims, assignment.dims);
624         try testing.expectEqual(dims.len, assignment.minor_to_major.len);
625         for (assignment.minor_to_major, 0..) |axis, position| {
626             try testing.expectEqual(dims.len - position - 1, axis);
627         }
628     }
629     try checkLayoutStorage(module.choir_module, module.context(), &cache, analysis);
630     try checkLayoutAdmission(module.choir_module, module.context());
631 }
632 
633 fn checkLayoutAdmission(op: *ir.Operation, ctx: *ir.Context) !void {
634     const revision = choir.product.revision;
635     const preparation = @import("root.zig");
636     var charge: u64 = 1;
637     for ([_]passes.AnalysisDescriptor{
638         preparation.shape.shape_layout_analysis_descriptor,
639         preparation.fusion.fusion_plan_analysis_descriptor,
640         preparation.schedule.schedule_plan_analysis_descriptor,
641         bufferization.buffer_plan_analysis_descriptor,
642         memory_space.memory_space_plan_analysis_descriptor,
643         layout_plan_analysis_descriptor,
644     }) |descriptor| {
645         const bounds = try descriptor.work_contract.?.estimate(.{ .operation = op });
646         charge = try accounting.add(charge, bounds.work.structural_visits);
647     }
648     for ([_]i8{ -1, 0, 1 }) |offset| {
649         var allowance = revision.WorkVector.uniform(std.math.maxInt(u64));
650         allowance.structural_visits = @intCast(@as(i128, charge) + offset);
651         const ledger = try revision.AccountingV1.create(testing.allocator, .{
652             .allowance = allowance,
653             .workspace = std.math.maxInt(u64),
654             .events = 14,
655         }, &.{.{ .name = layout_planning_pass_name, .version = 1 }});
656         defer ledger.destroy();
657         var cache = try passes.AnalysisCache.initAccounted(
658             testing.allocator,
659             null,
660             ledger,
661             .{},
662             6,
663         );
664         defer cache.deinit();
665         var manager = passes.PassManager.init(testing.allocator);
666         defer manager.deinit();
667         try manager.addPass(layoutPlanningPass());
668         const result = manager.runWithAnalysisCache(op, ctx, &cache, .{});
669         if (offset < 0) {
670             try testing.expectEqual(passes.PassResult.failure, result);
671             try testing.expectEqual(revision.receipt.Outcome.exhausted, ledger.view().outcome);
672             try testing.expectEqual(@as(usize, 3), cache.entries.count());
673         } else {
674             try testing.expectEqual(passes.PassResult.success, result);
675             try ledger.producersComplete();
676             try testing.expectEqual(@as(usize, 6), cache.entries.count());
677         }
678     }
679 }