cattykit 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.
cattykit/__init__.py ADDED
@@ -0,0 +1,10 @@
1
+ from .experiments import run_experiment
2
+ from .models import available_models, install_model, load_model, model_info
3
+
4
+ __all__ = [
5
+ "available_models",
6
+ "install_model",
7
+ "load_model",
8
+ "model_info",
9
+ "run_experiment",
10
+ ]
@@ -0,0 +1,99 @@
1
+ from collections import Counter
2
+ from dataclasses import dataclass
3
+ from pathlib import Path
4
+ from datetime import datetime
5
+ import sys
6
+
7
+
8
+ from .logging import SQLiteLogger
9
+ from .models import load_model
10
+
11
+
12
+ @dataclass
13
+ class ExperimentResult:
14
+ distributions: dict[str, Counter[str]]
15
+ logging_db: str
16
+
17
+
18
+ def run_experiment(
19
+ model_name: str,
20
+ problems: list[str],
21
+ iterations: int,
22
+ logging_db: str | Path | None = None,
23
+ verbose: bool = True,
24
+ ) -> ExperimentResult:
25
+ """Runs the model on each problem for the required number of iterations.
26
+ The iteration count is used as the random seed given to the model.
27
+ The same random seeds are used for each problem.
28
+ If logging_db is not set, the model name and run date/time are used.
29
+ Counters of answers and the logging_db are returned in ExperimentResult.
30
+ If verbose, problems and solutions will print to stdout."""
31
+ interactive = verbose and sys.stdout.isatty()
32
+ if logging_db is None:
33
+ now = datetime.now().strftime("%Y%m%d-%H%M%S")
34
+ logging_db = f"{model_name}-experiments-{now}.sqlite"
35
+ distributions: dict[str, Counter[str]] = {}
36
+ logger = SQLiteLogger(logging_db)
37
+ try:
38
+ for problem in problems:
39
+ if interactive:
40
+ print(problem)
41
+ display = RunningAnswerDisplay()
42
+ answers: Counter[str] = Counter()
43
+ for seed in range(iterations):
44
+ model = load_model(
45
+ model_name,
46
+ config={"seed": seed},
47
+ logger=logger,
48
+ )
49
+ try:
50
+ answer = model.solve(problem)
51
+ answers[answer] += 1
52
+ if interactive:
53
+ # print(f"{seed}: {answer}")
54
+ display.update(answer)
55
+ finally:
56
+ model.close()
57
+ distributions[problem] = answers
58
+ finally:
59
+ logger.close()
60
+ return ExperimentResult(
61
+ distributions=distributions,
62
+ logging_db=str(logging_db),
63
+ )
64
+
65
+
66
+ def total_variation_distance(
67
+ observed: dict[str, int], expected: dict[str, int]
68
+ ) -> float:
69
+ observed_total = sum(observed.values())
70
+ expected_total = sum(expected.values())
71
+ answers = observed.keys() | expected.keys()
72
+ return 0.5 * sum(
73
+ abs(
74
+ observed.get(answer, 0) / observed_total
75
+ - expected.get(answer, 0) / expected_total
76
+ )
77
+ for answer in answers
78
+ )
79
+
80
+
81
+ class RunningAnswerDisplay:
82
+ def __init__(self):
83
+ self.counts = Counter()
84
+ self.lines_printed = 0
85
+ self.total = 0
86
+
87
+ def update(self, answer: str) -> None:
88
+ self.counts[answer] += 1
89
+ self.total += 1
90
+
91
+ if self.lines_printed:
92
+ # Move cursor back up over the previously printed answer lines.
93
+ print(f"\033[{self.lines_printed}A", end="")
94
+
95
+ for answer, count in self.counts.most_common():
96
+ # Clear line, then rewrite it.
97
+ print(f"\033[2K{answer}: {count} / {self.total}")
98
+
99
+ self.lines_printed = len(self.counts)
@@ -0,0 +1,9 @@
1
+ """Structured logging for CattyKit model runs."""
2
+
3
+ from .event import ModelEvent
4
+ from .model_logger import ModelLogger
5
+ from .null_logger import NullLogger
6
+ from .print_logger import PrintLogger
7
+ from .sqlite_logger import SQLiteLogger
8
+
9
+ __all__ = ["ModelEvent", "ModelLogger", "NullLogger", "PrintLogger", "SQLiteLogger"]
@@ -0,0 +1,32 @@
1
+ """Structured events emitted by CattyKit models."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from collections.abc import Mapping
6
+ from dataclasses import dataclass
7
+ from datetime import UTC, datetime
8
+ from typing import Any
9
+
10
+
11
+ @dataclass(frozen=True, slots=True)
12
+ class ModelEvent:
13
+ """A structured event emitted during a model run."""
14
+
15
+ timestamp: datetime
16
+ model: str
17
+ kind: str
18
+ data: Mapping[str, Any]
19
+
20
+ @classmethod
21
+ def create(cls, model: str, kind: str, **data: Any) -> ModelEvent:
22
+ """Create an event timestamped in UTC."""
23
+ return cls(timestamp=datetime.now(UTC), model=model, kind=kind, data=data)
24
+
25
+ def as_dict(self) -> dict[str, Any]:
26
+ """Return a JSON-ready representation of the event."""
27
+ return {
28
+ "timestamp": self.timestamp.isoformat(),
29
+ "model": self.model,
30
+ "kind": self.kind,
31
+ "data": dict(self.data),
32
+ }
@@ -0,0 +1,18 @@
1
+ """Protocol implemented by CattyKit loggers."""
2
+
3
+ from typing import Protocol, runtime_checkable
4
+
5
+ from .event import ModelEvent
6
+
7
+
8
+ @runtime_checkable
9
+ class ModelLogger(Protocol):
10
+ """Receives structured events emitted by a model."""
11
+
12
+ def log(self, event: ModelEvent) -> None:
13
+ """Record an event."""
14
+ ...
15
+
16
+ def close(self) -> None:
17
+ """Release logger resources."""
18
+ ...
@@ -0,0 +1,13 @@
1
+ """No-op model logger."""
2
+
3
+ from .event import ModelEvent
4
+
5
+
6
+ class NullLogger:
7
+ """A logger that discards all events."""
8
+
9
+ def log(self, event: ModelEvent) -> None:
10
+ """Discard an event."""
11
+
12
+ def close(self) -> None:
13
+ """Release no resources."""
@@ -0,0 +1,24 @@
1
+ """JSON-lines model logger."""
2
+
3
+ import json
4
+ import sys
5
+ from typing import TextIO
6
+
7
+ from .event import ModelEvent
8
+
9
+
10
+ class PrintLogger:
11
+ """A logger that writes one JSON event per line to a text stream."""
12
+
13
+ def __init__(self, stream: TextIO | None = None) -> None:
14
+ self._stream = stream if stream is not None else sys.stdout
15
+
16
+ def log(self, event: ModelEvent) -> None:
17
+ """Print an event as a JSON line."""
18
+ print(
19
+ json.dumps(event.as_dict(), sort_keys=True), file=self._stream, flush=True
20
+ )
21
+
22
+ def close(self) -> None:
23
+ """Flush, but do not close, the caller-owned stream."""
24
+ self._stream.flush()