lib/pluck/src/profiling/internal/grammar.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2
3 const pluck = @import("pluck");
4 const bench_util = @import("util.zig");
5
6 const pexpr = pluck.pexpr;
7 const ToplevelContext = pluck.toplevel.ToplevelContext;
8 const QueryResult = pluck.toplevel.QueryResult;
9 const LimitReason = pluck.evaluator.LimitReason;
10
11 const model_defs = [_][]const u8{
12 "(define-type ab (A) (B))",
13 "(define-type param (P1) (P2) (P3) (P4) (P5) (P6) (P7) (P8) (P9))",
14 "(define (sample-param) (uniform (P1) (P2) (P3) (P4) (P5) (P6) (P7) (P8) (P9)))",
15 \\(define (pvalue p)
16 \\ (case p of
17 \\ P1 => 0.1
18 \\ | P2 => 0.2
19 \\ | P3 => 0.3
20 \\ | P4 => 0.4
21 \\ | P5 => 0.5
22 \\ | P6 => 0.6
23 \\ | P7 => 0.7
24 \\ | P8 => 0.8
25 \\ | P9 => 0.9))
26 ,
27 "(define (sample-ab p) (if (flip (pvalue p)) (A) (B)))",
28 "(define (sample-list p n) (case n of O => (Nil) | S m => (Cons (sample-ab p) (sample-list p m))))",
29 "(define (all-a n) (case n of O => (Nil) | S m => (Cons (A) (all-a m))))",
30 };
31
32 const Metrics = struct {
33 outcomes: usize,
34 variables: usize,
35 nodes: usize,
36 forward_calls: usize,
37 recursive_calls: usize,
38 ite_cache_hits: u64,
39 ite_cache_misses: u64,
40 unique_table_grows: u64,
41 ite_cache_grows: u64,
42 query_ns: u64,
43 total_ns: u64,
44 wmc_ns: u64,
45 max_error: f64,
46 limit_reason: ?LimitReason,
47 program_error: bool,
48 };
49
50 fn loadModel(ctx: *ToplevelContext) !void {
51 for (model_defs) |form| {
52 if (try ctx.processForm(form)) |result| {
53 var owned = result;
54 owned.deinit();
55 }
56 }
57 }
58
59 fn posteriorExpr(allocator: std.mem.Allocator, evidence_len: usize) ![]u8 {
60 return std.fmt.allocPrint(
61 allocator,
62 "(let ((p (sample-param))) (Posterior p (list=? constructor=? (sample-list p {d}) (all-a {d}))))",
63 .{ evidence_len, evidence_len },
64 );
65 }
66
67 fn power(base: f64, exponent: usize) f64 {
68 var result: f64 = 1.0;
69 for (0..exponent) |_| {
70 result *= base;
71 }
72 return result;
73 }
74
75 fn expectedProbability(param_index: usize, evidence_len: usize) f64 {
76 var total: f64 = 0.0;
77 for (1..10) |idx| {
78 total += power(@as(f64, @floatFromInt(idx)) / 10.0, evidence_len);
79 }
80 return power(@as(f64, @floatFromInt(param_index)) / 10.0, evidence_len) / total;
81 }
82
83 fn paramIndex(value: []const u8) ?usize {
84 if (value.len != 4) return null;
85 if (value[0] != '(' or value[1] != 'P' or value[3] != ')') return null;
86 if (value[2] < '1' or value[2] > '9') return null;
87 return value[2] - '0';
88 }
89
90 fn maxPosteriorError(result: *const QueryResult, evidence_len: usize) !f64 {
91 var seen = @as([10]bool, @splat(false));
92 var max_error: f64 = 0.0;
93 for (result.outcomes) |outcome| {
94 const idx = paramIndex(outcome.value_str) orelse return error.UnexpectedOutcome;
95 seen[idx] = true;
96 const expected = expectedProbability(idx, evidence_len);
97 max_error = @max(max_error, @abs(outcome.probability - expected));
98 }
99 for (1..10) |idx| {
100 if (!seen[idx]) {
101 max_error = @max(max_error, expectedProbability(idx, evidence_len));
102 }
103 }
104 return max_error;
105 }
106
107 fn limitTag(reason: ?LimitReason) []const u8 {
108 return if (reason) |r| switch (r) {
109 .max_depth => "max_depth",
110 .time_limit => "time_limit",
111 .ite_limit => "ite_limit",
112 .factor_weight_too_complex => "factor_weight_too_complex",
113 } else "";
114 }
115
116 fn runQuery(ctx: *ToplevelContext, allocator: std.mem.Allocator, evidence_len: usize) !Metrics {
117 const expr = try posteriorExpr(allocator, evidence_len);
118 defer allocator.free(expr);
119
120 const start = bench_util.nowNs();
121 const maybe_result = try ctx.processForm(expr);
122 const total_ns: u64 = bench_util.elapsedNs(start);
123 if (maybe_result == null) return error.MissingQueryResult;
124
125 var result = maybe_result.?;
126 defer result.deinit();
127
128 const max_error = if (!result.program_error and result.limit_reason == null)
129 try maxPosteriorError(&result, evidence_len)
130 else
131 1.0;
132
133 return .{
134 .outcomes = result.outcomes.len,
135 .variables = result.stats.variable_count,
136 .nodes = result.stats.node_count,
137 .forward_calls = result.stats.num_forward_calls,
138 .recursive_calls = result.stats.num_recursive_calls,
139 .ite_cache_hits = result.stats.ite_cache_hits,
140 .ite_cache_misses = result.stats.ite_cache_misses,
141 .unique_table_grows = result.stats.unique_table_grows,
142 .ite_cache_grows = result.stats.ite_cache_grows,
143 .query_ns = result.stats.time_ns,
144 .total_ns = total_ns,
145 .wmc_ns = result.stats.wmc_time_ns,
146 .max_error = max_error,
147 .limit_reason = result.limit_reason,
148 .program_error = result.program_error,
149 };
150 }
151
152 fn parsePosteriorExpr(
153 allocator: std.mem.Allocator,
154 ctx: *const ToplevelContext,
155 evidence_len: usize,
156 ) !*pexpr.PExpr {
157 const source = try posteriorExpr(allocator, evidence_len);
158 defer allocator.free(source);
159 return try pexpr.parseExpr(allocator, source, &ctx.types, &ctx.definitions);
160 }
161
162 fn runParsedQuery(ctx: *ToplevelContext, expr: *pexpr.PExpr, evidence_len: usize) !Metrics {
163 const start = bench_util.nowNs();
164 var result = try ctx.runQueryExpr(expr);
165 const total_ns: u64 = bench_util.elapsedNs(start);
166 defer result.deinit();
167
168 const max_error = if (!result.program_error and result.limit_reason == null)
169 try maxPosteriorError(&result, evidence_len)
170 else
171 1.0;
172
173 return .{
174 .outcomes = result.outcomes.len,
175 .variables = result.stats.variable_count,
176 .nodes = result.stats.node_count,
177 .forward_calls = result.stats.num_forward_calls,
178 .recursive_calls = result.stats.num_recursive_calls,
179 .ite_cache_hits = result.stats.ite_cache_hits,
180 .ite_cache_misses = result.stats.ite_cache_misses,
181 .unique_table_grows = result.stats.unique_table_grows,
182 .ite_cache_grows = result.stats.ite_cache_grows,
183 .query_ns = result.stats.time_ns,
184 .total_ns = total_ns,
185 .wmc_ns = result.stats.wmc_time_ns,
186 .max_error = max_error,
187 .limit_reason = result.limit_reason,
188 .program_error = result.program_error,
189 };
190 }
191
192 test "BENCHMARK: coarse-to-fine AB list posterior" {
193 const benchmark_latency = bench_util.benchmarkScope("coarse-to-fine AB list posterior");
194 defer benchmark_latency.end();
195 const allocator = bench_util.allocator();
196 const bench = bench_util.BenchConfig.init();
197
198 var ctx = try ToplevelContext.init(allocator);
199 defer ctx.deinit();
200
201 try loadModel(&ctx);
202
203 bench_util.stdout("\n=== Benchmark: CoarseToFine AB list posterior ===\n", .{});
204 bench_util.stdout("benchmark,evidence_len,outcomes,vars,nodes,forward_calls,recursive_calls,ite_cache_hits,ite_cache_misses,unique_table_grows,ite_cache_grows,query_ns,total_ns,wmc_ns,max_error,limit,program_error\n", .{});
205
206 const evidence_lengths = [_]usize{ 12, 14, 16 };
207 const len_count = bench.limit(evidence_lengths.len, 1);
208
209 for (evidence_lengths[0..len_count]) |evidence_len| {
210 const metrics = try runQuery(&ctx, allocator, evidence_len);
211 try std.testing.expect(!metrics.program_error);
212 try std.testing.expect(metrics.limit_reason == null);
213 try std.testing.expectEqual(@as(usize, 9), metrics.outcomes);
214 try std.testing.expect(metrics.max_error < 1e-12);
215 bench_util.stdout("ctf_ab,{d},{d},{d},{d},{d},{d},{d},{d},{d},{d},{d},{d},{d},{e},{s},{d}\n", .{
216 evidence_len,
217 metrics.outcomes,
218 metrics.variables,
219 metrics.nodes,
220 metrics.forward_calls,
221 metrics.recursive_calls,
222 metrics.ite_cache_hits,
223 metrics.ite_cache_misses,
224 metrics.unique_table_grows,
225 metrics.ite_cache_grows,
226 metrics.query_ns,
227 metrics.total_ns,
228 metrics.wmc_ns,
229 metrics.max_error,
230 limitTag(metrics.limit_reason),
231 @intFromBool(metrics.program_error),
232 });
233 }
234 }
235
236 test "BENCHMARK: coarse-to-fine AB parsed posterior core" {
237 const benchmark_latency = bench_util.benchmarkScope("coarse-to-fine AB parsed posterior core");
238 defer benchmark_latency.end();
239 const allocator = bench_util.allocator();
240 const bench = bench_util.BenchConfig.init();
241
242 var ctx = try ToplevelContext.init(allocator);
243 defer ctx.deinit();
244
245 try loadModel(&ctx);
246
247 bench_util.stdout("\n=== Benchmark: CoarseToFine AB parsed posterior core ===\n", .{});
248 bench_util.stdout("benchmark,evidence_len,iters,outcomes,vars,nodes,forward_calls,recursive_calls,ite_cache_hits,ite_cache_misses,unique_table_grows,ite_cache_grows,avg_query_ns,avg_total_ns,avg_wmc_ns,max_error,limit,program_error\n", .{});
249
250 const evidence_lengths = [_]usize{ 12, 14, 16 };
251 const len_count = bench.limit(evidence_lengths.len, 1);
252 const iters = bench.iters(12, 4);
253
254 var parse_arena = std.heap.ArenaAllocator.init(allocator);
255 defer parse_arena.deinit();
256 const parse_alloc = parse_arena.allocator();
257
258 for (evidence_lengths[0..len_count]) |evidence_len| {
259 const expr = try parsePosteriorExpr(parse_alloc, &ctx, evidence_len);
260 var last_metrics: Metrics = undefined;
261 var query_total: u64 = 0;
262 var total_total: u64 = 0;
263 var wmc_total: u64 = 0;
264 var max_error: f64 = 0.0;
265
266 for (0..iters) |_| {
267 const metrics = try runParsedQuery(&ctx, expr, evidence_len);
268 try std.testing.expect(!metrics.program_error);
269 try std.testing.expect(metrics.limit_reason == null);
270 try std.testing.expectEqual(@as(usize, 9), metrics.outcomes);
271 try std.testing.expect(metrics.max_error < 1e-12);
272 last_metrics = metrics;
273 query_total += metrics.query_ns;
274 total_total += metrics.total_ns;
275 wmc_total += metrics.wmc_ns;
276 max_error = @max(max_error, metrics.max_error);
277 }
278
279 bench_util.stdout("ctf_ab_parsed,{d},{d},{d},{d},{d},{d},{d},{d},{d},{d},{d},{d},{d},{d},{e},{s},{d}\n", .{
280 evidence_len,
281 iters,
282 last_metrics.outcomes,
283 last_metrics.variables,
284 last_metrics.nodes,
285 last_metrics.forward_calls,
286 last_metrics.recursive_calls,
287 last_metrics.ite_cache_hits,
288 last_metrics.ite_cache_misses,
289 last_metrics.unique_table_grows,
290 last_metrics.ite_cache_grows,
291 query_total / iters,
292 total_total / iters,
293 wmc_total / iters,
294 max_error,
295 limitTag(last_metrics.limit_reason),
296 @intFromBool(last_metrics.program_error),
297 });
298 }
299 }