lib/accy/src/choir/shape/test.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const tensor = @import("tensor.zig");
3 const namespace = @import("root.zig");
4 const SymbolId = namespace.SymbolId;
5 const Term = namespace.Term;
6 const Expression = namespace.Expression;
7 const Bounds = namespace.Bounds;
8 const Predicate = namespace.Predicate;
9 const Fact = namespace.Fact;
10 const FactMode = namespace.FactMode;
11 const Tensor = namespace.Tensor;
12 const Symbol = namespace.Symbol;
13 const Family = namespace.Family;
14 const Builder = namespace.Builder;
15 const Error = namespace.Error;
16 const constant = namespace.constant;
17 const fingerprint = namespace.fingerprint;
18
19 test {
20 @import("test_discovery").discover(namespace);
21 }
22
23 test "shape family records symbolic matmul constraints" {
24 var builder = try Builder.init(std.testing.allocator, "matmul");
25 errdefer builder.deinit();
26
27 const m = try builder.symbol("M");
28 const k = try builder.symbol("K");
29 const n = try builder.symbol("N");
30
31 const m_expr = try builder.symbolExpression(m);
32 const k_expr = try builder.symbolExpression(k);
33 const n_expr = try builder.symbolExpression(n);
34
35 _ = try builder.tensor("lhs", &.{ m_expr, k_expr });
36 _ = try builder.tensor("rhs", &.{ k_expr, n_expr });
37 _ = try builder.tensor("out", &.{ m_expr, n_expr });
38 try builder.assumeBounds(m_expr, .{ .min = 1, .opt = 128, .max = 4096 });
39 try builder.assumeBounds(k_expr, .{ .min = 1, .opt = 128, .max = 4096 });
40 try builder.assumeBounds(n_expr, .{ .min = 1, .opt = 128, .max = 4096 });
41 try builder.assumeDivisible(k_expr, 16);
42 try builder.assertBounds(m_expr, .{ .min = 1 });
43
44 var value = builder.finish();
45 defer value.deinit();
46
47 try std.testing.expectEqual(@as(usize, 3), value.symbols.len);
48 try std.testing.expectEqual(@as(usize, 3), value.tensors.len);
49 try std.testing.expectEqual(@as(usize, 4), value.assumedFactCount());
50 try std.testing.expectEqual(@as(usize, 1), value.assertedFactCount());
51 try std.testing.expectEqual(k, value.symbolIndex("K").?);
52
53 const out = value.tensor("out").?;
54 try std.testing.expectEqual(@as(usize, 2), out.rank());
55 try std.testing.expect(out.extents[0].eql(m_expr));
56 try std.testing.expect(out.extents[1].eql(n_expr));
57 }
58
59 test "shape family expressions compose affine extent values" {
60 var builder = try Builder.init(std.testing.allocator, "flatten");
61 errdefer builder.deinit();
62
63 const n = try builder.symbol("N");
64 const n_expr = try builder.symbolExpression(n);
65 const four_n = try builder.scaledSymbolExpression(n, 4, 0);
66 const four_n_plus_one = try builder.addExpression(four_n, builder.constantExpression(1));
67
68 _ = try builder.tensor("input", &.{ n_expr, builder.constantExpression(4) });
69 _ = try builder.tensor("flat", &.{four_n});
70 try builder.assertEqual(four_n_plus_one, try builder.addExpression(four_n, builder.constantExpression(1)));
71
72 var value = builder.finish();
73 defer value.deinit();
74
75 const flat = value.tensor("flat").?;
76 try std.testing.expectEqual(@as(i64, 4), flat.extents[0].terms[0].coefficient);
77 try std.testing.expectEqual(@as(usize, 1), value.assertedFactCount());
78 }
79
80 test "shape family validates names bounds symbols and divisors" {
81 try std.testing.expectError(Error.EmptyName, Builder.init(std.testing.allocator, ""));
82
83 var builder = try Builder.init(std.testing.allocator, "invalid");
84 defer builder.deinit();
85
86 const n = try builder.symbol("N");
87 try std.testing.expectError(Error.DuplicateSymbol, builder.symbol("N"));
88
89 const n_expr = try builder.symbolExpression(n);
90 try std.testing.expectError(Error.InvalidBounds, builder.assumeBounds(n_expr, .{}));
91 try std.testing.expectError(Error.InvalidBounds, builder.assumeBounds(n_expr, .{ .min = 8, .opt = 4 }));
92 try std.testing.expectError(Error.InvalidDivisor, builder.assumeDivisible(n_expr, 0));
93 try std.testing.expectError(Error.UnknownSymbol, builder.symbolExpression(999));
94 try std.testing.expectError(Error.UnknownSymbol, builder.addExpression(n_expr, .{ .terms = &.{.{ .symbol = 999 }} }));
95 try std.testing.expectError(Error.InvalidExtent, builder.tensor("bad", &.{builder.constantExpression(-1)}));
96 _ = try builder.tensor("valid", &.{builder.constantExpression(1)});
97 try std.testing.expectError(Error.DuplicateTensor, builder.tensor("valid", &.{builder.constantExpression(1)}));
98 }
99
100 test "shape family owns expression terms at storage boundaries" {
101 var builder = try Builder.init(std.testing.allocator, "ownership");
102 errdefer builder.deinit();
103
104 const n = try builder.symbol("N");
105 var borrowed_terms = [_]Term{.{ .symbol = n, .coefficient = 2 }};
106 const borrowed = Expression{ .constant = 3, .terms = &borrowed_terms };
107 _ = try builder.tensor("scaled", &.{borrowed});
108 try builder.assertEqual(borrowed, builder.constantExpression(7));
109
110 var value = builder.finish();
111 defer value.deinit();
112
113 borrowed_terms[0].coefficient = 99;
114
115 const scaled = value.tensor("scaled").?;
116 try std.testing.expectEqual(@as(i64, 2), scaled.extents[0].terms[0].coefficient);
117 switch (value.facts[0].predicate) {
118 .equal => |equal| try std.testing.expectEqual(@as(i64, 2), equal.lhs.terms[0].coefficient),
119 else => return error.UnexpectedPredicate,
120 }
121 }
122
123 test "shape family fingerprint is stable and semantic" {
124 var first_builder = try Builder.init(std.testing.allocator, "matmul");
125 errdefer first_builder.deinit();
126 const first_m = try first_builder.symbol("M");
127 const first_n = try first_builder.symbol("N");
128 const first_m_expr = try first_builder.symbolExpression(first_m);
129 const first_n_expr = try first_builder.symbolExpression(first_n);
130 _ = try first_builder.tensor("out", &.{ first_m_expr, first_n_expr });
131 try first_builder.assumeBounds(first_n_expr, .{ .min = 1, .max = 4096 });
132 var first = first_builder.finish();
133 defer first.deinit();
134
135 var same_builder = try Builder.init(std.testing.allocator, "matmul");
136 errdefer same_builder.deinit();
137 const same_m = try same_builder.symbol("M");
138 const same_n = try same_builder.symbol("N");
139 const same_m_expr = try same_builder.symbolExpression(same_m);
140 const same_n_expr = try same_builder.symbolExpression(same_n);
141 _ = try same_builder.tensor("out", &.{ same_m_expr, same_n_expr });
142 try same_builder.assumeBounds(same_n_expr, .{ .min = 1, .max = 4096 });
143 var same = same_builder.finish();
144 defer same.deinit();
145
146 var changed_builder = try Builder.init(std.testing.allocator, "matmul");
147 errdefer changed_builder.deinit();
148 const changed_m = try changed_builder.symbol("M");
149 const changed_n = try changed_builder.symbol("N");
150 const changed_m_expr = try changed_builder.symbolExpression(changed_m);
151 const changed_n_expr = try changed_builder.symbolExpression(changed_n);
152 _ = try changed_builder.tensor("out", &.{ changed_m_expr, changed_n_expr });
153 try changed_builder.assumeBounds(changed_n_expr, .{ .min = 1, .max = 2048 });
154 var changed = changed_builder.finish();
155 defer changed.deinit();
156
157 try std.testing.expectEqual(fingerprint(first), fingerprint(same));
158 try std.testing.expect(fingerprint(first) != fingerprint(changed));
159 }
160
161 test "accy choir shape declaration coverage" {
162 std.testing.refAllDecls(namespace);
163 }