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 }