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 }