lib/accy/src/choir/shape/family.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const alloc_arena = @import("alloc_arena");
3 const expression = @import("expression.zig");
4 const fact_mod = @import("fact.zig");
5 const tensor_mod = @import("tensor.zig");
6
7 pub const Error = error{
8 DuplicateSymbol,
9 DuplicateTensor,
10 EmptyName,
11 InvalidBounds,
12 InvalidDivisor,
13 InvalidExtent,
14 UnknownSymbol,
15 };
16
17 pub const Symbol = struct {
18 name: []const u8,
19 };
20
21 pub const Family = struct {
22 backing_allocator: std.mem.Allocator,
23 arena_state: *alloc_arena.Arena,
24 name: []const u8,
25 symbols: []const Symbol,
26 tensors: []const tensor_mod.Tensor,
27 facts: []const fact_mod.Fact,
28
29 pub fn deinit(self: *Family) void {
30 self.arena_state.deinit();
31 self.backing_allocator.destroy(self.arena_state);
32 self.* = undefined;
33 }
34
35 pub fn symbolIndex(self: Family, name: []const u8) ?expression.SymbolId {
36 for (self.symbols, 0..) |item, index| {
37 if (std.mem.eql(u8, item.name, name)) return @intCast(index);
38 }
39 return null;
40 }
41
42 pub fn tensor(self: Family, name: []const u8) ?tensor_mod.Tensor {
43 for (self.tensors) |item| {
44 if (std.mem.eql(u8, item.name, name)) return item;
45 }
46 return null;
47 }
48
49 pub fn assumedFactCount(self: Family) usize {
50 var count: usize = 0;
51 for (self.facts) |item| {
52 if (item.mode == .assume) count += 1;
53 }
54 return count;
55 }
56
57 pub fn assertedFactCount(self: Family) usize {
58 var count: usize = 0;
59 for (self.facts) |item| {
60 if (item.mode == .assert) count += 1;
61 }
62 return count;
63 }
64 };
65
66 pub const Builder = struct {
67 backing_allocator: std.mem.Allocator,
68 arena_state: *alloc_arena.Arena,
69 name: []const u8,
70 symbols: std.ArrayListUnmanaged(Symbol) = .empty,
71 tensors: std.ArrayListUnmanaged(tensor_mod.Tensor) = .empty,
72 facts: std.ArrayListUnmanaged(fact_mod.Fact) = .empty,
73
74 pub fn init(backing_allocator: std.mem.Allocator, name: []const u8) !Builder {
75 if (name.len == 0) return Error.EmptyName;
76 const arena_state = try backing_allocator.create(alloc_arena.Arena);
77 errdefer backing_allocator.destroy(arena_state);
78 arena_state.* = alloc_arena.Arena.init(backing_allocator);
79 errdefer arena_state.deinit();
80 const alloc = arena_state.allocator();
81 return .{
82 .backing_allocator = backing_allocator,
83 .arena_state = arena_state,
84 .name = try alloc.dupe(u8, name),
85 };
86 }
87
88 pub fn deinit(self: *Builder) void {
89 self.arena_state.deinit();
90 self.backing_allocator.destroy(self.arena_state);
91 self.* = undefined;
92 }
93
94 pub fn symbol(self: *Builder, name: []const u8) !expression.SymbolId {
95 if (name.len == 0) return Error.EmptyName;
96 if (self.findSymbol(name) != null) return Error.DuplicateSymbol;
97 const alloc = self.arena();
98 const id: expression.SymbolId = @intCast(self.symbols.items.len);
99 try self.symbols.append(alloc, .{ .name = try alloc.dupe(u8, name) });
100 return id;
101 }
102
103 pub fn symbolExpression(self: *Builder, id: expression.SymbolId) !expression.Expression {
104 try self.requireSymbol(id);
105 const terms = try self.arena().alloc(expression.Term, 1);
106 terms[0] = .{ .symbol = id };
107 return .{ .terms = terms };
108 }
109
110 pub fn scaledSymbolExpression(self: *Builder, id: expression.SymbolId, coefficient: i64, constant: i64) !expression.Expression {
111 try self.requireSymbol(id);
112 const terms = try self.arena().alloc(expression.Term, 1);
113 terms[0] = .{ .symbol = id, .coefficient = coefficient };
114 return .{ .constant = constant, .terms = terms };
115 }
116
117 pub fn constantExpression(_: *Builder, value: i64) expression.Expression {
118 return expression.constant(value);
119 }
120
121 pub fn addExpression(self: *Builder, lhs: expression.Expression, rhs: expression.Expression) !expression.Expression {
122 try self.requireExpression(lhs);
123 try self.requireExpression(rhs);
124 const terms = try self.arena().alloc(expression.Term, lhs.terms.len + rhs.terms.len);
125 @memcpy(terms[0..lhs.terms.len], lhs.terms);
126 @memcpy(terms[lhs.terms.len..], rhs.terms);
127 return .{ .constant = lhs.constant + rhs.constant, .terms = terms };
128 }
129
130 pub fn tensor(self: *Builder, name: []const u8, extents: []const expression.Expression) !usize {
131 if (name.len == 0) return Error.EmptyName;
132 if (self.findTensor(name) != null) return Error.DuplicateTensor;
133 const alloc = self.arena();
134 const id = self.tensors.items.len;
135 const owned_extents = try alloc.alloc(expression.Expression, extents.len);
136 for (extents, owned_extents) |extent, *owned_extent| owned_extent.* = try self.ownTensorExtent(extent);
137 try self.tensors.append(alloc, .{
138 .name = try alloc.dupe(u8, name),
139 .extents = owned_extents,
140 });
141 return id;
142 }
143
144 pub fn assumeEqual(self: *Builder, lhs: expression.Expression, rhs: expression.Expression) !void {
145 try self.appendFact(.assume, .{ .equal = .{ .lhs = lhs, .rhs = rhs } });
146 }
147
148 pub fn assertEqual(self: *Builder, lhs: expression.Expression, rhs: expression.Expression) !void {
149 try self.appendFact(.assert, .{ .equal = .{ .lhs = lhs, .rhs = rhs } });
150 }
151
152 pub fn assumeBounds(self: *Builder, value: expression.Expression, bounds: fact_mod.Bounds) !void {
153 try self.appendBounds(.assume, value, bounds);
154 }
155
156 pub fn assertBounds(self: *Builder, value: expression.Expression, bounds: fact_mod.Bounds) !void {
157 try self.appendBounds(.assert, value, bounds);
158 }
159
160 pub fn assumeDivisible(self: *Builder, value: expression.Expression, divisor: u64) !void {
161 try self.appendDivisible(.assume, value, divisor);
162 }
163
164 pub fn assertDivisible(self: *Builder, value: expression.Expression, divisor: u64) !void {
165 try self.appendDivisible(.assert, value, divisor);
166 }
167
168 pub fn finish(self: *Builder) Family {
169 const result = Family{
170 .backing_allocator = self.backing_allocator,
171 .arena_state = self.arena_state,
172 .name = self.name,
173 .symbols = self.symbols.items,
174 .tensors = self.tensors.items,
175 .facts = self.facts.items,
176 };
177 self.* = undefined;
178 return result;
179 }
180
181 fn appendBounds(self: *Builder, mode: fact_mod.Mode, value: expression.Expression, bounds: fact_mod.Bounds) !void {
182 if (!bounds.valid()) return Error.InvalidBounds;
183 try self.appendFact(mode, .{ .bound = .{ .value = value, .bounds = bounds } });
184 }
185
186 fn appendDivisible(self: *Builder, mode: fact_mod.Mode, value: expression.Expression, divisor: u64) !void {
187 if (divisor == 0) return Error.InvalidDivisor;
188 try self.appendFact(mode, .{ .divisible = .{ .value = value, .divisor = divisor } });
189 }
190
191 fn appendFact(self: *Builder, mode: fact_mod.Mode, predicate: fact_mod.Predicate) !void {
192 const owned_predicate = try self.ownPredicate(predicate);
193 try self.facts.append(self.arena(), .{ .mode = mode, .predicate = owned_predicate });
194 }
195
196 fn arena(self: *Builder) std.mem.Allocator {
197 return self.arena_state.allocator();
198 }
199
200 fn findSymbol(self: *const Builder, name: []const u8) ?expression.SymbolId {
201 for (self.symbols.items, 0..) |item, index| {
202 if (std.mem.eql(u8, item.name, name)) return @intCast(index);
203 }
204 return null;
205 }
206
207 fn findTensor(self: *const Builder, name: []const u8) ?usize {
208 for (self.tensors.items, 0..) |item, index| {
209 if (std.mem.eql(u8, item.name, name)) return index;
210 }
211 return null;
212 }
213
214 fn requireSymbol(self: *const Builder, id: expression.SymbolId) !void {
215 if (id >= self.symbols.items.len) return Error.UnknownSymbol;
216 }
217
218 fn requireExpression(self: *const Builder, value: expression.Expression) !void {
219 for (value.terms) |term| try self.requireSymbol(term.symbol);
220 }
221
222 fn ownTensorExtent(self: *Builder, value: expression.Expression) !expression.Expression {
223 const owned = try self.ownExpression(value);
224 if (owned.terms.len == 0 and owned.constant < 0) return Error.InvalidExtent;
225 return owned;
226 }
227
228 fn ownExpression(self: *Builder, value: expression.Expression) !expression.Expression {
229 try self.requireExpression(value);
230 if (value.terms.len == 0) return .{ .constant = value.constant };
231 return .{
232 .constant = value.constant,
233 .terms = try self.arena().dupe(expression.Term, value.terms),
234 };
235 }
236
237 fn ownPredicate(self: *Builder, predicate: fact_mod.Predicate) !fact_mod.Predicate {
238 switch (predicate) {
239 .equal => |value| return .{ .equal = .{
240 .lhs = try self.ownExpression(value.lhs),
241 .rhs = try self.ownExpression(value.rhs),
242 } },
243 .bound => |value| return .{ .bound = .{
244 .value = try self.ownExpression(value.value),
245 .bounds = value.bounds,
246 } },
247 .divisible => |value| return .{ .divisible = .{
248 .value = try self.ownExpression(value.value),
249 .divisor = value.divisor,
250 } },
251 }
252 }
253 };