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 }