lib/pluck/src/lpsmc.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const random_seed = @import("seed.zig");
3 const Allocator = std.mem.Allocator;
4
5 const bdd = @import("bdd.zig");
6 const Bdd = bdd.Bdd;
7 const VarLabel = bdd.VarLabel;
8 const Manager = bdd.Manager;
9 const WmcParams = bdd.WmcParams;
10
11 const runtime = @import("runtime.zig");
12 const RuntimeValue = runtime.RuntimeValue;
13 const GuardedWorlds = runtime.GuardedWorlds;
14
15 pub const World = runtime.GuardedWorld;
16
17 pub const WorldsResult = GuardedWorlds;
18
19 const VarLabelSet = bdd.VarLabelSet;
20
21 pub const EvaluateThunkFn = *const fn (
22 allocator: Allocator,
23 val: *RuntimeValue,
24 path_condition: Bdd,
25 state: *anyopaque,
26 ) anyerror!WorldsResult;
27
28 pub const FreeWorldsSliceFn = *const fn (allocator: Allocator, worlds: []World) void;
29
30 pub const SetWeightFn = *const fn (wmc_params: *WmcParams, variable: VarLabel, low: f64, high: f64) Allocator.Error!void;
31
32 pub const EvaluatorOps = struct {
33 evaluateThunk: EvaluateThunkFn,
34 freeWorldsSlice: FreeWorldsSliceFn,
35 setWeight: SetWeightFn,
36 state: *anyopaque,
37 wmc_params: *WmcParams,
38 };
39
40 fn setWeightForTests(wmc_params: *WmcParams, variable: VarLabel, low: f64, high: f64) Allocator.Error!void {
41 return wmc_params.setWeight(variable, low, high);
42 }
43
44 pub const PathChoice = struct {
45 top_k_bdd: Bdd,
46 sampled_bdd: ?Bdd,
47 sampled_probability: f64,
48 k_used: usize,
49 ess_ratio: f64,
50 depends_on_vars: VarLabelSet,
51
52 const Self = @This();
53
54 pub fn deinit(self: *Self, allocator: Allocator) void {
55 self.depends_on_vars.deinit(allocator);
56 }
57 };
58
59 pub const SubproblemCache = struct {
60 available_info: Bdd,
61 return_worlds: std.ArrayList(World),
62 multiplier: f64,
63 depends_on_vars: VarLabelSet,
64
65 const Self = @This();
66
67 pub fn deinit(self: *Self, allocator: Allocator) void {
68 self.return_worlds.deinit(allocator);
69 self.depends_on_vars.deinit(allocator);
70 }
71 };
72
73 pub const SuspendedComputation = struct {
74 continuation: *RuntimeValue,
75 guard: Bdd,
76 multiplier: f64,
77 iteration: u32,
78
79 const Self = @This();
80
81 pub fn init(continuation: *RuntimeValue, guard: Bdd, multiplier: f64, iteration: u32) Self {
82 return Self{
83 .continuation = continuation,
84 .guard = guard,
85 .multiplier = multiplier,
86 .iteration = iteration,
87 };
88 }
89 };
90
91 pub const WorkItem = struct {
92 continuation: *RuntimeValue,
93 available_info: Bdd,
94 multiplier: f64,
95 iteration: u32,
96 };
97
98 pub const RunToSuspendResult = struct {
99 return_worlds: []World,
100 suspended: []SuspendedComputation,
101 validity_guard: Bdd,
102
103 const Self = @This();
104
105 pub fn hasSuspensions(self: *const Self) bool {
106 return self.suspended.len > 0;
107 }
108
109 pub fn isComplete(self: *const Self) bool {
110 return self.suspended.len == 0;
111 }
112 };
113
114 pub const LPSMCVarianceStats = struct {
115 sum_weights: f64 = 0.0,
116 sum_squared_weights: f64 = 0.0,
117 num_samples: u32 = 0,
118 max_multiplier: f64 = 1.0,
119 high_variance_warning: bool = false,
120
121 pub const DEFAULT_ESS_THRESHOLD: f64 = 0.1;
122
123 pub fn effectiveSampleSize(self: *const LPSMCVarianceStats) f64 {
124 if (self.sum_squared_weights == 0.0) return 0.0;
125 return (self.sum_weights * self.sum_weights) / self.sum_squared_weights;
126 }
127
128 pub fn essRatio(self: *const LPSMCVarianceStats) f64 {
129 if (self.num_samples == 0) return 1.0;
130 return self.effectiveSampleSize() / @as(f64, @floatFromInt(self.num_samples));
131 }
132
133 pub fn recordWeight(self: *LPSMCVarianceStats, weight: f64) void {
134 self.sum_weights += weight;
135 self.sum_squared_weights += weight * weight;
136 self.num_samples += 1;
137 if (weight > self.max_multiplier) {
138 self.max_multiplier = weight;
139 }
140 }
141
142 pub fn isHighVariance(self: *const LPSMCVarianceStats, threshold: f64) bool {
143 return self.essRatio() < threshold;
144 }
145 };
146
147 pub const AdaptiveKPolicy = struct {
148 enabled: bool = false,
149 min_k: usize = 1,
150 max_k: usize = 64,
151 ess_low: f64 = 0.25,
152 ess_high: f64 = 0.8,
153 min_samples: u32 = 2,
154 growth_factor: f64 = 2.0,
155 shrink_factor: f64 = 0.5,
156 sampled_prob_low: f64 = 0.05,
157 };
158
159 pub fn normalizeAdaptiveKPolicy(policy: AdaptiveKPolicy) AdaptiveKPolicy {
160 var out = policy;
161 if (out.min_k == 0) out.min_k = 1;
162 if (out.max_k < out.min_k) out.max_k = out.min_k;
163 if (out.ess_low <= 0.0 or out.ess_low >= 1.0) out.ess_low = 0.25;
164 if (out.ess_high <= out.ess_low or out.ess_high > 1.0) out.ess_high = 0.8;
165 if (out.min_samples == 0) out.min_samples = 1;
166 if (out.growth_factor < 1.0) out.growth_factor = 1.0;
167 if (out.shrink_factor <= 0.0 or out.shrink_factor >= 1.0) out.shrink_factor = 0.5;
168 if (out.sampled_prob_low <= 0.0 or out.sampled_prob_low > 1.0) out.sampled_prob_low = 0.05;
169 return out;
170 }
171
172 pub fn clampAdaptiveK(policy: AdaptiveKPolicy, k: usize) usize {
173 if (!policy.enabled) return k;
174 return std.math.clamp(k, policy.min_k, policy.max_k);
175 }
176
177 pub fn nextAdaptiveK(
178 policy: AdaptiveKPolicy,
179 current_k: usize,
180 stats: *const LPSMCVarianceStats,
181 used_sampling: bool,
182 sampled_probability: f64,
183 ) usize {
184 if (!policy.enabled) return current_k;
185 var k = clampAdaptiveK(policy, current_k);
186
187 if (stats.num_samples < policy.min_samples and !used_sampling) {
188 return k;
189 }
190
191 const ess_ratio = stats.essRatio();
192 const low_sample_prob = used_sampling and sampled_probability > 0.0 and sampled_probability < policy.sampled_prob_low;
193
194 if (low_sample_prob or ess_ratio < policy.ess_low) {
195 const grown = @ceil(@as(f64, @floatFromInt(k)) * policy.growth_factor);
196 const grown_k = @as(usize, @intFromFloat(grown));
197 k = std.math.clamp(grown_k, policy.min_k, policy.max_k);
198 } else if (ess_ratio > policy.ess_high) {
199 const shrunk = @floor(@as(f64, @floatFromInt(k)) * policy.shrink_factor);
200 const shrunk_k = @as(usize, @intFromFloat(@max(shrunk, 1.0)));
201 k = std.math.clamp(shrunk_k, policy.min_k, policy.max_k);
202 }
203
204 return k;
205 }
206
207 pub const LpsmcRunStats = struct {
208 total_k: usize = 0,
209 subproblem_count: u32 = 0,
210 min_k: usize = 0,
211 max_k: usize = 0,
212 avg_k: f64 = 0.0,
213 iterations: u32 = 0,
214 };
215
216 pub const IncrementalLPSMC = struct {
217 subproblem_cache: std.AutoHashMapUnmanaged(u32, SubproblemCache),
218 path_choices: std.AutoHashMapUnmanaged(u32, PathChoice),
219 var_to_subproblems: std.AutoHashMapUnmanaged(VarLabel, std.AutoHashMapUnmanaged(u32, void)),
220 last_iteration_count: u32,
221 final_return_worlds: []World,
222 final_multiplier: f64,
223 variance_stats: LPSMCVarianceStats,
224 last_run_stats: LpsmcRunStats,
225 allocator: Allocator,
226 prng: std.Random.DefaultPrng,
227 cache_valid: bool,
228
229 pub const DEFAULT_ESS_THRESHOLD: f64 = 0.1;
230 };
231
232 pub fn init(allocator: Allocator) IncrementalLPSMC {
233 const prng = std.Random.DefaultPrng.init(random_seed.systemSeed());
234
235 return .{
236 .subproblem_cache = .{},
237 .path_choices = .{},
238 .var_to_subproblems = .{},
239 .last_iteration_count = 0,
240 .final_return_worlds = &[_]World{},
241 .final_multiplier = 1.0,
242 .variance_stats = .{},
243 .last_run_stats = .{},
244 .allocator = allocator,
245 .prng = prng,
246 .cache_valid = false,
247 };
248 }
249
250 pub fn deinit(lpsmc: *IncrementalLPSMC) void {
251 var cache_iter = lpsmc.subproblem_cache.iterator();
252 while (cache_iter.next()) |entry| {
253 entry.value_ptr.deinit(lpsmc.allocator);
254 }
255 lpsmc.subproblem_cache.deinit(lpsmc.allocator);
256
257 var path_iter = lpsmc.path_choices.iterator();
258 while (path_iter.next()) |entry| {
259 entry.value_ptr.deinit(lpsmc.allocator);
260 }
261 lpsmc.path_choices.deinit(lpsmc.allocator);
262
263 var var_iter = lpsmc.var_to_subproblems.iterator();
264 while (var_iter.next()) |entry| {
265 entry.value_ptr.deinit(lpsmc.allocator);
266 }
267 lpsmc.var_to_subproblems.deinit(lpsmc.allocator);
268
269 if (lpsmc.final_return_worlds.len > 0) {
270 lpsmc.allocator.free(lpsmc.final_return_worlds);
271 }
272 }
273
274 pub fn recordDependency(lpsmc: *IncrementalLPSMC, subproblem_idx: u32, var_label: VarLabel) !void {
275 const entry = try lpsmc.var_to_subproblems.getOrPut(lpsmc.allocator, var_label);
276 if (!entry.found_existing) {
277 entry.value_ptr.* = .{};
278 }
279 try entry.value_ptr.put(lpsmc.allocator, subproblem_idx, {});
280 }
281
282 fn recordBddDependencies(lpsmc: *IncrementalLPSMC, subproblem_idx: u32, bdd_val: Bdd, manager: *const Manager) !VarLabelSet {
283 var vars = VarLabelSet{};
284 try collectBddVars(lpsmc, bdd_val, manager, &vars);
285
286 var it = vars.iterator();
287 while (it.next()) |entry| {
288 try recordDependency(lpsmc, subproblem_idx, entry.key_ptr.*);
289 }
290
291 return vars;
292 }
293
294 fn collectBddVars(lpsmc: *IncrementalLPSMC, bdd_val: Bdd, manager: *const Manager, vars: *VarLabelSet) !void {
295 if (bdd_val.isConst()) return;
296
297 const var_label = manager.topVar(bdd_val);
298 try vars.put(lpsmc.allocator, var_label, {});
299
300 const low_bdd = manager.low(bdd_val);
301 const high_bdd = manager.high(bdd_val);
302 try collectBddVars(lpsmc, low_bdd, manager, vars);
303 try collectBddVars(lpsmc, high_bdd, manager, vars);
304 }
305
306 pub fn affectsPathChoices(lpsmc: *const IncrementalLPSMC, affected_var: VarLabel) bool {
307 var path_iter = lpsmc.path_choices.iterator();
308 while (path_iter.next()) |entry| {
309 if (entry.value_ptr.depends_on_vars.contains(affected_var)) {
310 return true;
311 }
312 }
313 return false;
314 }
315
316 pub fn getAffectedSubproblems(lpsmc: *const IncrementalLPSMC, affected_var: VarLabel, out: *std.ArrayList(u32)) !void {
317 if (lpsmc.var_to_subproblems.get(affected_var)) |subproblems| {
318 var it = subproblems.iterator();
319 while (it.next()) |entry| {
320 try out.append(lpsmc.allocator, entry.key_ptr.*);
321 }
322 }
323 }
324
325 pub fn invalidateForRefinement(lpsmc: *IncrementalLPSMC, affected_var: VarLabel) bool {
326 if (affectsPathChoices(lpsmc, affected_var)) {
327 clearCaches(lpsmc);
328 return true;
329 }
330
331 if (lpsmc.var_to_subproblems.get(affected_var)) |subproblems| {
332 var it = subproblems.iterator();
333 while (it.next()) |entry| {
334 const subproblem_id = entry.key_ptr.*;
335 if (lpsmc.subproblem_cache.getPtr(subproblem_id)) |cache_entry| {
336 cache_entry.deinit(lpsmc.allocator);
337 }
338 _ = lpsmc.subproblem_cache.remove(subproblem_id);
339 if (lpsmc.path_choices.getPtr(subproblem_id)) |path_choice| {
340 path_choice.deinit(lpsmc.allocator);
341 }
342 _ = lpsmc.path_choices.remove(subproblem_id);
343 }
344 if (lpsmc.var_to_subproblems.getPtr(affected_var)) |var_entry| {
345 var_entry.deinit(lpsmc.allocator);
346 }
347 _ = lpsmc.var_to_subproblems.remove(affected_var);
348 }
349 lpsmc.cache_valid = false;
350 return false;
351 }
352
353 pub fn clearCaches(lpsmc: *IncrementalLPSMC) void {
354 var cache_iter = lpsmc.subproblem_cache.iterator();
355 while (cache_iter.next()) |entry| {
356 entry.value_ptr.deinit(lpsmc.allocator);
357 }
358 lpsmc.subproblem_cache.clearRetainingCapacity();
359
360 var path_iter = lpsmc.path_choices.iterator();
361 while (path_iter.next()) |entry| {
362 entry.value_ptr.deinit(lpsmc.allocator);
363 }
364 lpsmc.path_choices.clearRetainingCapacity();
365
366 var var_iter = lpsmc.var_to_subproblems.iterator();
367 while (var_iter.next()) |entry| {
368 entry.value_ptr.deinit(lpsmc.allocator);
369 }
370 lpsmc.var_to_subproblems.clearRetainingCapacity();
371
372 if (lpsmc.final_return_worlds.len > 0) {
373 lpsmc.allocator.free(lpsmc.final_return_worlds);
374 lpsmc.final_return_worlds = &[_]World{};
375 }
376 lpsmc.last_iteration_count = 0;
377 lpsmc.final_multiplier = 1.0;
378 lpsmc.variance_stats = .{};
379 lpsmc.cache_valid = false;
380 }
381
382 fn copyWorlds(allocator: Allocator, worlds: []const World) ![]World {
383 const copy = try allocator.alloc(World, worlds.len);
384 @memcpy(copy, worlds);
385 return copy;
386 }
387
388 fn cachedRunResult(lpsmc: *IncrementalLPSMC, result_allocator: Allocator) !?[]World {
389 if (!lpsmc.cache_valid or lpsmc.final_return_worlds.len == 0) {
390 return null;
391 }
392
393 return try copyWorlds(result_allocator, lpsmc.final_return_worlds);
394 }
395
396 fn resetRunState(lpsmc: *IncrementalLPSMC) void {
397 lpsmc.last_iteration_count = 0;
398 lpsmc.variance_stats = .{};
399 }
400
401 fn evidenceGuard(
402 allocator: Allocator,
403 evidence_thunk: ?*RuntimeValue,
404 ops: EvaluatorOps,
405 manager: *Manager,
406 ) !Bdd {
407 var guard = Bdd.TRUE;
408 if (evidence_thunk) |ev_thunk| {
409 const evidence_result = try ops.evaluateThunk(allocator, ev_thunk, Bdd.TRUE, ops.state);
410 defer ops.freeWorldsSlice(allocator, evidence_result.worlds);
411
412 guard = Bdd.FALSE;
413 for (evidence_result.worlds) |world| {
414 if (world.value.data == .constructed) {
415 const c = world.value.data.constructed;
416 if (std.mem.eql(u8, c.constructor, "True") and c.args.len == 0) {
417 guard = try manager.bddOr(guard, world.guard);
418 }
419 }
420 }
421 }
422 return guard;
423 }
424
425 fn appendInitialWork(
426 worklist: *std.ArrayList(WorkItem),
427 allocator: Allocator,
428 suspendible_thunk: *RuntimeValue,
429 guard: Bdd,
430 ) !void {
431 try worklist.append(allocator, .{
432 .continuation = suspendible_thunk,
433 .available_info = guard,
434 .multiplier = 1.0,
435 .iteration = 0,
436 });
437 }
438
439 fn appendSubproblemResult(
440 lpsmc: *IncrementalLPSMC,
441 id: u32,
442 item: WorkItem,
443 subproblem_vars: VarLabelSet,
444 worlds: []const World,
445 return_worlds: *std.ArrayList(World),
446 ) !void {
447 var cache_entry = SubproblemCache{
448 .available_info = item.available_info,
449 .return_worlds = .empty,
450 .multiplier = item.multiplier,
451 .depends_on_vars = subproblem_vars,
452 };
453 errdefer cache_entry.deinit(lpsmc.allocator);
454
455 for (worlds) |world| {
456 try return_worlds.append(lpsmc.allocator, world);
457 try cache_entry.return_worlds.append(lpsmc.allocator, world);
458 }
459
460 try lpsmc.subproblem_cache.put(lpsmc.allocator, id, cache_entry);
461 }
462
463 fn recordSuspensionStats(run_stats: *LpsmcRunStats, k_used: usize) void {
464 run_stats.total_k += k_used;
465 run_stats.subproblem_count += 1;
466 run_stats.min_k = if (run_stats.subproblem_count == 1) k_used else @min(run_stats.min_k, k_used);
467 run_stats.max_k = @max(run_stats.max_k, k_used);
468 }
469
470 fn recordSelectionState(
471 lpsmc: *IncrementalLPSMC,
472 id: u32,
473 selection: SuspensionSelection,
474 k_used: usize,
475 manager: *Manager,
476 ) !void {
477 if (selection.used_sampling) {
478 lpsmc.variance_stats.recordWeight(selection.new_multiplier);
479 }
480 const ess_ratio = lpsmc.variance_stats.essRatio();
481
482 var path_vars = try recordBddDependencies(lpsmc, id, selection.top_k_bdd, manager);
483 errdefer path_vars.deinit(lpsmc.allocator);
484
485 try lpsmc.path_choices.put(lpsmc.allocator, id, PathChoice{
486 .top_k_bdd = selection.top_k_bdd,
487 .sampled_bdd = selection.sampled_bdd,
488 .sampled_probability = selection.sampled_probability,
489 .k_used = k_used,
490 .ess_ratio = ess_ratio,
491 .depends_on_vars = path_vars,
492 });
493
494 if (selection.used_sampling and lpsmc.variance_stats.isHighVariance(LPSMCVarianceStats.DEFAULT_ESS_THRESHOLD)) {
495 lpsmc.variance_stats.high_variance_warning = true;
496 }
497 }
498
499 fn enqueueSelectionChildren(
500 allocator: Allocator,
501 worklist: *std.ArrayList(WorkItem),
502 suspended: []const SuspendedComputation,
503 selection: SuspensionSelection,
504 next_iteration: u32,
505 manager: *Manager,
506 ) !void {
507 for (suspended) |susp| {
508 const child_guard = try manager.bddAnd(selection.selection_guard, susp.guard);
509 if (child_guard.isFalse()) continue;
510 try worklist.append(allocator, .{
511 .continuation = susp.continuation,
512 .available_info = child_guard,
513 .multiplier = selection.new_multiplier,
514 .iteration = next_iteration,
515 });
516 }
517 }
518
519 fn finishRunStats(
520 lpsmc: *IncrementalLPSMC,
521 subproblem_count: u32,
522 multiplier: f64,
523 max_iterations: u32,
524 run_stats: LpsmcRunStats,
525 ) void {
526 var stats = run_stats;
527 lpsmc.last_iteration_count = subproblem_count;
528 lpsmc.final_multiplier = multiplier;
529 if (lpsmc.last_iteration_count == 0) {
530 lpsmc.last_iteration_count = max_iterations;
531 }
532 stats.iterations = lpsmc.last_iteration_count;
533 if (stats.subproblem_count > 0) {
534 stats.avg_k = @as(f64, @floatFromInt(stats.total_k)) /
535 @as(f64, @floatFromInt(stats.subproblem_count));
536 }
537 lpsmc.last_run_stats = stats;
538 }
539
540 fn publishRunResult(
541 lpsmc: *IncrementalLPSMC,
542 result_allocator: Allocator,
543 return_worlds: *std.ArrayList(World),
544 ) ![]World {
545 const cached_result = try return_worlds.toOwnedSlice(lpsmc.allocator);
546
547 if (lpsmc.final_return_worlds.len > 0) {
548 lpsmc.allocator.free(lpsmc.final_return_worlds);
549 }
550 lpsmc.final_return_worlds = cached_result;
551 lpsmc.cache_valid = true;
552
553 return try copyWorlds(result_allocator, cached_result);
554 }
555
556 pub fn run(
557 lpsmc: *IncrementalLPSMC,
558 result_allocator: Allocator,
559 suspendible_thunk: *RuntimeValue,
560 evidence_thunk: ?*RuntimeValue,
561 k: usize,
562 k_policy: AdaptiveKPolicy,
563 ops: EvaluatorOps,
564 manager: *Manager,
565 ) ![]World {
566 if (try cachedRunResult(lpsmc, result_allocator)) |copy| {
567 return copy;
568 }
569
570 resetRunState(lpsmc);
571
572 const max_iterations: u32 = 1000;
573 const policy = normalizeAdaptiveKPolicy(k_policy);
574 var current_k: usize = clampAdaptiveK(policy, k);
575 var run_stats = LpsmcRunStats{};
576
577 const random = lpsmc.prng.random();
578
579 const initial_guard = try evidenceGuard(result_allocator, evidence_thunk, ops, manager);
580 if (initial_guard.isFalse()) {
581 lpsmc.last_run_stats = .{};
582 return &[_]World{};
583 }
584
585 var return_worlds = std.ArrayList(World).empty;
586 errdefer return_worlds.deinit(lpsmc.allocator);
587
588 var worklist: std.ArrayList(WorkItem) = .empty;
589 defer worklist.deinit(lpsmc.allocator);
590 try appendInitialWork(&worklist, lpsmc.allocator, suspendible_thunk, initial_guard);
591
592 var multiplier: f64 = 1.0;
593 var subproblem_id: u32 = 0;
594 var processed: u32 = 0;
595 while (worklist.pop()) |item| {
596 processed += 1;
597 if (processed > max_iterations) break;
598 const id = subproblem_id;
599 subproblem_id += 1;
600 multiplier = item.multiplier;
601
602 const run_result = try runToSuspend(
603 result_allocator,
604 item.continuation,
605 item.available_info,
606 item.multiplier,
607 item.iteration,
608 ops,
609 manager,
610 );
611 defer if (run_result.return_worlds.len > 0) result_allocator.free(run_result.return_worlds);
612 defer if (run_result.suspended.len > 0) result_allocator.free(run_result.suspended);
613
614 const subproblem_vars = try recordBddDependencies(lpsmc, id, item.available_info, manager);
615 try appendSubproblemResult(lpsmc, id, item, subproblem_vars, run_result.return_worlds, &return_worlds);
616
617 if (!run_result.hasSuspensions() or item.iteration + 1 >= max_iterations) {
618 continue;
619 }
620
621 const k_used = current_k;
622 recordSuspensionStats(&run_stats, k_used);
623
624 const selection_opt = try selectSuspension(
625 lpsmc.allocator,
626 run_result.suspended,
627 item.available_info,
628 item.multiplier,
629 k_used,
630 ops,
631 manager,
632 random,
633 );
634
635 if (selection_opt == null) {
636 continue;
637 }
638
639 const selection = selection_opt.?;
640
641 try recordSelectionState(lpsmc, id, selection, k_used, manager);
642 current_k = nextAdaptiveK(policy, current_k, &lpsmc.variance_stats, selection.used_sampling, selection.sampled_probability);
643
644 const next_iteration = item.iteration + 1;
645 try enqueueSelectionChildren(lpsmc.allocator, &worklist, run_result.suspended, selection, next_iteration, manager);
646 }
647
648 finishRunStats(lpsmc, subproblem_id, multiplier, max_iterations, run_stats);
649 return try publishRunResult(lpsmc, result_allocator, &return_worlds);
650 }
651
652 fn appendWeightedReturnWorld(
653 allocator: Allocator,
654 output: *std.ArrayList(World),
655 value: *RuntimeValue,
656 guard: Bdd,
657 available_info: Bdd,
658 multiplier: f64,
659 ops: EvaluatorOps,
660 manager: *Manager,
661 ) !void {
662 const new_var = try manager.newVar(true);
663 const new_guard = try manager.bddAnd(new_var, try manager.bddAnd(guard, available_info));
664 try ops.setWeight(ops.wmc_params, manager.topVar(new_var), 1.0, multiplier);
665 if (!new_guard.isFalse()) {
666 try output.append(allocator, World{
667 .value = value,
668 .guard = new_guard,
669 });
670 }
671 }
672
673 fn handleConstructedRunWorld(
674 allocator: Allocator,
675 world: World,
676 available_info: Bdd,
677 multiplier: f64,
678 iteration: u32,
679 ops: EvaluatorOps,
680 manager: *Manager,
681 return_worlds: *std.ArrayList(World),
682 suspended: *std.ArrayList(SuspendedComputation),
683 ) !bool {
684 if (world.value.data != .constructed) return false;
685
686 const c = world.value.data.constructed;
687 if (std.mem.eql(u8, c.constructor, "Suspend")) {
688 if (c.args.len >= 1) {
689 try suspended.append(allocator, SuspendedComputation.init(
690 c.args[0],
691 world.guard,
692 multiplier,
693 iteration,
694 ));
695 }
696 return true;
697 }
698
699 if (std.mem.eql(u8, c.constructor, "FinallyTrue")) {
700 const true_val = try RuntimeValue.initConstructed(allocator, "True", &[_]*RuntimeValue{});
701 try appendWeightedReturnWorld(allocator, return_worlds, true_val, world.guard, available_info, multiplier, ops, manager);
702 return true;
703 }
704
705 if (std.mem.eql(u8, c.constructor, "FinallyFalse")) {
706 const false_val = try RuntimeValue.initConstructed(allocator, "False", &[_]*RuntimeValue{});
707 try appendWeightedReturnWorld(allocator, return_worlds, false_val, world.guard, available_info, multiplier, ops, manager);
708 return true;
709 }
710
711 return false;
712 }
713
714 pub fn runToSuspend(
715 allocator: Allocator,
716 thunk: *RuntimeValue,
717 available_info: Bdd,
718 multiplier: f64,
719 iteration: u32,
720 ops: EvaluatorOps,
721 manager: *Manager,
722 ) !RunToSuspendResult {
723 const eval_result = try ops.evaluateThunk(allocator, thunk, available_info, ops.state);
724 defer ops.freeWorldsSlice(allocator, eval_result.worlds);
725
726 var return_worlds_list: std.ArrayList(World) = .empty;
727 errdefer return_worlds_list.deinit(allocator);
728
729 var suspended_list: std.ArrayList(SuspendedComputation) = .empty;
730 errdefer suspended_list.deinit(allocator);
731
732 for (eval_result.worlds) |world| {
733 if (try handleConstructedRunWorld(allocator, world, available_info, multiplier, iteration, ops, manager, &return_worlds_list, &suspended_list)) {
734 continue;
735 }
736
737 try appendWeightedReturnWorld(allocator, &return_worlds_list, world.value, world.guard, available_info, multiplier, ops, manager);
738 }
739
740 return RunToSuspendResult{
741 .return_worlds = try return_worlds_list.toOwnedSlice(allocator),
742 .suspended = try suspended_list.toOwnedSlice(allocator),
743 .validity_guard = eval_result.validity_guard,
744 };
745 }
746
747 pub const SuspensionSelection = struct {
748 selection_guard: Bdd,
749 new_multiplier: f64,
750 used_sampling: bool,
751 top_k_bdd: Bdd,
752 sampled_bdd: ?Bdd,
753 sampled_probability: f64,
754 };
755
756 const MixResult = struct {
757 mixed_bdd: Bdd,
758 mult_increment: f64,
759 };
760
761 fn mixTopKAndSample(
762 manager: *Manager,
763 ops: EvaluatorOps,
764 top_k_bdd: Bdd,
765 sampled_bdd: Bdd,
766 sampled_probability: f64,
767 ) !MixResult {
768 const mult_increment = 1.0 + (1.0 / sampled_probability);
769 const topk_prob = 1.0 / mult_increment;
770 const sampled_prob = 1.0 - topk_prob;
771
772 const mix_var = try manager.newVar(true);
773 try ops.setWeight(ops.wmc_params, manager.topVar(mix_var), sampled_prob, topk_prob);
774
775 const mixed_bdd = try manager.ite(mix_var, top_k_bdd, sampled_bdd);
776
777 return MixResult{
778 .mixed_bdd = mixed_bdd,
779 .mult_increment = mult_increment,
780 };
781 }
782
783 pub fn selectSuspension(
784 allocator: Allocator,
785 suspensions: []const SuspendedComputation,
786 available_info: Bdd,
787 current_multiplier: f64,
788 k: usize,
789 ops: EvaluatorOps,
790 manager: *Manager,
791 random: std.Random,
792 ) !?SuspensionSelection {
793 if (suspensions.len == 0) return null;
794
795 var combined_susp_guard = Bdd.FALSE;
796 for (suspensions) |susp| {
797 combined_susp_guard = try manager.bddOr(combined_susp_guard, susp.guard);
798 }
799
800 const new_available = try manager.bddAnd(available_info, combined_susp_guard);
801
802 if (new_available.isFalse()) {
803 return null;
804 }
805
806 const top_k_bdd = try bdd.topKPaths(manager, new_available, k, ops.wmc_params, allocator);
807 const posterior_guard = try manager.bddAnd(new_available, manager.bddNot(top_k_bdd));
808
809 var new_multiplier = current_multiplier;
810 var combined_guard = top_k_bdd;
811 var used_sampling = false;
812 var sampled_bdd: ?Bdd = null;
813 var sampled_probability: f64 = 0.0;
814
815 if (!posterior_guard.isFalse()) {
816 const posterior_mass = bdd.wmc(manager, posterior_guard, ops.wmc_params);
817 if (posterior_mass > 0.0) {
818 const sampled = try bdd.weightedSample(manager, posterior_guard, ops.wmc_params, random);
819
820 if (!sampled.sample.isFalse() and sampled.probability > 0.0) {
821 const conditional_prob_raw = sampled.probability / posterior_mass;
822 const conditional_prob = if (conditional_prob_raw > 1.0) 1.0 else conditional_prob_raw;
823 if (conditional_prob > 0.0) {
824 const sampled_guard = try manager.bddAnd(posterior_guard, sampled.sample);
825 sampled_bdd = sampled_guard;
826 sampled_probability = conditional_prob;
827
828 const mix_result = try mixTopKAndSample(manager, ops, top_k_bdd, sampled_guard, conditional_prob);
829 new_multiplier *= mix_result.mult_increment;
830 combined_guard = mix_result.mixed_bdd;
831 used_sampling = true;
832 }
833 }
834 }
835 }
836
837 const selection_guard = try manager.bddAnd(available_info, combined_guard);
838 if (selection_guard.isFalse()) return null;
839
840 return SuspensionSelection{
841 .selection_guard = selection_guard,
842 .new_multiplier = new_multiplier,
843 .used_sampling = used_sampling,
844 .top_k_bdd = top_k_bdd,
845 .sampled_bdd = sampled_bdd,
846 .sampled_probability = sampled_probability,
847 };
848 }
849
850 pub fn subproblemMonteCarloImpl(
851 allocator: Allocator,
852 suspendible_thunk: *RuntimeValue,
853 evidence_thunk: ?*RuntimeValue,
854 k: usize,
855 k_policy: AdaptiveKPolicy,
856 ops: EvaluatorOps,
857 manager: *Manager,
858 external_rng: ?std.Random,
859 ) ![]World {
860 const max_iterations = 1000;
861 const policy = normalizeAdaptiveKPolicy(k_policy);
862 var current_k: usize = clampAdaptiveK(policy, k);
863 var variance_stats = LPSMCVarianceStats{};
864
865 var internal_prng: ?std.Random.DefaultPrng = null;
866 const random = if (external_rng) |rng| rng else blk: {
867 internal_prng = std.Random.DefaultPrng.init(random_seed.systemSeed());
868 break :blk internal_prng.?.random();
869 };
870
871 const initial_guard = try evidenceGuard(allocator, evidence_thunk, ops, manager);
872 if (initial_guard.isFalse()) {
873 return &[_]World{};
874 }
875
876 var return_worlds: std.ArrayList(World) = .empty;
877 defer return_worlds.deinit(allocator);
878
879 var worklist: std.ArrayList(WorkItem) = .empty;
880 defer worklist.deinit(allocator);
881 try appendInitialWork(&worklist, allocator, suspendible_thunk, initial_guard);
882
883 var processed: usize = 0;
884 while (worklist.pop()) |item| {
885 processed += 1;
886 if (processed > max_iterations) break;
887
888 const run_result = try runToSuspend(
889 allocator,
890 item.continuation,
891 item.available_info,
892 item.multiplier,
893 item.iteration,
894 ops,
895 manager,
896 );
897 defer if (run_result.return_worlds.len > 0) allocator.free(run_result.return_worlds);
898 defer if (run_result.suspended.len > 0) allocator.free(run_result.suspended);
899
900 for (run_result.return_worlds) |world| {
901 try return_worlds.append(allocator, world);
902 }
903
904 if (!run_result.hasSuspensions() or item.iteration + 1 >= max_iterations) {
905 continue;
906 }
907
908 const k_used = current_k;
909 const selection_opt = try selectSuspension(
910 allocator,
911 run_result.suspended,
912 item.available_info,
913 item.multiplier,
914 k_used,
915 ops,
916 manager,
917 random,
918 );
919
920 if (selection_opt == null) {
921 continue;
922 }
923
924 const selection = selection_opt.?;
925
926 if (selection.used_sampling) {
927 variance_stats.recordWeight(selection.new_multiplier);
928 }
929 current_k = nextAdaptiveK(policy, current_k, &variance_stats, selection.used_sampling, selection.sampled_probability);
930
931 const next_iteration = item.iteration + 1;
932 for (run_result.suspended) |susp| {
933 const child_guard = try manager.bddAnd(selection.selection_guard, susp.guard);
934 if (child_guard.isFalse()) continue;
935 try worklist.append(allocator, .{
936 .continuation = susp.continuation,
937 .available_info = child_guard,
938 .multiplier = selection.new_multiplier,
939 .iteration = next_iteration,
940 });
941 }
942 }
943
944 return try return_worlds.toOwnedSlice(allocator);
945 }
946
947 test "LPSMCVarianceStats basic operations" {
948 var stats = LPSMCVarianceStats{};
949
950 try std.testing.expectEqual(@as(u32, 0), stats.num_samples);
951 try std.testing.expectEqual(@as(f64, 1.0), stats.essRatio());
952
953 stats.recordWeight(1.0);
954 stats.recordWeight(1.0);
955 stats.recordWeight(1.0);
956
957 try std.testing.expectEqual(@as(u32, 3), stats.num_samples);
958 try std.testing.expectApproxEqAbs(@as(f64, 3.0), stats.effectiveSampleSize(), 1e-10);
959 try std.testing.expectApproxEqAbs(@as(f64, 1.0), stats.essRatio(), 1e-10);
960
961 try std.testing.expect(!stats.isHighVariance(LPSMCVarianceStats.DEFAULT_ESS_THRESHOLD));
962 }
963
964 test "LPSMCVarianceStats high variance detection" {
965 var stats = LPSMCVarianceStats{};
966
967 stats.recordWeight(100.0);
968 stats.recordWeight(1.0);
969 stats.recordWeight(1.0);
970 stats.recordWeight(1.0);
971
972 try std.testing.expectEqual(@as(u32, 4), stats.num_samples);
973 const ess = stats.effectiveSampleSize();
974 try std.testing.expect(ess < 2.0);
975
976 try std.testing.expect(stats.isHighVariance(0.5));
977 }
978
979 test "AdaptiveKPolicy adjusts k based on ESS and sampled probability" {
980 var policy = AdaptiveKPolicy{
981 .enabled = true,
982 .min_k = 2,
983 .max_k = 16,
984 .ess_low = 0.3,
985 .ess_high = 0.8,
986 .min_samples = 1,
987 .growth_factor = 2.0,
988 .shrink_factor = 0.5,
989 .sampled_prob_low = 0.1,
990 };
991 policy = normalizeAdaptiveKPolicy(policy);
992
993 var stats = LPSMCVarianceStats{};
994 stats.recordWeight(100.0);
995 stats.recordWeight(1.0);
996 stats.recordWeight(1.0);
997 stats.recordWeight(1.0);
998 const k_up = nextAdaptiveK(policy, 4, &stats, true, 0.2);
999 try std.testing.expect(k_up > 4);
1000
1001 var stats_high = LPSMCVarianceStats{};
1002 stats_high.recordWeight(1.0);
1003 stats_high.recordWeight(1.0);
1004 const k_down = nextAdaptiveK(policy, 8, &stats_high, true, 0.5);
1005 try std.testing.expect(k_down < 8);
1006
1007 const k_prob = nextAdaptiveK(policy, 4, &stats_high, true, 0.01);
1008 try std.testing.expect(k_prob > 4);
1009 }
1010
1011 test "PathChoice init and deinit" {
1012 const allocator = std.testing.allocator;
1013
1014 var vars = VarLabelSet{};
1015 try vars.put(allocator, 1, {});
1016 try vars.put(allocator, 2, {});
1017
1018 var choice = PathChoice{
1019 .top_k_bdd = Bdd.TRUE,
1020 .sampled_bdd = null,
1021 .sampled_probability = 0.0,
1022 .k_used = 1,
1023 .ess_ratio = 1.0,
1024 .depends_on_vars = vars,
1025 };
1026
1027 try std.testing.expect(choice.depends_on_vars.contains(1));
1028 try std.testing.expect(choice.depends_on_vars.contains(2));
1029
1030 choice.deinit(allocator);
1031 }
1032
1033 test "IncrementalLPSMC init and deinit" {
1034 const allocator = std.testing.allocator;
1035
1036 var lpsmc = init(allocator);
1037 defer deinit(&lpsmc);
1038
1039 try std.testing.expectEqual(@as(u32, 0), lpsmc.last_iteration_count);
1040 try std.testing.expectEqual(@as(f64, 1.0), lpsmc.final_multiplier);
1041 try std.testing.expectEqual(@as(u32, 0), lpsmc.variance_stats.num_samples);
1042 }
1043
1044 test "IncrementalLPSMC clearCaches resets state" {
1045 const allocator = std.testing.allocator;
1046
1047 var lpsmc = init(allocator);
1048 defer deinit(&lpsmc);
1049
1050 lpsmc.last_iteration_count = 5;
1051 lpsmc.final_multiplier = 2.5;
1052 lpsmc.variance_stats.recordWeight(1.0);
1053 lpsmc.variance_stats.recordWeight(2.0);
1054
1055 clearCaches(&lpsmc);
1056
1057 try std.testing.expectEqual(@as(u32, 0), lpsmc.last_iteration_count);
1058 try std.testing.expectEqual(@as(f64, 1.0), lpsmc.final_multiplier);
1059 try std.testing.expectEqual(@as(u32, 0), lpsmc.variance_stats.num_samples);
1060 }
1061
1062 test "deterministic RNG produces reproducible results" {
1063 var prng1 = std.Random.DefaultPrng.init(12345);
1064 var prng2 = std.Random.DefaultPrng.init(12345);
1065
1066 const random1 = prng1.random();
1067 const random2 = prng2.random();
1068
1069 var i: usize = 0;
1070 while (i < 10) : (i += 1) {
1071 const val1 = random1.float(f64);
1072 const val2 = random2.float(f64);
1073 try std.testing.expectEqual(val1, val2);
1074 }
1075 }
1076
1077 test "SubproblemCache init and deinit" {
1078 const allocator = std.testing.allocator;
1079
1080 var vars = VarLabelSet{};
1081 try vars.put(allocator, 1, {});
1082
1083 var cache = SubproblemCache{
1084 .available_info = Bdd.TRUE,
1085 .return_worlds = .empty,
1086 .multiplier = 1.0,
1087 .depends_on_vars = vars,
1088 };
1089
1090 try std.testing.expect(cache.depends_on_vars.contains(1));
1091 try std.testing.expectEqual(@as(f64, 1.0), cache.multiplier);
1092
1093 cache.deinit(allocator);
1094 }
1095
1096 test "external_rng parameter is used when provided" {
1097 var prng1 = std.Random.DefaultPrng.init(42);
1098 var prng2 = std.Random.DefaultPrng.init(42);
1099 var prng3 = std.Random.DefaultPrng.init(99);
1100
1101 const rng1 = prng1.random();
1102 const rng2 = prng2.random();
1103 const rng3 = prng3.random();
1104
1105 const val1a = rng1.float(f64);
1106 const val2a = rng2.float(f64);
1107 try std.testing.expectEqual(val1a, val2a);
1108
1109 const val3a = rng3.float(f64);
1110 try std.testing.expect(val1a != val3a);
1111
1112 const val1b = rng1.float(f64);
1113 const val2b = rng2.float(f64);
1114 try std.testing.expectEqual(val1b, val2b);
1115 }
1116
1117 test "subproblemMonteCarloImpl accepts external RNG (API contract)" {
1118 const FnType = @TypeOf(subproblemMonteCarloImpl);
1119 const fn_info = @typeInfo(FnType).@"fn";
1120
1121 try std.testing.expectEqual(@as(usize, 8), fn_info.param_types.len);
1122
1123 try std.testing.expect(fn_info.param_types[7].? == ?std.Random);
1124 }
1125
1126 test "SuspendedComputation init" {
1127 const allocator = std.testing.allocator;
1128 _ = allocator;
1129
1130 const susp = SuspendedComputation{
1131 .continuation = undefined,
1132 .guard = Bdd.TRUE,
1133 .multiplier = 1.5,
1134 .iteration = 3,
1135 };
1136
1137 try std.testing.expectEqual(Bdd.TRUE, susp.guard);
1138 try std.testing.expectEqual(@as(f64, 1.5), susp.multiplier);
1139 try std.testing.expectEqual(@as(u32, 3), susp.iteration);
1140 }
1141
1142 test "RunToSuspendResult helper methods" {
1143 const empty_result = RunToSuspendResult{
1144 .return_worlds = &[_]World{},
1145 .suspended = &[_]SuspendedComputation{},
1146 .validity_guard = Bdd.TRUE,
1147 };
1148 try std.testing.expect(empty_result.isComplete());
1149 try std.testing.expect(!empty_result.hasSuspensions());
1150
1151 var suspensions = [_]SuspendedComputation{
1152 SuspendedComputation{
1153 .continuation = undefined,
1154 .guard = Bdd.TRUE,
1155 .multiplier = 1.0,
1156 .iteration = 0,
1157 },
1158 };
1159 const suspended_result = RunToSuspendResult{
1160 .return_worlds = &[_]World{},
1161 .suspended = &suspensions,
1162 .validity_guard = Bdd.TRUE,
1163 };
1164 try std.testing.expect(!suspended_result.isComplete());
1165 try std.testing.expect(suspended_result.hasSuspensions());
1166 }
1167
1168 test "SuspensionSelection struct fields" {
1169 const selection = SuspensionSelection{
1170 .selection_guard = Bdd.TRUE,
1171 .new_multiplier = 2.5,
1172 .used_sampling = true,
1173 .top_k_bdd = Bdd.TRUE,
1174 .sampled_bdd = Bdd.FALSE,
1175 .sampled_probability = 0.25,
1176 };
1177
1178 try std.testing.expectEqual(Bdd.TRUE, selection.selection_guard);
1179 try std.testing.expectEqual(@as(f64, 2.5), selection.new_multiplier);
1180 try std.testing.expect(selection.used_sampling);
1181 try std.testing.expectEqual(Bdd.TRUE, selection.top_k_bdd);
1182 try std.testing.expectEqual(Bdd.FALSE, selection.sampled_bdd.?);
1183 try std.testing.expectEqual(@as(f64, 0.25), selection.sampled_probability);
1184 }
1185
1186 test "runToSuspend API contract" {
1187 const FnType = @TypeOf(runToSuspend);
1188 const fn_info = @typeInfo(FnType).@"fn";
1189
1190 try std.testing.expectEqual(@as(usize, 7), fn_info.param_types.len);
1191
1192 const return_info = @typeInfo(fn_info.return_type.?);
1193 try std.testing.expect(return_info == .error_union);
1194 }
1195
1196 test "selectSuspension API contract" {
1197 const FnType = @TypeOf(selectSuspension);
1198 const fn_info = @typeInfo(FnType).@"fn";
1199
1200 try std.testing.expectEqual(@as(usize, 8), fn_info.param_types.len);
1201
1202 const return_info = @typeInfo(fn_info.return_type.?);
1203 try std.testing.expect(return_info == .error_union);
1204 }
1205
1206 test "selectSuspension returns null for empty suspensions" {
1207 const allocator = std.testing.allocator;
1208
1209 var manager = try Manager.init(allocator);
1210 defer manager.deinit();
1211
1212 var wmc_params = WmcParams.init(allocator);
1213 defer wmc_params.deinit();
1214
1215 const ops = EvaluatorOps{
1216 .evaluateThunk = undefined,
1217 .freeWorldsSlice = undefined,
1218 .setWeight = undefined,
1219 .state = undefined,
1220 .wmc_params = &wmc_params,
1221 };
1222
1223 var prng = std.Random.DefaultPrng.init(12345);
1224 const random = prng.random();
1225
1226 const empty_suspensions = &[_]SuspendedComputation{};
1227 const result = try selectSuspension(
1228 allocator,
1229 empty_suspensions,
1230 Bdd.TRUE,
1231 1.0,
1232 5,
1233 ops,
1234 &manager,
1235 random,
1236 );
1237
1238 try std.testing.expect(result == null);
1239 }
1240
1241 test "mixTopKAndSample preserves unbiased weighting" {
1242 const allocator = std.testing.allocator;
1243
1244 var manager = try Manager.init(allocator);
1245 defer manager.deinit();
1246
1247 var wmc_params = WmcParams.init(allocator);
1248 defer wmc_params.deinit();
1249
1250 const a = try manager.newVar(true);
1251 const b = try manager.newVar(true);
1252
1253 try wmc_params.setWeight(manager.topVar(a), 0.3, 0.7);
1254 try wmc_params.setWeight(manager.topVar(b), 0.6, 0.4);
1255
1256 const top_k_bdd = a;
1257 const sampled_bdd = try manager.bddAnd(manager.bddNot(a), b);
1258
1259 const ops = EvaluatorOps{
1260 .evaluateThunk = undefined,
1261 .freeWorldsSlice = undefined,
1262 .setWeight = setWeightForTests,
1263 .state = undefined,
1264 .wmc_params = &wmc_params,
1265 };
1266
1267 const sampled_probability: f64 = 0.25;
1268 const mix = try mixTopKAndSample(&manager, ops, top_k_bdd, sampled_bdd, sampled_probability);
1269
1270 const wmc_top = bdd.wmc(&manager, top_k_bdd, &wmc_params);
1271 const wmc_sampled = bdd.wmc(&manager, sampled_bdd, &wmc_params);
1272 const wmc_mixed = bdd.wmc(&manager, mix.mixed_bdd, &wmc_params);
1273
1274 const expected = wmc_top + (1.0 / sampled_probability) * wmc_sampled;
1275 const actual = mix.mult_increment * wmc_mixed;
1276
1277 try std.testing.expectApproxEqAbs(expected, actual, 1e-12);
1278 }