cdclkit 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
cdclkit/preprocess.py ADDED
@@ -0,0 +1,500 @@
1
+ # SPDX-License-Identifier: Apache-2.0
2
+ # Copyright (c) 2026 Carlo Perassi. Licensed under the Apache License 2.0.
3
+ """SatELite-style preprocessing: subsumption, strengthening, variable elimination.
4
+
5
+ Preprocessing is where the biggest wins on structured industrial instances come
6
+ from, and it is also where correctness gets slippery, because the transformed
7
+ formula is **not logically equivalent** to the original -- only
8
+ *equisatisfiable*. Bounded variable elimination in particular throws away
9
+ clauses whose information is not recoverable from the reduced formula alone.
10
+ So every elimination is recorded on a stack, and :meth:`Preprocessor.reconstruct`
11
+ replays that stack backwards to turn a model of the reduced formula into a
12
+ model of the original. Getting reconstruction wrong is the classic
13
+ preprocessing bug: the solver reports SAT and hands back an assignment that
14
+ does not satisfy the user's input.
15
+
16
+ The techniques, with their justifications:
17
+
18
+ **Unit propagation.** Trivially equivalence preserving.
19
+
20
+ **Subsumption.** ``C`` subsumes ``D`` when ``C`` is a subset of ``D``; then
21
+ ``D`` is implied by ``C`` and can be dropped. Equivalence preserving.
22
+ Implemented with 64-bit signature filtering: a cheap superset test on the
23
+ signature rejects the vast majority of candidate pairs before any set
24
+ operation happens.
25
+
26
+ **Self-subsuming resolution** (strengthening). If ``C \\ {l}`` is a subset of
27
+ ``D`` and ``~l`` is in ``D``, then resolving ``C`` and ``D`` on ``l`` gives a
28
+ clause that subsumes ``D``; so ``~l`` may be deleted from ``D`` in place.
29
+ Equivalence preserving, and it is the engine that makes subsumption keep
30
+ firing: every strengthened clause is a new subsumption candidate.
31
+
32
+ **Bounded variable elimination** (Davis-Putnam, bounded). Replace all clauses
33
+ containing ``v`` by all their non-tautological resolvents on ``v``. The
34
+ resolvents are logically implied, so adding them is sound; removing the
35
+ originals is what breaks equivalence and requires reconstruction. Elimination
36
+ is only performed when it does not increase the clause count (the "bounded"
37
+ part) and when resolvent lengths stay within a cap, otherwise the formula
38
+ explodes -- unbounded Davis-Putnam is exponential, which is exactly why DPLL
39
+ replaced it in 1962.
40
+
41
+ **Pure literal elimination.** If ``v`` occurs with only one polarity, fix it.
42
+ A special case of BVE (the resolvent set is empty), listed separately because
43
+ it is worth a dedicated cheap pass.
44
+
45
+ **Blocked clause elimination.** ``C`` is blocked on ``l in C`` when every
46
+ resolvent of ``C`` on ``l`` is a tautology. Removing a blocked clause
47
+ preserves satisfiability (and, on the reconstruction side, is handled exactly
48
+ like an elimination). This is the operation whose *inverse* -- blocked clause
49
+ addition -- needs the RAT rule rather than RUP, and it is why DRAT has an "A"
50
+ in it.
51
+
52
+ Proof logging
53
+ -------------
54
+ Every clause added is emitted before it is used, every clause removed is
55
+ emitted as a deletion afterwards, in that order. BVE resolvents are RUP (unit
56
+ propagation on the two parents derives them), so a plain DRAT checker verifies
57
+ the entire preprocessing phase with no special support.
58
+ """
59
+
60
+ from __future__ import annotations
61
+
62
+ from typing import Iterable, Sequence
63
+
64
+ from dratify.cnf import CNF
65
+ from dratify.lits import neg
66
+
67
+ __all__ = ["Preprocessor", "PreprocessStats", "preprocess"]
68
+
69
+
70
+ def _sig(lits: Iterable[int]) -> int:
71
+ """A 64-bit occurrence signature used to reject subsumption candidates."""
72
+ s = 0
73
+ for l in lits:
74
+ s |= 1 << ((l >> 1) & 63)
75
+ return s
76
+
77
+
78
+ class PreprocessStats:
79
+ __slots__ = (
80
+ "rounds",
81
+ "units",
82
+ "subsumed",
83
+ "strengthened",
84
+ "eliminated_vars",
85
+ "resolvents",
86
+ "removed_clauses",
87
+ "blocked",
88
+ "pure",
89
+ "tautologies",
90
+ "tried_vars",
91
+ )
92
+
93
+ def __init__(self) -> None:
94
+ for k in PreprocessStats.__slots__:
95
+ setattr(self, k, 0)
96
+
97
+ def as_dict(self) -> dict:
98
+ return {k: getattr(self, k) for k in PreprocessStats.__slots__}
99
+
100
+ def report(self) -> str:
101
+ d = self.as_dict()
102
+ return "\n".join(
103
+ [
104
+ f"c preprocess rounds : {d['rounds']}",
105
+ f"c units propagated : {d['units']}",
106
+ f"c clauses subsumed : {d['subsumed']}",
107
+ f"c literals strengthened : {d['strengthened']}",
108
+ f"c variables eliminated : {d['eliminated_vars']} "
109
+ f"(of {d['tried_vars']} tried, {d['resolvents']} resolvents added)",
110
+ f"c pure literals : {d['pure']}",
111
+ f"c blocked clauses : {d['blocked']}",
112
+ f"c clauses removed total : {d['removed_clauses']}",
113
+ ]
114
+ )
115
+
116
+
117
+ class Preprocessor:
118
+ """Simplifies a :class:`CNF` in place-ish, returning a new reduced formula.
119
+
120
+ Usage::
121
+
122
+ pre = Preprocessor(formula, proof=writer)
123
+ reduced = pre.run()
124
+ ... solve `reduced`, get `model` ...
125
+ full_model = pre.reconstruct(model)
126
+
127
+ ``reconstruct`` is mandatory: the reduced formula's models are, in general,
128
+ *partial* with respect to the original variables.
129
+ """
130
+
131
+ def __init__(
132
+ self,
133
+ formula: CNF,
134
+ proof=None,
135
+ max_resolvent_len: int = 20,
136
+ elim_growth: int = 0,
137
+ do_bve: bool = True,
138
+ do_bce: bool = True,
139
+ subsumption_limit: int = 2_000_000,
140
+ ) -> None:
141
+ self.orig_nvars = formula.nvars
142
+ self.proof = proof
143
+ self.max_resolvent_len = max_resolvent_len
144
+ self.elim_growth = elim_growth
145
+ self.do_bve = do_bve
146
+ self.do_bce = do_bce
147
+ self.subsumption_limit = subsumption_limit
148
+ self.stats = PreprocessStats()
149
+
150
+ # clause store: parallel arrays, index = clause id
151
+ self.cls: list[tuple[int, ...] | None] = [tuple(c) for c in formula.clauses]
152
+ self.sig: list[int] = [_sig(c) for c in self.cls]
153
+ self.occ: list[set[int]] = [set() for _ in range(2 * formula.nvars)]
154
+ for i, c in enumerate(self.cls):
155
+ for l in c:
156
+ self.occ[l].add(i)
157
+
158
+ self.value: dict[int, bool] = {} # var -> fixed value
159
+ self.eliminated: list[tuple[int, list[tuple[int, ...]]]] = []
160
+ self.frozen: set[int] = set() # variables that must survive
161
+ self.unsat = False
162
+ self._touched: set[int] = set(range(2 * formula.nvars))
163
+
164
+ # -- freezing -----------------------------------------------------------
165
+
166
+ def freeze(self, variables: Iterable[int]) -> None:
167
+ """Protect variables from elimination (needed for incremental use or
168
+ when the caller wants to read their value out of the model)."""
169
+ self.frozen.update(variables)
170
+
171
+ # -- low-level clause ops ----------------------------------------------
172
+
173
+ def _emit_add(self, lits: Sequence[int]) -> None:
174
+ if self.proof is not None:
175
+ self.proof.add(lits)
176
+
177
+ def _emit_del(self, lits: Sequence[int]) -> None:
178
+ if self.proof is not None:
179
+ self.proof.delete(lits)
180
+
181
+ def _add_clause(self, lits: Sequence[int], log: bool = True) -> int:
182
+ c = tuple(lits)
183
+ if log:
184
+ self._emit_add(c)
185
+ i = len(self.cls)
186
+ self.cls.append(c)
187
+ self.sig.append(_sig(c))
188
+ for l in c:
189
+ self.occ[l].add(i)
190
+ self._touched.add(l)
191
+ return i
192
+
193
+ def _remove_clause(self, i: int, log: bool = True) -> None:
194
+ c = self.cls[i]
195
+ if c is None:
196
+ return
197
+ for l in c:
198
+ self.occ[l].discard(i)
199
+ self._touched.add(l)
200
+ self.cls[i] = None
201
+ if log:
202
+ self._emit_del(c)
203
+ self.stats.removed_clauses += 1
204
+
205
+ def _replace_clause(self, i: int, lits: Sequence[int]) -> None:
206
+ """Strengthen clause ``i`` to ``lits`` (a proper subset)."""
207
+ old = self.cls[i]
208
+ self._emit_add(lits)
209
+ for l in old:
210
+ self.occ[l].discard(i)
211
+ self._touched.add(l)
212
+ self.cls[i] = tuple(lits)
213
+ self.sig[i] = _sig(lits)
214
+ for l in lits:
215
+ self.occ[l].add(i)
216
+ self._touched.add(l)
217
+ self._emit_del(old)
218
+
219
+ def alive(self) -> Iterable[int]:
220
+ return (i for i, c in enumerate(self.cls) if c is not None)
221
+
222
+ # -- unit propagation ---------------------------------------------------
223
+
224
+ def propagate(self) -> bool:
225
+ """Propagate unit clauses to fixpoint. False means UNSAT."""
226
+ queue = [self.cls[i][0] for i in self.alive() if len(self.cls[i]) == 1]
227
+ while queue:
228
+ l = queue.pop()
229
+ v, positive = l >> 1, not (l & 1)
230
+ if v in self.value:
231
+ if self.value[v] != positive:
232
+ self.unsat = True
233
+ self._emit_add(())
234
+ return False
235
+ continue
236
+ self.value[v] = positive
237
+ self.stats.units += 1
238
+ # clauses containing l are satisfied
239
+ for i in list(self.occ[l]):
240
+ self._remove_clause(i)
241
+ # ~l is removed from the clauses that contain it
242
+ for i in list(self.occ[l ^ 1]):
243
+ c = self.cls[i]
244
+ if c is None:
245
+ continue
246
+ rest = tuple(x for x in c if x != (l ^ 1))
247
+ if not rest:
248
+ self.unsat = True
249
+ self._emit_add(())
250
+ return False
251
+ self._replace_clause(i, rest)
252
+ if len(rest) == 1:
253
+ queue.append(rest[0])
254
+ return True
255
+
256
+ # -- subsumption --------------------------------------------------------
257
+
258
+ def subsume(self) -> None:
259
+ """Remove subsumed clauses and strengthen by self-subsuming resolution."""
260
+ work = sorted(self.alive(), key=lambda i: len(self.cls[i]))
261
+ budget = self.subsumption_limit
262
+ for i in work:
263
+ c = self.cls[i]
264
+ if c is None:
265
+ continue
266
+ # pick the literal with the smallest occurrence list to scan
267
+ best = min(c, key=lambda l: len(self.occ[l]) + len(self.occ[l ^ 1]))
268
+ si = self.sig[i]
269
+ cset = set(c)
270
+ for pol in (best, best ^ 1):
271
+ for j in list(self.occ[pol]):
272
+ if j == i:
273
+ continue
274
+ d = self.cls[j]
275
+ if d is None or len(d) < len(c):
276
+ continue
277
+ budget -= 1
278
+ if budget < 0:
279
+ return
280
+ if si & ~self.sig[j]:
281
+ continue
282
+ dset = set(d)
283
+ if cset <= dset:
284
+ self._remove_clause(j)
285
+ self.stats.subsumed += 1
286
+ continue
287
+ # self-subsuming resolution: C\{l} subset of D and ~l in D
288
+ diff = cset - dset
289
+ if len(diff) == 1:
290
+ l = next(iter(diff))
291
+ if (l ^ 1) in dset:
292
+ new = tuple(x for x in d if x != (l ^ 1))
293
+ if not new:
294
+ self.unsat = True
295
+ self._emit_add(())
296
+ return
297
+ self._replace_clause(j, new)
298
+ self.stats.strengthened += 1
299
+
300
+ # -- blocked clause elimination ----------------------------------------
301
+
302
+ def _resolvent(self, c: Sequence[int], d: Sequence[int], l: int) -> list[int] | None:
303
+ """Resolve on ``l`` (in ``c``); None when the resolvent is a tautology."""
304
+ out = [x for x in c if x != l]
305
+ seen = set(out)
306
+ for x in d:
307
+ if x == (l ^ 1):
308
+ continue
309
+ if (x ^ 1) in seen:
310
+ return None
311
+ if x not in seen:
312
+ seen.add(x)
313
+ out.append(x)
314
+ return out
315
+
316
+ def block_eliminate(self) -> None:
317
+ """Remove clauses that are blocked on one of their literals."""
318
+ for i in list(self.alive()):
319
+ c = self.cls[i]
320
+ if c is None or len(c) > self.max_resolvent_len:
321
+ continue
322
+ for l in c:
323
+ if (l >> 1) in self.frozen:
324
+ continue
325
+ blocked = True
326
+ for j in self.occ[l ^ 1]:
327
+ d = self.cls[j]
328
+ if d is None:
329
+ continue
330
+ if self._resolvent(c, d, l) is not None:
331
+ blocked = False
332
+ break
333
+ if blocked and self.occ[l ^ 1]:
334
+ self.eliminated.append((l >> 1, [c]))
335
+ self._remove_clause(i)
336
+ self.stats.blocked += 1
337
+ break
338
+
339
+ # -- variable elimination ----------------------------------------------
340
+
341
+ def _pure_literals(self) -> None:
342
+ for v in range(self.orig_nvars):
343
+ if v in self.value or v in self.frozen:
344
+ continue
345
+ pos, negs = self.occ[v << 1], self.occ[(v << 1) | 1]
346
+ if pos and not negs:
347
+ self._eliminate_pure(v, True)
348
+ elif negs and not pos:
349
+ self._eliminate_pure(v, False)
350
+
351
+ def _eliminate_pure(self, v: int, positive: bool) -> None:
352
+ lit = (v << 1) | (0 if positive else 1)
353
+ clauses = [self.cls[i] for i in list(self.occ[lit]) if self.cls[i] is not None]
354
+ if not clauses:
355
+ return
356
+ self.eliminated.append((v, clauses))
357
+ for i in list(self.occ[lit]):
358
+ self._remove_clause(i)
359
+ self.stats.pure += 1
360
+ self.stats.eliminated_vars += 1
361
+
362
+ def eliminate_vars(self) -> bool:
363
+ """Bounded variable elimination. False means UNSAT was derived."""
364
+ order = sorted(
365
+ (v for v in range(self.orig_nvars) if v not in self.value and v not in self.frozen),
366
+ key=lambda v: len(self.occ[v << 1]) * len(self.occ[(v << 1) | 1]),
367
+ )
368
+ for v in order:
369
+ if v in self.value:
370
+ continue
371
+ pos = [self.cls[i] for i in self.occ[v << 1] if self.cls[i] is not None]
372
+ negs = [self.cls[i] for i in self.occ[(v << 1) | 1] if self.cls[i] is not None]
373
+ self.stats.tried_vars += 1
374
+ if not pos and not negs:
375
+ continue
376
+ if len(pos) * len(negs) > 400:
377
+ continue # cheap guard against quadratic blowup
378
+ resolvents = []
379
+ too_big = False
380
+ for c in pos:
381
+ for d in negs:
382
+ r = self._resolvent(c, d, v << 1)
383
+ if r is None:
384
+ self.stats.tautologies += 1
385
+ continue
386
+ if not r:
387
+ # empty resolvent: the formula is unsatisfiable
388
+ self._emit_add(())
389
+ self.unsat = True
390
+ return False
391
+ if len(r) > self.max_resolvent_len:
392
+ too_big = True
393
+ break
394
+ resolvents.append(r)
395
+ if too_big:
396
+ break
397
+ if too_big or len(resolvents) > len(pos) + len(negs) + self.elim_growth:
398
+ continue
399
+ # commit: add resolvents first (they are RUP given the parents),
400
+ # then delete the parents
401
+ for r in resolvents:
402
+ self._add_clause(r)
403
+ self.stats.resolvents += 1
404
+ self.eliminated.append((v, pos + negs))
405
+ for i in list(self.occ[v << 1]) + list(self.occ[(v << 1) | 1]):
406
+ self._remove_clause(i)
407
+ self.stats.eliminated_vars += 1
408
+ return True
409
+
410
+ # -- driver -------------------------------------------------------------
411
+
412
+ def run(self, rounds: int = 3) -> CNF:
413
+ """Run the simplification loop and return the reduced formula."""
414
+ for _ in range(rounds):
415
+ self.stats.rounds += 1
416
+ before = (self.stats.subsumed, self.stats.strengthened, self.stats.eliminated_vars)
417
+ if not self.propagate():
418
+ break
419
+ self.subsume()
420
+ if self.unsat:
421
+ break
422
+ self._pure_literals()
423
+ if self.do_bce:
424
+ self.block_eliminate()
425
+ if self.do_bve and not self.eliminate_vars():
426
+ break
427
+ if not self.propagate():
428
+ break
429
+ after = (self.stats.subsumed, self.stats.strengthened, self.stats.eliminated_vars)
430
+ if after == before:
431
+ break
432
+ return self.to_cnf()
433
+
434
+ def to_cnf(self) -> CNF:
435
+ out = CNF(self.orig_nvars)
436
+ if self.unsat:
437
+ out.add([])
438
+ return out
439
+ for i in self.alive():
440
+ out.add(self.cls[i])
441
+ out.nvars = self.orig_nvars
442
+ return out
443
+
444
+ # -- model reconstruction ----------------------------------------------
445
+
446
+ def reconstruct(self, model: Sequence[bool]) -> list[bool]:
447
+ """Extend a model of the reduced formula to the original variables.
448
+
449
+ Walk the elimination stack in reverse. For each recorded variable, if
450
+ any of its stored clauses is currently unsatisfied, flip the variable
451
+ to the polarity that satisfies it -- this always works, because every
452
+ stored clause containing the opposite polarity was already accounted
453
+ for by the resolvents that remain in the reduced formula.
454
+ """
455
+ full = [False] * self.orig_nvars
456
+ for v in range(min(len(model), self.orig_nvars)):
457
+ full[v] = model[v]
458
+ for v, val in self.value.items():
459
+ full[v] = val
460
+
461
+ def sat(clause: Sequence[int]) -> bool:
462
+ return any(full[l >> 1] != bool(l & 1) for l in clause)
463
+
464
+ for v, clauses in reversed(self.eliminated):
465
+ unsat_clauses = [c for c in clauses if not sat(c)]
466
+ if not unsat_clauses:
467
+ continue
468
+ # Every unsatisfied clause must contain the *same* polarity of v --
469
+ # that is the invariant the elimination establishes. If clauses of
470
+ # both polarities were unsatisfied, their resolvent on v would be
471
+ # unsatisfied too, and that resolvent is still in the reduced
472
+ # formula, contradicting the fact that `model` satisfies it.
473
+ pos = any((v << 1) in c for c in unsat_clauses)
474
+ neg_ = any(((v << 1) | 1) in c for c in unsat_clauses)
475
+ if pos and neg_:
476
+ raise AssertionError(
477
+ f"reconstruction invariant violated for x{v}: clauses of both "
478
+ "polarities are unsatisfied; the elimination stack is corrupt"
479
+ )
480
+ full[v] = pos
481
+ if not all(sat(c) for c in clauses):
482
+ raise AssertionError(
483
+ f"reconstruction failed for x{v}: flipping it did not satisfy "
484
+ "its stored clauses"
485
+ )
486
+ return full
487
+
488
+ # -- reporting ----------------------------------------------------------
489
+
490
+ def summary(self, reduced: CNF) -> str:
491
+ return (
492
+ f"c preprocessing: {self.orig_nvars} vars / {len(self.cls)} clause slots "
493
+ f"-> {reduced.nvars} vars / {reduced.nclauses} clauses\n" + self.stats.report()
494
+ )
495
+
496
+
497
+ def preprocess(formula: CNF, proof=None, **kw) -> tuple[CNF, Preprocessor]:
498
+ """Convenience wrapper: returns ``(reduced_formula, preprocessor)``."""
499
+ pre = Preprocessor(formula, proof=proof, **kw)
500
+ return pre.run(), pre