tiny-datalog 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.
@@ -0,0 +1,383 @@
1
+ #!/usr/bin/env python3
2
+ """
3
+ semiring.py — Datalog over semirings: provenance and recursive aggregation.
4
+
5
+ A semiring (S, +, x, 0, 1) generalises what a derivation *carries*. Every
6
+ fact holds a value from S; a rule instance multiplies the values of its
7
+ body facts; alternative derivations of the same fact add. Choosing the
8
+ semiring changes the question the same program answers:
9
+
10
+ bool (or, and) plain Datalog — is the fact derivable?
11
+ minplus (min, +) cheapest derivation — shortest paths
12
+ count (+, x) how many distinct derivations?
13
+ why (set union) which base facts support it? (minimal witnesses)
14
+ viterbi (max, x) probability of the most likely derivation
15
+
16
+ Facts take weights with `edge(a, b) @ 3.`; unweighted facts get the
17
+ semiring's `one`. Run semirings on ordinary programs, not on magic-sets
18
+ rewritings: the magic guard predicates would multiply into every product
19
+ and corrupt the values.
20
+
21
+ Evaluation is Kleene fixpoint iteration of the immediate-consequence
22
+ operator — correct for the omega-continuous semirings provided here, with
23
+ a round cap that catches genuinely divergent combinations (counting
24
+ derivations in a cyclic graph really is infinite; so is the cheapest
25
+ route round a negative-cost cycle). Positive programs
26
+ only: negation over semiring values needs more theory than this file
27
+ carries.
28
+
29
+ This is the entry point to the "Datalog over semirings" research thread:
30
+ Green–Karvounarakis–Tannen's provenance semirings (PODS 2007) and the
31
+ convergence theory of Abo Khamis–Ngo–Suciu and colleagues.
32
+
33
+ CLI
34
+ ---
35
+ python3 semiring.py --semiring minplus programs/routes.dl
36
+ python3 semiring.py --semiring why -q 'path(a, e)' programs/routes.dl
37
+ """
38
+
39
+ from __future__ import annotations
40
+
41
+ import argparse
42
+ import sys
43
+
44
+ from tiny_datalog.datalog import (
45
+ Const, DatalogError, Program, format_atom, parse, read_program,
46
+ validate, _aggregate_of, _match, _sort_key)
47
+
48
+
49
+ # ---------------------------------------------------------------------------
50
+ # Semirings
51
+ # ---------------------------------------------------------------------------
52
+
53
+ class Semiring:
54
+ """Interface: zero, one, plus, times, fact_value, fmt — and
55
+ `divergence`, the reason to give when values never settle."""
56
+ name = "abstract"
57
+ zero = None
58
+ one = None
59
+ divergence = "its values keep changing forever"
60
+
61
+ def plus(self, a, b):
62
+ raise NotImplementedError
63
+
64
+ def times(self, a, b):
65
+ raise NotImplementedError
66
+
67
+ def fact_value(self, pred, tup, weight):
68
+ """Value of a base fact; `weight` is its @ annotation or None."""
69
+ return self.one if weight is None else weight
70
+
71
+ def fmt(self, v):
72
+ return str(v)
73
+
74
+
75
+ class BoolSemiring(Semiring):
76
+ """(or, and): recovers ordinary Datalog."""
77
+ name = "bool"
78
+ zero, one = False, True
79
+
80
+ def plus(self, a, b):
81
+ return a or b
82
+
83
+ def times(self, a, b):
84
+ return a and b
85
+
86
+ def fact_value(self, pred, tup, weight):
87
+ return True
88
+
89
+ def fmt(self, v):
90
+ return "true" if v else "false"
91
+
92
+
93
+ class MinPlusSemiring(Semiring):
94
+ """(min, +), the tropical semiring: cost of the cheapest derivation.
95
+ Unweighted facts cost 0."""
96
+ name = "minplus"
97
+ zero, one = float("inf"), 0
98
+ divergence = ("a negative-cost cycle makes costs fall without bound, "
99
+ "so there is no cheapest derivation")
100
+
101
+ def plus(self, a, b):
102
+ return min(a, b)
103
+
104
+ def times(self, a, b):
105
+ return a + b
106
+
107
+ def fmt(self, v):
108
+ return "%g" % v
109
+
110
+
111
+ class CountSemiring(Semiring):
112
+ """(+, x) over the naturals: number of distinct derivations. Weights
113
+ are ignored — every base fact counts once, however often or with
114
+ whatever weight it is written, and so does every rule: a program is a
115
+ *set* of clauses, and saying one twice adds no new way to derive
116
+ anything. (Reading weights as multiplicities would give bag
117
+ semantics; see lessons/08-semirings.md.) Diverges when derivations
118
+ are unbounded (cycles) — by design."""
119
+ name = "count"
120
+ zero, one = 0, 1
121
+ divergence = ("a cycle among the derivations can be pumped forever, "
122
+ "so there are infinitely many of them to count")
123
+
124
+ def plus(self, a, b):
125
+ return a + b
126
+
127
+ def times(self, a, b):
128
+ return a * b
129
+
130
+ def fact_value(self, pred, tup, weight):
131
+ return 1
132
+
133
+
134
+ class ViterbiSemiring(Semiring):
135
+ """(max, x) over [0, 1]: probability of the most likely single
136
+ derivation. See lessons/09-probabilistic.md for why this — and not
137
+ "add up the probabilities" — is the honest semiring."""
138
+ name = "viterbi"
139
+ zero, one = 0.0, 1.0
140
+
141
+ def plus(self, a, b):
142
+ return max(a, b)
143
+
144
+ def times(self, a, b):
145
+ return a * b
146
+
147
+ def fact_value(self, pred, tup, weight):
148
+ if weight is None:
149
+ return 1.0
150
+ if not 0 <= weight <= 1:
151
+ # outside [0, 1] a "probability" is meaningless, and max no
152
+ # longer tames cycles (a loop of weight 2 grows forever)
153
+ raise DatalogError(
154
+ "viterbi weights are probabilities and must lie in "
155
+ "[0, 1]: %s @ %s" % (format_atom(pred, tup), weight))
156
+ return float(weight)
157
+
158
+ def fmt(self, v):
159
+ return "%.4g" % v
160
+
161
+
162
+ def _minimal(sets):
163
+ """Keep only the minimal witness sets (absorption: A + A.B = A)."""
164
+ return frozenset(s for s in sets
165
+ if not any(o < s for o in sets))
166
+
167
+
168
+ class WhySemiring(Semiring):
169
+ """Why-provenance: each value is a set of minimal witness sets — the
170
+ alternative sets of base facts sufficient to derive the fact."""
171
+ name = "why"
172
+ zero = frozenset()
173
+ one = frozenset([frozenset()])
174
+
175
+ def plus(self, a, b):
176
+ return _minimal(a | b)
177
+
178
+ def times(self, a, b):
179
+ return _minimal(frozenset(x | y for x in a for y in b))
180
+
181
+ def fact_value(self, pred, tup, weight):
182
+ return frozenset([frozenset([format_atom(pred, tup)])])
183
+
184
+ def fmt(self, v):
185
+ parts = sorted("{%s}" % ", ".join(sorted(w)) for w in v)
186
+ return " | ".join(parts) if parts else "{}"
187
+
188
+
189
+ SEMIRINGS = {sr.name: sr for sr in (
190
+ BoolSemiring(), MinPlusSemiring(), CountSemiring(),
191
+ ViterbiSemiring(), WhySemiring())}
192
+
193
+
194
+ # ---------------------------------------------------------------------------
195
+ # Evaluation: Kleene iteration of the immediate-consequence operator
196
+ # ---------------------------------------------------------------------------
197
+
198
+ def _eval_rule(rule, rels, sr):
199
+ """Yield (head_tuple, value): every instantiation of the rule, with
200
+ the product of its body facts' values."""
201
+ pairs = [({}, sr.one)]
202
+ for lit in rule.body:
203
+ rel = rels.get(lit.atom.pred, {})
204
+ args = lit.atom.args
205
+ new = []
206
+ for s, v in pairs:
207
+ for tup, val in rel.items():
208
+ m = _match(args, tup, s)
209
+ if m is not None:
210
+ new.append((m, sr.times(v, val)))
211
+ pairs = new
212
+ if not pairs:
213
+ return
214
+ for s, v in pairs:
215
+ yield (tuple(a.value if isinstance(a, Const) else s[a.name]
216
+ for a in rule.head.args), v)
217
+
218
+
219
+ class SemiringEngine:
220
+ """Evaluates a positive program over a semiring. After run(), `rels`
221
+ maps each predicate to {tuple: value} and `rounds` records how many
222
+ Kleene rounds the fixpoint took."""
223
+
224
+ def __init__(self, clauses, semiring):
225
+ if isinstance(semiring, str):
226
+ try:
227
+ semiring = SEMIRINGS[semiring]
228
+ except KeyError:
229
+ # the CLI is guarded by argparse choices; this is the
230
+ # library caller's path, and a bare KeyError names the
231
+ # typo without saying what would have been right
232
+ raise DatalogError(
233
+ "unknown semiring %r — choose one of %s"
234
+ % (semiring, ", ".join(sorted(SEMIRINGS))))
235
+ self.sr = semiring
236
+ self.arity = validate(clauses)
237
+ for r in clauses:
238
+ if r.retract:
239
+ # `q~.` is an update, not a fact; let the base engine's
240
+ # Program reject it with its own explanation
241
+ Program([r])
242
+ if _aggregate_of(r.head):
243
+ raise DatalogError(
244
+ "semiring evaluation does not compose with head "
245
+ "aggregation (a semiring already IS the aggregation "
246
+ "— see lessons 8 and 13): %s" % r)
247
+ for lit in r.body:
248
+ if lit.negated:
249
+ raise DatalogError(
250
+ "semiring evaluation is defined for positive "
251
+ "programs only (negated literal `%s` in: %s)"
252
+ % (lit, r))
253
+ # A program is a set of clauses: a rule written twice is one rule
254
+ # (otherwise count would see two derivations where there is one).
255
+ self.rules = list(dict.fromkeys(r for r in clauses if r.body))
256
+ self.idb = {r.head.pred for r in self.rules}
257
+ self.base = {}
258
+ seen_facts = set()
259
+ for r in clauses:
260
+ if r.body:
261
+ continue
262
+ tup = tuple(a.value for a in r.head.args)
263
+ v = semiring.fact_value(r.head.pred, tup, r.weight)
264
+ # Likewise a fact written twice with the same value is one
265
+ # fact. Keying on the *value*, not the written weight, is
266
+ # what makes count ignore weights consistently.
267
+ if (r.head.pred, tup, v) in seen_facts:
268
+ continue
269
+ seen_facts.add((r.head.pred, tup, v))
270
+ d = self.base.setdefault(r.head.pred, {})
271
+ # distinct values for the same tuple (parallel edges) combine
272
+ d[tup] = semiring.plus(d[tup], v) if tup in d else v
273
+ self.rels = {}
274
+ self.rounds = 0
275
+
276
+ def run(self, max_rounds=200):
277
+ if max_rounds < 1:
278
+ raise DatalogError(
279
+ "max_rounds must be at least 1, not %r" % (max_rounds,))
280
+ sr = self.sr
281
+ current = {p: dict(d) for p, d in self.base.items()}
282
+ for n in range(1, max_rounds + 1):
283
+ new = {p: dict(d) for p, d in self.base.items()}
284
+ for rule in self.rules:
285
+ d = new.setdefault(rule.head.pred, {})
286
+ for tup, val in _eval_rule(rule, current, sr):
287
+ d[tup] = sr.plus(d[tup], val) if tup in d else val
288
+ new = {p: {t: v for t, v in d.items() if v != sr.zero}
289
+ for p, d in new.items()}
290
+ new = {p: d for p, d in new.items() if d}
291
+ if new == current:
292
+ self.rels = current
293
+ self.rounds = n
294
+ return self
295
+ current = new
296
+ # Can we tell divergence from a small budget? Round n accounts
297
+ # for every derivation of height <= n. A derivation that repeats
298
+ # no fact on any branch has height below the number of facts, so
299
+ # once the rounds outnumber the facts, values that still change
300
+ # can only come from pumping a cycle — and that never stops.
301
+ nfacts = sum(len(d) for d in current.values())
302
+ if max_rounds > nfacts:
303
+ raise DatalogError(
304
+ "no fixpoint over the %s semiring: values still change "
305
+ "after %d rounds although only %d facts exist, so no "
306
+ "round budget will do — %s"
307
+ % (sr.name, max_rounds, nfacts, sr.divergence))
308
+ raise DatalogError(
309
+ "no fixpoint after %d rounds over the %s semiring — with %d "
310
+ "facts derived so far that is too few rounds to tell a deep "
311
+ "derivation from a divergent one; raise it with --max-rounds"
312
+ % (max_rounds, sr.name, nfacts))
313
+
314
+ def value(self, pred, tup):
315
+ return self.rels.get(pred, {}).get(tup, self.sr.zero)
316
+
317
+
318
+ def run_semiring(text, semiring, max_rounds=200):
319
+ """Parse and evaluate a program over the named (or given) semiring."""
320
+ return SemiringEngine(parse(text), semiring).run(max_rounds)
321
+
322
+
323
+ # ---------------------------------------------------------------------------
324
+ # CLI
325
+ # ---------------------------------------------------------------------------
326
+
327
+ def _positive_int(text):
328
+ n = int(text)
329
+ if n < 1:
330
+ raise argparse.ArgumentTypeError("must be at least 1, not %d" % n)
331
+ return n
332
+
333
+
334
+ def main(argv=None):
335
+ ap = argparse.ArgumentParser(
336
+ description="Evaluate a positive Datalog program over a semiring.")
337
+ ap.add_argument("file", help="Datalog program (.dl); facts may carry "
338
+ "`@ weight` annotations")
339
+ ap.add_argument("-s", "--semiring", default="minplus",
340
+ choices=sorted(SEMIRINGS),
341
+ help="semiring to evaluate over (default: minplus)")
342
+ ap.add_argument("-q", "--query", action="append", default=[],
343
+ metavar="ATOM", help="query atom (repeatable)")
344
+ ap.add_argument("--max-rounds", type=_positive_int, default=200,
345
+ help="Kleene round budget before giving up (default 200)")
346
+ args = ap.parse_args(argv)
347
+
348
+ try:
349
+ engine = run_semiring(read_program(args.file),
350
+ args.semiring, args.max_rounds)
351
+ except DatalogError as exc:
352
+ print("error: %s" % exc, file=sys.stderr)
353
+ return 1
354
+
355
+ sr = engine.sr
356
+ print("Semiring: %s (fixpoint after %d rounds)" % (sr.name, engine.rounds))
357
+ if args.query:
358
+ from tiny_datalog.datalog import _parse_query_atom
359
+ for q in args.query:
360
+ try:
361
+ atom = _parse_query_atom(q, engine.arity)
362
+ except DatalogError as exc:
363
+ print("error: %s" % exc, file=sys.stderr)
364
+ return 1
365
+ print("?- %s" % atom)
366
+ hits = [(t, v) for t, v in engine.rels.get(atom.pred, {}).items()
367
+ if _match(atom.args, t, {}) is not None]
368
+ for t, v in sorted(hits, key=lambda x: _sort_key(x[0])):
369
+ print(" %s = %s" % (format_atom(atom.pred, t), sr.fmt(v)))
370
+ print(" (%d answer%s)" % (len(hits), "" if len(hits) == 1 else "s"))
371
+ return 0
372
+ for pred in sorted(engine.idb):
373
+ d = engine.rels.get(pred, {})
374
+ print("%% %s/%d — %d fact%s" % (pred, engine.arity[pred], len(d),
375
+ "" if len(d) == 1 else "s"))
376
+ for tup in sorted(d, key=_sort_key):
377
+ print("%s = %s" % (format_atom(pred, tup), sr.fmt(d[tup])))
378
+ print()
379
+ return 0
380
+
381
+
382
+ if __name__ == "__main__":
383
+ sys.exit(main())