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 };