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 }