lib/smt/src/profiling/sat.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const bench = @import("bench");
3 const smt = @import("smt");
4
5 const assert = std.debug.assert;
6 const Literal = smt.Literal;
7 const Solver = smt.Solver;
8
9 const churn_seed: u64 = 0x5eed_5a7_0001;
10 const churn_variable_count: u32 = 96;
11 const churn_clause_count: usize = 408;
12 const churn_conflict_budget: usize = 2_000;
13
14 const pigeon_count: u32 = 8;
15 const hole_count: u32 = 7;
16
17 const incremental_seed: u64 = 0x5eed_5a7_0002;
18 const incremental_variable_count: u32 = 64;
19 const incremental_clause_count: usize = 192;
20 const incremental_rounds: usize = 32;
21 const incremental_conflict_budget: usize = 2_000;
22
23 const bounded_store_limit: usize = 256;
24
25 pub fn addTo(suite: *bench.Suite) !void {
26 try suite.add("smt sat random3 conflict churn", randomConflictChurn, .{});
27 try suite.add("smt sat random3 churn bounded store", boundedConflictChurn, .{});
28 try suite.add("smt sat pigeonhole unsat", pigeonholeUnsat, .{});
29 try suite.add("smt sat incremental assumptions", incrementalAssumptions, .{});
30 }
31
32 fn distinctTriple(random: std.Random, variable_count: u32) [3]u32 {
33 assert(variable_count >= 3);
34 const first = random.uintLessThan(u32, variable_count);
35 var second = random.uintLessThan(u32, variable_count - 1);
36 if (second >= first) second += 1;
37 var third = random.uintLessThan(u32, variable_count - 2);
38 const low = @min(first, second);
39 const high = @max(first, second);
40 if (third >= low) third += 1;
41 if (third >= high) third += 1;
42 assert(first != second);
43 assert(second != third);
44 assert(first != third);
45 return .{ first, second, third };
46 }
47
48 fn buildRandomFormula(
49 solver: *Solver,
50 seed: u64,
51 variable_count: u32,
52 clause_count: usize,
53 ) void {
54 assert(variable_count >= 3);
55 var prng = std.Random.DefaultPrng.init(seed);
56 const random = prng.random();
57 for (0..variable_count) |_| {
58 _ = solver.addVariable() catch @panic("smt bench addVariable failed");
59 }
60 for (0..clause_count) |_| {
61 const variables = distinctTriple(random, variable_count);
62 var literals: [3]Literal = undefined;
63 for (&literals, variables) |*literal, variable| {
64 literal.* = Literal.init(variable, random.boolean());
65 }
66 solver.addClause(&literals) catch @panic("smt bench addClause failed");
67 }
68 assert(solver.variableCount() == variable_count);
69 }
70
71 fn pigeonVariable(pigeon: u32, hole: u32) u32 {
72 assert(pigeon < pigeon_count);
73 assert(hole < hole_count);
74 return pigeon * hole_count + hole;
75 }
76
77 fn buildPigeonholeFormula(solver: *Solver) void {
78 for (0..pigeon_count * hole_count) |_| {
79 _ = solver.addVariable() catch @panic("smt bench addVariable failed");
80 }
81 for (0..pigeon_count) |pigeon| {
82 var placement: [hole_count]Literal = undefined;
83 for (&placement, 0..) |*literal, hole| {
84 literal.* = Literal.positive(pigeonVariable(@intCast(pigeon), @intCast(hole)));
85 }
86 solver.addClause(&placement) catch @panic("smt bench addClause failed");
87 }
88 for (0..hole_count) |hole| {
89 for (0..pigeon_count) |first| {
90 for (first + 1..pigeon_count) |second| {
91 solver.addClause(&.{
92 Literal.negative(pigeonVariable(@intCast(first), @intCast(hole))),
93 Literal.negative(pigeonVariable(@intCast(second), @intCast(hole))),
94 }) catch @panic("smt bench addClause failed");
95 }
96 }
97 }
98 assert(solver.variableCount() == pigeon_count * hole_count);
99 }
100
101 fn randomConflictChurn(sample_allocator: std.mem.Allocator) void {
102 var solver = Solver.init(sample_allocator);
103 defer solver.deinit();
104 buildRandomFormula(&solver, churn_seed, churn_variable_count, churn_clause_count);
105 solver.conflict_budget = churn_conflict_budget;
106 const phase = bench.phaseAt("smt.sat.random3.solve", @src());
107 const status = solver.solve() catch @panic("smt bench random3 solve failed");
108 phase.end();
109 bench.coz.progressNamed("smt.sat.random3.complete");
110 const stats = solver.lastSolveStats();
111 assert(status == .unknown);
112 assert(stats.conflicts == churn_conflict_budget);
113 assert(stats.learned_clauses == churn_conflict_budget);
114 std.mem.doNotOptimizeAway(stats.learned_clauses);
115 }
116
117 fn boundedConflictChurn(sample_allocator: std.mem.Allocator) void {
118 var solver = Solver.init(sample_allocator);
119 defer solver.deinit();
120 buildRandomFormula(&solver, churn_seed, churn_variable_count, churn_clause_count);
121 solver.conflict_budget = churn_conflict_budget;
122 solver.max_learned_clauses = bounded_store_limit;
123 const phase = bench.phaseAt("smt.sat.random3.bounded.solve", @src());
124 const status = solver.solve() catch @panic("smt bench bounded churn solve failed");
125 phase.end();
126 bench.coz.progressNamed("smt.sat.random3.bounded.complete");
127 const stats = solver.lastSolveStats();
128 assert(status != .unsat);
129 if (status == .unknown) {
130 assert(stats.conflicts == churn_conflict_budget);
131 }
132 assert(stats.evicted_clauses > 0);
133 assert(solver.replaceableLearnedClauses() <= bounded_store_limit + 1);
134 std.mem.doNotOptimizeAway(stats.evicted_clauses);
135 }
136
137 fn pigeonholeUnsat(sample_allocator: std.mem.Allocator) void {
138 var solver = Solver.init(sample_allocator);
139 defer solver.deinit();
140 buildPigeonholeFormula(&solver);
141 const phase = bench.phaseAt("smt.sat.pigeonhole.solve", @src());
142 const status = solver.solve() catch @panic("smt bench pigeonhole solve failed");
143 phase.end();
144 bench.coz.progressNamed("smt.sat.pigeonhole.complete");
145 const stats = solver.lastSolveStats();
146 assert(status == .unsat);
147 assert(stats.learned_clauses > 0);
148 std.mem.doNotOptimizeAway(stats.conflicts);
149 }
150
151 fn incrementalAssumptions(sample_allocator: std.mem.Allocator) void {
152 var solver = Solver.init(sample_allocator);
153 defer solver.deinit();
154 buildRandomFormula(
155 &solver,
156 incremental_seed,
157 incremental_variable_count,
158 incremental_clause_count,
159 );
160 solver.conflict_budget = incremental_conflict_budget;
161 const phase = bench.phaseAt("smt.sat.incremental.solve", @src());
162 var decided_rounds: usize = 0;
163 for (0..incremental_rounds) |round| {
164 const pivot: u32 = @intCast(round % incremental_variable_count);
165 const partner: u32 = @intCast((round * 7 + 3) % incremental_variable_count);
166 const assumptions = [_]Literal{
167 Literal.init(pivot, round % 2 == 0),
168 Literal.init(partner, round % 3 == 0),
169 };
170 const status = solver.solveWithAssumptions(&assumptions) catch
171 @panic("smt bench incremental solve failed");
172 if (status != .unknown) decided_rounds += 1;
173 bench.coz.progressNamed("smt.sat.incremental.round");
174 }
175 phase.end();
176 assert(decided_rounds <= incremental_rounds);
177 std.mem.doNotOptimizeAway(decided_rounds);
178 }