lib/pluck/src/toplevel/lifecycle.zig

daab053ee43316e1809a84551d573ddd1e5bf3d2

  1 const std = @import("std");
  2 const alloc_arena = @import("alloc_arena");
  3 const pluck = @import("../root.zig");
  4 const log = pluck.logger;
  5 const Allocator = std.mem.Allocator;
  6 
  7 const pexpr = pluck.pexpr;
  8 const Definitions = pexpr.Definitions;
  9 const TypeRegistry = pexpr.TypeRegistry;
 10 
 11 const lpsmc_module = pluck.lpsmc;
 12 
 13 const bdd = pluck.bdd;
 14 const Manager = bdd.Manager;
 15 
 16 const context_owner = @import("context.zig");
 17 const ToplevelContext = context_owner.ToplevelContext;
 18 const top_types = @import("types.zig");
 19 const ToplevelError = top_types.ToplevelError;
 20 const config_owner = @import("config.zig");
 21 const ToplevelConfig = config_owner.ToplevelConfig;
 22 const ToplevelConfigPatch = config_owner.ToplevelConfigPatch;
 23 const ConfigError = config_owner.ConfigError;
 24 const shared_owner = @import("shared.zig");
 25 const source_owner = @import("source.zig");
 26 
 27 const stdlib_source = @embedFile("../stdlib.pluck");
 28 
 29 pub fn init(allocator: Allocator) !ToplevelContext {
 30     const arena_ptr = try allocator.create(alloc_arena.Arena);
 31     errdefer allocator.destroy(arena_ptr);
 32     arena_ptr.* = alloc_arena.Arena.init(allocator);
 33     errdefer arena_ptr.deinit();
 34 
 35     const query_arena_ptr = try allocator.create(alloc_arena.Arena);
 36     errdefer allocator.destroy(query_arena_ptr);
 37     query_arena_ptr.* = alloc_arena.Arena.init(allocator);
 38     errdefer query_arena_ptr.deinit();
 39 
 40     const arena_alloc = arena_ptr.allocator();
 41 
 42     const types = try TypeRegistry.initWithDefaults(arena_alloc);
 43 
 44     const definitions = Definitions.init(arena_alloc);
 45 
 46     const manager = try allocator.create(Manager);
 47     errdefer allocator.destroy(manager);
 48     manager.* = try Manager.init(arena_alloc);
 49     errdefer manager.deinit();
 50 
 51     var self = ToplevelContext{
 52         .allocator = allocator,
 53         .owns_shared = true,
 54         .owns_arena = true,
 55         .arena = arena_ptr,
 56         .query_arena = query_arena_ptr,
 57         .types = types,
 58         .definitions = definitions,
 59         .manager = manager,
 60         .config = .{},
 61     };
 62 
 63     try loadStdlib(&self);
 64 
 65     return self;
 66 }
 67 
 68 pub fn initWorker(self: *const ToplevelContext, allocator: Allocator) !ToplevelContext {
 69     const arena_ptr = try allocator.create(alloc_arena.Arena);
 70     errdefer allocator.destroy(arena_ptr);
 71     arena_ptr.* = alloc_arena.Arena.init(allocator);
 72     errdefer arena_ptr.deinit();
 73 
 74     const query_arena_ptr = try allocator.create(alloc_arena.Arena);
 75     errdefer allocator.destroy(query_arena_ptr);
 76     query_arena_ptr.* = alloc_arena.Arena.init(allocator);
 77     errdefer query_arena_ptr.deinit();
 78 
 79     const manager = try allocator.create(Manager);
 80     errdefer allocator.destroy(manager);
 81     manager.* = try Manager.init(arena_ptr.allocator());
 82     errdefer manager.deinit();
 83 
 84     const arena_alloc = arena_ptr.allocator();
 85 
 86     var cloned_defs = try shared_owner.cloneDefinitionsShallow(arena_alloc, &self.definitions);
 87     errdefer cloned_defs.defs.deinit();
 88 
 89     var cloned_types = try shared_owner.cloneTypesShallow(arena_alloc, &self.types);
 90     errdefer {
 91         cloned_types.type_of_constructor.deinit();
 92         cloned_types.constructors_of_type.deinit();
 93         cloned_types.args_of_constructor.deinit();
 94     }
 95 
 96     var worker = self.*;
 97     worker.allocator = allocator;
 98     worker.owns_shared = false;
 99     worker.owns_arena = true;
100     worker.arena = arena_ptr;
101     worker.query_arena = query_arena_ptr;
102     worker.manager = manager;
103     worker.definitions = cloned_defs;
104     worker.types = cloned_types;
105     worker.incremental_lpsmc = null;
106     worker.last_query_constructor_typo = null;
107     worker.current_source = null;
108     worker.loading_stdlib = false;
109     return worker;
110 }
111 
112 pub fn loadStdlib(self: *ToplevelContext) error{ OutOfMemory, StdlibLoadFailed }!void {
113     self.loading_stdlib = true;
114     defer self.loading_stdlib = false;
115 
116     source_owner.processSourceSexprNoWriter(self, stdlib_source) catch |err| {
117         if (@import("builtin").mode == .debug) {
118             log.err("stdlib load failed: {}", .{err});
119         }
120         return switch (err) {
121             ToplevelError.OutOfMemory => error.OutOfMemory,
122             else => error.StdlibLoadFailed,
123         };
124     };
125 }
126 
127 pub fn deinit(self: *ToplevelContext) void {
128     if (!self.owns_shared) {
129         self.definitions.defs.deinit();
130         self.types.type_of_constructor.deinit();
131         self.types.constructors_of_type.deinit();
132         self.types.args_of_constructor.deinit();
133     }
134 
135     if (self.incremental_lpsmc) |lpsmc| {
136         lpsmc_module.deinit(lpsmc);
137         self.allocator.destroy(lpsmc);
138     }
139 
140     self.manager.deinit();
141     self.allocator.destroy(self.manager);
142     if (self.owns_arena) {
143         self.arena.deinit();
144         self.allocator.destroy(self.arena);
145     }
146     self.query_arena.deinit();
147     self.allocator.destroy(self.query_arena);
148 }
149 
150 pub fn setConfig(self: *ToplevelContext, config: ToplevelConfig) void {
151     self.config = config;
152 }
153 
154 pub fn patchConfig(self: *ToplevelContext, patch: ToplevelConfigPatch) ConfigError!void {
155     try config_owner.applyPatch(&self.config, patch);
156 }
157 
158 const query_context = pluck.query_context;
159 pub const SharedContext = query_context.SharedContext;
160 pub const QueryContext = query_context.QueryContext;
161 pub const RunContext = query_context.RunContext;
162 pub const QueryConfig = query_context.QueryConfig;
163 pub const SessionConfig = query_context.SessionConfig;
164 
165 pub fn asSharedContext(self: *const ToplevelContext) SharedContext {
166     return SharedContext{
167         .allocator = self.allocator,
168         .arena = self.arena,
169         .types = &self.types,
170         .definitions = &self.definitions,
171         .config = SessionConfig{
172             .max_depth = self.config.max_depth,
173             .ite_limit = self.config.ite_limit,
174             .time_limit = self.config.time_limit,
175             .sample_after_max_depth = self.config.sample_after_max_depth,
176             .use_strict_order = self.config.use_strict_order,
177             .use_reverse_order = self.config.use_reverse_order,
178             .definition_order_mode = self.config.definition_order_mode,
179             .parallel_wmc = self.config.parallel_wmc,
180             .verbose = self.config.verbose,
181             .rng_seed = self.config.rng_seed,
182             .lpsmc_rng_seed = self.config.lpsmc_rng_seed,
183         },
184     };
185 }
186 
187 pub fn createQueryContext(self: *const ToplevelContext, query_config: QueryConfig) !QueryContext {
188     return try query_context.initQueryContext(asSharedContext(self), query_config);
189 }
190 
191 pub fn reset(self: *ToplevelContext) error{ OutOfMemory, StdlibLoadFailed, MacroAlreadyExists }!void {
192     if (self.incremental_lpsmc) |lpsmc| {
193         lpsmc_module.clearCaches(lpsmc);
194     }
195 
196     _ = self.arena.reset(.retain_capacity);
197     const alloc = self.arena.allocator();
198 
199     self.types = try TypeRegistry.initWithDefaults(alloc);
200     self.definitions = Definitions.init(alloc);
201 
202     self.manager.* = try Manager.init(alloc);
203 
204     self.current_source = null;
205     self.last_query_constructor_typo = null;
206 
207     try loadStdlib(self);
208 }
209 
210 pub fn resetManagerForQuery(self: *ToplevelContext) !void {
211     const manager_alloc = self.manager.allocator;
212     self.manager.deinit();
213     self.manager.* = try Manager.init(manager_alloc);
214 }
215 
216 test "query type registry matches Pluck.jl public constructors" {
217     const allocator = std.testing.allocator;
218     var ctx = try init(allocator);
219     defer ctx.deinit();
220 
221     try std.testing.expect(ctx.types.hasConstructor("Marginal"));
222     try std.testing.expect(ctx.types.hasConstructor("Posterior"));
223     try std.testing.expect(ctx.types.hasConstructor("PosteriorSamples"));
224     try std.testing.expect(ctx.types.hasConstructor("AdaptiveRejection"));
225     try std.testing.expect(!ctx.types.hasConstructor("SubproblemMonteCarlo"));
226 
227     try ctx.reset();
228     try std.testing.expect(ctx.types.hasConstructor("Marginal"));
229     try std.testing.expect(ctx.types.hasConstructor("Posterior"));
230     try std.testing.expect(ctx.types.hasConstructor("PosteriorSamples"));
231     try std.testing.expect(ctx.types.hasConstructor("AdaptiveRejection"));
232     try std.testing.expect(!ctx.types.hasConstructor("SubproblemMonteCarlo"));
233 }
234 
235 test "patchConfig updates context config" {
236     var ctx = try init(std.testing.allocator);
237     defer ctx.deinit();
238 
239     try ctx.patchConfig(.{
240         .time_limit = .{ .set = 0.5 },
241         .var_order = .reverse,
242         .fallback_lpsmc_k = 3,
243     });
244 
245     try std.testing.expectEqual(@as(?f64, 0.5), ctx.config.time_limit);
246     try std.testing.expect(ctx.config.use_strict_order);
247     try std.testing.expect(ctx.config.use_reverse_order);
248     try std.testing.expectEqual(@as(usize, 3), ctx.config.fallback_lpsmc_k);
249 }
250 
251 test "worker context validates definitions without mutating parent" {
252     const allocator = std.testing.allocator;
253     var ctx = try init(allocator);
254     defer ctx.deinit();
255 
256     try ctx.processSource("(define p 0.25)");
257 
258     var worker = try ctx.initWorker(allocator);
259     defer worker.deinit();
260     try worker.processSource(
261         \\(define-type result (Yes) (No))
262         \\(define p 0.75)
263     );
264 
265     try std.testing.expect(worker.types.hasConstructor("Yes"));
266     try std.testing.expect(!ctx.types.hasConstructor("Yes"));
267 
268     var parent_result = try ctx.processForm("(Marginal (flip p))") orelse return error.TestExpectedEqual;
269     defer parent_result.deinit();
270     var parent_true: f64 = 0;
271     for (parent_result.outcomes) |outcome| {
272         if (std.mem.eql(u8, outcome.value_str, "True")) parent_true = outcome.probability;
273     }
274     try std.testing.expectApproxEqAbs(@as(f64, 0.25), parent_true, 1e-9);
275 
276     var worker_result = try worker.processForm("(Marginal (flip p))") orelse return error.TestExpectedEqual;
277     defer worker_result.deinit();
278     var worker_true: f64 = 0;
279     for (worker_result.outcomes) |outcome| {
280         if (std.mem.eql(u8, outcome.value_str, "True")) worker_true = outcome.probability;
281     }
282     try std.testing.expectApproxEqAbs(@as(f64, 0.75), worker_true, 1e-9);
283 }
284 
285 test "stdlib list and nat helpers support comparison programs" {
286     const allocator = std.testing.allocator;
287     var ctx = try init(allocator);
288     defer ctx.deinit();
289 
290     const maybe_result = try ctx.processForm(
291         "(Marginal (and (eq_nat (length (Cons True (Cons False Nil))) 2) (list_eq eq_nat (take 2 (Cons 0 (Cons 1 (Cons 2 Nil)))) (Cons 0 (Cons 1 Nil)))))",
292     );
293     try std.testing.expect(maybe_result != null);
294     var result = maybe_result.?;
295     defer result.deinit();
296 
297     try std.testing.expect(!result.program_error);
298     try std.testing.expectEqual(@as(usize, 1), result.outcomes.len);
299     try std.testing.expectEqualStrings("True", result.outcomes[0].value_str);
300     try std.testing.expectApproxEqAbs(@as(f64, 1.0), result.outcomes[0].probability, 1e-10);
301 }
302 
303 test "stdlib operator names support Pluck.jl programs" {
304     const allocator = std.testing.allocator;
305     var ctx = try init(allocator);
306     defer ctx.deinit();
307 
308     const maybe_result = try ctx.processForm(
309         "(Marginal (and (nat=? (+ 1 2) 3) (and (nat=? (- 3 1) 2) (and (nat=? (* 2 3) 6) (and (> 3 2) (and (< 2 3) (and (>= 3 3) (and (<= 3 3) (and (list=? nat=? [1 2] (Cons 1 (Cons 2 Nil))) (=? (Cons True Nil) (Cons True Nil)))))))))))",
310     );
311     try std.testing.expect(maybe_result != null);
312     var result = maybe_result.?;
313     defer result.deinit();
314 
315     try std.testing.expect(!result.program_error);
316     try std.testing.expectEqual(@as(usize, 1), result.outcomes.len);
317     try std.testing.expectEqualStrings("True", result.outcomes[0].value_str);
318     try std.testing.expectApproxEqAbs(@as(f64, 1.0), result.outcomes[0].probability, 1e-10);
319 }
320 
321 test "stdlib reference helpers support Pluck.jl programs" {
322     const allocator = std.testing.allocator;
323     var ctx = try init(allocator);
324     defer ctx.deinit();
325 
326     const maybe_result = try ctx.processForm(
327         "(Marginal (and (nat=? (dec 2) 1) (and (iseven 2) (list=? nat=? (zip_with + [1 2] [3 4]) [4 6]))))",
328     );
329     try std.testing.expect(maybe_result != null);
330     var result = maybe_result.?;
331     defer result.deinit();
332 
333     try std.testing.expect(!result.program_error);
334     try std.testing.expectEqual(@as(usize, 1), result.outcomes.len);
335     try std.testing.expectEqualStrings("True", result.outcomes[0].value_str);
336     try std.testing.expectApproxEqAbs(@as(f64, 1.0), result.outcomes[0].probability, 1e-10);
337 }
338 
339 test "stdlib reference list helpers match Pluck.jl names" {
340     const allocator = std.testing.allocator;
341     var ctx = try init(allocator);
342     defer ctx.deinit();
343 
344     const maybe_result = try ctx.processForm(
345         "(Marginal (and (nat=? (mod 7 3) 1) (and (isempty (cdr_safe Nil)) (and (nat=? (fold + 0 [1 2 3]) 6) (and (list=? nat=? (mapi + [2 3]) [2 4]) (and (list=? nat=? (filter (fn x -> (< x 3)) [1 3 2]) [1 2]) (and (list=? nat=? (filteri (fn x i -> (< i 2)) [5 6 7]) [5 6]) (and (list=? nat=? (range 3) [0 1 2]) (and (list=? nat=? (append_one [1 2] 3) [1 2 3]) (and (nat=? (head-or Nil 4) 4) (and (nat=? (head-or [5] 4) 5) (and (int2=? (Int2 True False) (Int2 True False)) (constructor=? (suspendible-list=? nat=? [1] [1]) (Suspend FinallyTrue))))))))))))))",
346     );
347     try std.testing.expect(maybe_result != null);
348     var result = maybe_result.?;
349     defer result.deinit();
350 
351     try std.testing.expect(!result.program_error);
352     try std.testing.expectEqual(@as(usize, 1), result.outcomes.len);
353     try std.testing.expectEqualStrings("True", result.outcomes[0].value_str);
354     try std.testing.expectApproxEqAbs(@as(f64, 1.0), result.outcomes[0].probability, 1e-10);
355 }
356 
357 test "stdlib reference distribution helpers parse and evaluate" {
358     const allocator = std.testing.allocator;
359     var ctx = try init(allocator);
360     defer ctx.deinit();
361 
362     const uniform_result = try ctx.processForm("(Marginal (make-uniform [True False]))");
363     try std.testing.expect(uniform_result != null);
364     var uniform = uniform_result.?;
365     defer uniform.deinit();
366 
367     try std.testing.expect(!uniform.program_error);
368     try std.testing.expectEqual(@as(usize, 1), uniform.outcomes.len);
369     try std.testing.expectEqualStrings("(Cons (Pair True 0.5) (Cons (Pair False 0.5) (Nil)))", uniform.outcomes[0].value_str);
370     try std.testing.expectApproxEqAbs(@as(f64, 1.0), uniform.outcomes[0].probability, 1e-10);
371 
372     const normalized_result = try ctx.processForm("(Marginal (normalize-seq [(Pair True 1.0) (Pair False 1.0)]))");
373     try std.testing.expect(normalized_result != null);
374     var normalized = normalized_result.?;
375     defer normalized.deinit();
376 
377     try std.testing.expect(!normalized.program_error);
378     try std.testing.expectEqual(@as(usize, 1), normalized.outcomes.len);
379     try std.testing.expectEqualStrings("(Cons (Pair True 0.5) (Cons (Pair False 1) (Nil)))", normalized.outcomes[0].value_str);
380     try std.testing.expectApproxEqAbs(@as(f64, 1.0), normalized.outcomes[0].probability, 1e-10);
381 
382     const sample_result = try ctx.processForm("(Marginal (sample-seq (normalize-seq [(Pair True 1.0) (Pair False 1.0)])))");
383     try std.testing.expect(sample_result != null);
384     var sample = sample_result.?;
385     defer sample.deinit();
386 
387     try std.testing.expect(!sample.program_error);
388     try std.testing.expectEqual(@as(usize, 2), sample.outcomes.len);
389     try std.testing.expectEqualStrings("False", sample.outcomes[0].value_str);
390     try std.testing.expectApproxEqAbs(@as(f64, 0.5), sample.outcomes[0].probability, 1e-10);
391 
392     const randnat_result = try ctx.processForm("(Marginal (< (randnat) 2))");
393     try std.testing.expect(randnat_result != null);
394     var randnat = randnat_result.?;
395     defer randnat.deinit();
396 
397     try std.testing.expect(!randnat.program_error);
398     try std.testing.expectEqual(@as(usize, 2), randnat.outcomes.len);
399     try std.testing.expectEqualStrings("False", randnat.outcomes[0].value_str);
400     try std.testing.expectApproxEqAbs(@as(f64, 0.25), randnat.outcomes[0].probability, 1e-10);
401     try std.testing.expectEqualStrings("True", randnat.outcomes[1].value_str);
402     try std.testing.expectApproxEqAbs(@as(f64, 0.75), randnat.outcomes[1].probability, 1e-10);
403 }
404 
405 test "lowercase constructors from Pluck.jl programs parse and evaluate" {
406     const allocator = std.testing.allocator;
407     var ctx = try init(allocator);
408     defer ctx.deinit();
409 
410     try std.testing.expect(try ctx.processForm("(define-type char (a_) (b_))") == null);
411 
412     const maybe_result = try ctx.processForm("(Marginal (constructor=? (a_) (a_)))");
413     try std.testing.expect(maybe_result != null);
414     var result = maybe_result.?;
415     defer result.deinit();
416 
417     try std.testing.expect(!result.program_error);
418     try std.testing.expectEqual(@as(usize, 1), result.outcomes.len);
419     try std.testing.expectEqualStrings("True", result.outcomes[0].value_str);
420     try std.testing.expectApproxEqAbs(@as(f64, 1.0), result.outcomes[0].probability, 1e-10);
421 }
422 
423 test "full distribution resolves finite lazy values beyond one thousand thunks" {
424     const allocator = std.testing.allocator;
425     var ctx = try init(allocator);
426     defer ctx.deinit();
427 
428     try std.testing.expect(try ctx.processForm("(define (long-list n) (case n of O => Nil | S k => (Cons O (long-list k))))") == null);
429 
430     const maybe_result = try ctx.processForm("(Marginal (long-list 1005))");
431     try std.testing.expect(maybe_result != null);
432     var result = maybe_result.?;
433     defer result.deinit();
434 
435     try std.testing.expect(!result.program_error);
436     try std.testing.expect(result.limit_reason == null);
437     try std.testing.expectEqual(@as(usize, 1), result.outcomes.len);
438     try std.testing.expectApproxEqAbs(@as(f64, 1.0), result.outcomes[0].probability, 1e-10);
439 }
440 
441 test "stdlib ADT equality helpers compare constructed values" {
442     const allocator = std.testing.allocator;
443     var ctx = try init(allocator);
444     defer ctx.deinit();
445 
446     const maybe_result = try ctx.processForm(
447         "(Marginal (adt_eq (Cons True (Cons False Nil)) (Cons True (Cons False Nil))))",
448     );
449     try std.testing.expect(maybe_result != null);
450     var result = maybe_result.?;
451     defer result.deinit();
452 
453     try std.testing.expect(!result.program_error);
454     try std.testing.expectEqual(@as(usize, 1), result.outcomes.len);
455     try std.testing.expectEqualStrings("True", result.outcomes[0].value_str);
456     try std.testing.expectApproxEqAbs(@as(f64, 1.0), result.outcomes[0].probability, 1e-10);
457 }