blut-core 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.
Files changed (57) hide show
  1. blut_core/__init__.py +54 -0
  2. blut_core/async_io.py +246 -0
  3. blut_core/checkpoint.py +172 -0
  4. blut_core/evaluator.py +98 -0
  5. blut_core/ingredients/__init__.py +20 -0
  6. blut_core/ingredients/_specs.py +2 -0
  7. blut_core/ingredients/checkpoint/__init__.py +0 -0
  8. blut_core/ingredients/checkpoint/_specs.py +162 -0
  9. blut_core/ingredients/data/__init__.py +0 -0
  10. blut_core/ingredients/data/_specs.py +87 -0
  11. blut_core/ingredients/ema/__init__.py +0 -0
  12. blut_core/ingredients/ema/_specs.py +62 -0
  13. blut_core/ingredients/eval/__init__.py +0 -0
  14. blut_core/ingredients/eval/_specs.py +150 -0
  15. blut_core/ingredients/forward/__init__.py +0 -0
  16. blut_core/ingredients/forward/_specs.py +35 -0
  17. blut_core/ingredients/logging/__init__.py +0 -0
  18. blut_core/ingredients/logging/_specs.py +49 -0
  19. blut_core/ingredients/loss/__init__.py +0 -0
  20. blut_core/ingredients/loss/_specs.py +142 -0
  21. blut_core/ingredients/model/__init__.py +0 -0
  22. blut_core/ingredients/model/_specs.py +138 -0
  23. blut_core/ingredients/optimizer/__init__.py +0 -0
  24. blut_core/ingredients/optimizer/_specs.py +164 -0
  25. blut_core/ingredients/optimizer/cautious_wd.py +23 -0
  26. blut_core/ingredients/optimizer/muon_optimizer.py +131 -0
  27. blut_core/ingredients/optimizer/soap_optimizer.py +301 -0
  28. blut_core/ingredients/sampler/__init__.py +0 -0
  29. blut_core/ingredients/sampler/_specs.py +92 -0
  30. blut_core/ingredients/scheduler/__init__.py +0 -0
  31. blut_core/ingredients/scheduler/_specs.py +195 -0
  32. blut_core/ingredients/scheduler/wsd.py +147 -0
  33. blut_core/ingredients/step/__init__.py +0 -0
  34. blut_core/ingredients/step/_specs.py +123 -0
  35. blut_core/load_dataset.py +68 -0
  36. blut_core/metric_log.py +136 -0
  37. blut_core/read_metric.py +296 -0
  38. blut_core/registry.py +109 -0
  39. blut_core/run_ledger.py +346 -0
  40. blut_core/run_manifest.py +204 -0
  41. blut_core/runctx.py +68 -0
  42. blut_core/spec.py +102 -0
  43. blut_core/status.py +90 -0
  44. blut_core/sysgauge.py +82 -0
  45. blut_core/tests/test_async_compute.py +161 -0
  46. blut_core/tests/test_async_io_contract.py +466 -0
  47. blut_core/tests/test_import_contract.py +32 -0
  48. blut_core/tests/test_load_dataset.py +39 -0
  49. blut_core/tests/test_parallel_strategy.py +64 -0
  50. blut_core/tests/test_primitives.py +111 -0
  51. blut_core/tests/test_read_metric.py +129 -0
  52. blut_core/tests/test_run_ledger.py +73 -0
  53. blut_core/trainer.py +631 -0
  54. blut_core-0.1.0.dist-info/METADATA +38 -0
  55. blut_core-0.1.0.dist-info/RECORD +57 -0
  56. blut_core-0.1.0.dist-info/WHEEL +4 -0
  57. blut_core-0.1.0.dist-info/licenses/LICENSE +202 -0
blut_core/__init__.py ADDED
@@ -0,0 +1,54 @@
1
+ """Shared Python runtime and ingredient registry for BLUT cookbooks.
2
+
3
+ The package contains domain-agnostic building blocks reusable by any cookbook:
4
+
5
+ - ``async_io`` — bounded, backpressured delivery for persistence work
6
+ - ``runctx`` — run identity + filesystem anchors from the BLUT env (P10)
7
+ - ``MetricLog`` — atomic per-epoch CSV/Parquet metric writer
8
+ - ``read_metric`` — verbatim metric/log reader (no LLM in the path; ADR 0038)
9
+ - ``status`` — emit/read the StatusUpdate wire protocol (P4)
10
+ - ``RunManifest`` — run provenance, written even on crash (P8)
11
+ - ``checkpoint`` — corruption-safe save/load + resume payload contract (P7)
12
+ - ``sysgauge`` — best-effort GPU/host snapshot for the metric stream (P10)
13
+ - ingredient specs and builders used by generic training recipes
14
+
15
+ ``torch`` is lazy-imported only inside ``checkpoint``/``sysgauge`` functions, so
16
+ ``import blut_core`` stays cheap and dependency-light. See ``blut/docs/metrics.md``
17
+ and decisions/0037, 0044.
18
+ """
19
+
20
+ from blut_core.registry import (
21
+ build_ingredient,
22
+ get_spec,
23
+ list_ingredients,
24
+ register_ingredient,
25
+ )
26
+ from blut_core.spec import KINDS, IngredientSpec
27
+
28
+ # Register generic ingredient specs on package import.
29
+ from blut_core.ingredients import _specs as _ingredient_specs # noqa: F401
30
+
31
+ from . import async_io, checkpoint, runctx, status, sysgauge
32
+ from .metric_log import MetricLog
33
+ from .run_manifest import RunManifest
34
+
35
+ # read_metric is a CLI module (`python -m blut_core.read_metric`); it is NOT
36
+ # eager-imported here (that triggers a runpy double-import warning under -m).
37
+ # `from blut_core import read_metric` still works (submodule import).
38
+
39
+ __all__ = [
40
+ "IngredientSpec",
41
+ "KINDS",
42
+ "MetricLog",
43
+ "RunManifest",
44
+ "async_io",
45
+ "build_ingredient",
46
+ "checkpoint",
47
+ "get_spec",
48
+ "list_ingredients",
49
+ "read_metric",
50
+ "register_ingredient",
51
+ "runctx",
52
+ "status",
53
+ "sysgauge",
54
+ ]
blut_core/async_io.py ADDED
@@ -0,0 +1,246 @@
1
+ """Bounded, backpressured asynchronous I/O for cookbook runtimes.
2
+
3
+ Bounded workers are daemons so an exceptional interpreter shutdown cannot hang
4
+ on an idle worker. Successful callers must still call :meth:`AsyncSink.close`;
5
+ that is the only boundary that drains accepted work and surfaces worker errors.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ from dataclasses import dataclass
11
+ from queue import Queue
12
+ from threading import BoundedSemaphore, Condition, Thread, local
13
+ from typing import Callable, Generic, TypeVar, cast
14
+
15
+
16
+ __all__ = [
17
+ "AsyncSink",
18
+ "AsyncSinkWorkerError",
19
+ "Bounded",
20
+ "Inline",
21
+ "ItemTooLargeError",
22
+ "SinkClosedError",
23
+ ]
24
+
25
+
26
+ ItemT = TypeVar("ItemT")
27
+
28
+
29
+ def _identity(item: ItemT) -> ItemT:
30
+ return item
31
+
32
+
33
+ class ItemTooLargeError(ValueError):
34
+ """A submitted item exceeds the sink's declared retention limit."""
35
+
36
+ def __init__(self, size_bytes: int, max_item_bytes: int) -> None:
37
+ self.size_bytes = size_bytes
38
+ self.max_item_bytes = max_item_bytes
39
+ super().__init__(f"item is {size_bytes} bytes; limit is {max_item_bytes}")
40
+
41
+
42
+ class SinkClosedError(RuntimeError):
43
+ """Work was submitted after sink shutdown began."""
44
+
45
+
46
+ class AsyncSinkWorkerError(RuntimeError):
47
+ """A sink worker failed while delivering an accepted item."""
48
+
49
+ def __init__(self, worker_error_type: str, message: str) -> None:
50
+ self.worker_error_type = worker_error_type
51
+ super().__init__(message)
52
+
53
+
54
+ @dataclass(frozen=True)
55
+ class _WorkerFailure:
56
+ worker_error_type: str
57
+ message: str
58
+
59
+ @classmethod
60
+ def capture(cls, error: BaseException) -> _WorkerFailure:
61
+ error_type = type(error)
62
+ worker_error_type = f"{error_type.__module__}.{error_type.__qualname__}"
63
+ try:
64
+ message = str(error)
65
+ except BaseException:
66
+ message = "worker failure with an unprintable message"
67
+ return cls(worker_error_type=worker_error_type, message=message)
68
+
69
+ def to_exception(self) -> AsyncSinkWorkerError:
70
+ return AsyncSinkWorkerError(self.worker_error_type, self.message)
71
+
72
+
73
+ @dataclass(frozen=True)
74
+ class Inline:
75
+ """Run the installed handler synchronously on the submitting thread."""
76
+
77
+
78
+ @dataclass(frozen=True)
79
+ class Bounded:
80
+ """Run work on one worker with bounded retained residency."""
81
+
82
+ capacity: int
83
+ max_item_bytes: int
84
+
85
+ def __post_init__(self) -> None:
86
+ if isinstance(self.capacity, bool) or not isinstance(self.capacity, int):
87
+ raise TypeError("capacity must be an integer")
88
+ if isinstance(self.max_item_bytes, bool) or not isinstance(
89
+ self.max_item_bytes, int
90
+ ):
91
+ raise TypeError("max_item_bytes must be an integer")
92
+ if self.capacity <= 0:
93
+ raise ValueError("capacity must be greater than zero")
94
+ if self.max_item_bytes <= 0:
95
+ raise ValueError("max_item_bytes must be greater than zero")
96
+
97
+
98
+ _STOP = object()
99
+
100
+
101
+ class AsyncSink(Generic[ItemT]):
102
+ """Deliver items to one handler under an explicit, close-required mode."""
103
+
104
+ def __init__(
105
+ self,
106
+ handler: Callable[[ItemT], None],
107
+ mode: Inline | Bounded,
108
+ *,
109
+ item_size: Callable[[ItemT], int] | None = None,
110
+ retain: Callable[[ItemT], ItemT] = _identity,
111
+ retained_size: Callable[[ItemT], int] | None = None,
112
+ thread_name: str = "blut-async-io",
113
+ ) -> None:
114
+ self._handler = handler
115
+ self._mode = mode
116
+ self._item_size = item_size
117
+ self._retain = retain
118
+ self._retained_size = retained_size
119
+ self._closed = False
120
+ self._failure: _WorkerFailure | None = None
121
+ self._state = Condition()
122
+ self._active_submits = 0
123
+ self._stop_enqueued = False
124
+ self._callback_state = local()
125
+ self._queue: Queue[ItemT | object] | None = None
126
+ self._slots: BoundedSemaphore | None = None
127
+ self._worker: Thread | None = None
128
+ if isinstance(mode, Bounded):
129
+ self._queue = Queue(maxsize=mode.capacity)
130
+ self._slots = BoundedSemaphore(mode.capacity)
131
+ self._worker = Thread(target=self._run, name=thread_name, daemon=True)
132
+ self._worker.start()
133
+
134
+ def submit(self, item: ItemT, *, size_bytes: int | None = None) -> None:
135
+ """Deliver one item according to the configured mode."""
136
+
137
+ if getattr(self._callback_state, "active", False):
138
+ raise RuntimeError("submit cannot be called from a sink callback")
139
+
140
+ with self._state:
141
+ self._raise_if_failed_locked()
142
+ if self._closed:
143
+ raise SinkClosedError("sink is closed")
144
+ self._active_submits += 1
145
+
146
+ acquired_slot = False
147
+ try:
148
+ if isinstance(self._mode, Inline):
149
+ self._invoke_handler(self._invoke_callback(self._retain, item))
150
+ return
151
+
152
+ if size_bytes is None and self._item_size is not None:
153
+ size_bytes = self._invoke_callback(self._item_size, item)
154
+ if size_bytes is None and self._retained_size is None:
155
+ raise TypeError(
156
+ "bounded submit requires size_bytes, item_size, or retained_size"
157
+ )
158
+ if size_bytes is not None:
159
+ if isinstance(size_bytes, bool) or not isinstance(size_bytes, int):
160
+ raise TypeError("item size must be an integer byte count")
161
+ if size_bytes < 0:
162
+ raise ValueError("item size must not be negative")
163
+ if size_bytes > self._mode.max_item_bytes:
164
+ raise ItemTooLargeError(size_bytes, self._mode.max_item_bytes)
165
+
166
+ slots = cast(BoundedSemaphore, self._slots)
167
+ queue = cast(Queue[ItemT | object], self._queue)
168
+ slots.acquire()
169
+ acquired_slot = True
170
+ with self._state:
171
+ self._raise_if_failed_locked()
172
+ if self._closed:
173
+ raise SinkClosedError("sink is closed")
174
+ retained = self._invoke_callback(self._retain, item)
175
+ if self._retained_size is not None:
176
+ retained_bytes = self._invoke_callback(self._retained_size, retained)
177
+ if isinstance(retained_bytes, bool) or not isinstance(retained_bytes, int):
178
+ raise TypeError("retained item size must be an integer byte count")
179
+ if retained_bytes < 0:
180
+ raise ValueError("retained item size must not be negative")
181
+ if retained_bytes > self._mode.max_item_bytes:
182
+ raise ItemTooLargeError(retained_bytes, self._mode.max_item_bytes)
183
+ queue.put(retained)
184
+ acquired_slot = False
185
+ finally:
186
+ if acquired_slot:
187
+ cast(BoundedSemaphore, self._slots).release()
188
+ with self._state:
189
+ self._active_submits -= 1
190
+ self._state.notify_all()
191
+
192
+ def close(self) -> None:
193
+ """Drain accepted work and stop the sink."""
194
+
195
+ if getattr(self._callback_state, "active", False):
196
+ raise RuntimeError("close cannot be called from a sink callback")
197
+
198
+ with self._state:
199
+ self._closed = True
200
+ while self._active_submits:
201
+ self._state.wait()
202
+ if self._worker is not None and not self._stop_enqueued:
203
+ cast(Queue[ItemT | object], self._queue).put(_STOP)
204
+ self._stop_enqueued = True
205
+
206
+ if self._worker is not None:
207
+ self._worker.join()
208
+ self._raise_if_failed()
209
+
210
+ def _run(self) -> None:
211
+ queue = cast(Queue[ItemT | object], self._queue)
212
+ slots = cast(BoundedSemaphore, self._slots)
213
+ while True:
214
+ item = queue.get()
215
+ if item is _STOP:
216
+ return
217
+ try:
218
+ self._invoke_handler(cast(ItemT, item))
219
+ except BaseException as error:
220
+ with self._state:
221
+ if self._failure is None:
222
+ self._failure = _WorkerFailure.capture(error)
223
+ finally:
224
+ del item
225
+ slots.release()
226
+
227
+ def _invoke_handler(self, item: ItemT) -> None:
228
+ self._invoke_callback(self._handler, item)
229
+
230
+ def _invoke_callback(self, callback, item):
231
+ previous = getattr(self._callback_state, "active", False)
232
+ self._callback_state.active = True
233
+ try:
234
+ return callback(item)
235
+ finally:
236
+ self._callback_state.active = previous
237
+
238
+ def _raise_if_failed_locked(self) -> None:
239
+ if self._failure is not None:
240
+ raise self._failure.to_exception()
241
+
242
+ def _raise_if_failed(self) -> None:
243
+ with self._state:
244
+ failure = self._failure
245
+ if failure is not None:
246
+ raise failure.to_exception()
@@ -0,0 +1,172 @@
1
+ """checkpoint — corruption-safe torch checkpoint I/O + a resume payload contract
2
+ (ADR 0044 P7, "durability"; the single most important reliability gap).
3
+
4
+ - ``save(payload, path)`` — free-space preflight → write to a unique tmp in the
5
+ same dir → ``fsync`` → atomic ``os.replace`` → write a sidecar ``<path>.sha256``.
6
+ A mid-write kill leaves the PRIOR checkpoint intact, never a truncated file
7
+ (the failure mode that silently corrupts on a disk-full ``torch.save``).
8
+ Optional ``keep_last_n`` rotation of ``<stem>_ep*.pt`` siblings.
9
+ - ``load(path, validate=True)`` — verify the sidecar SHA (or, absent a sidecar,
10
+ a dry ``torch.load`` of the bytes) BEFORE returning → a corrupt/truncated
11
+ checkpoint raises ``CheckpointError``, never silent garbage.
12
+ - ``validate_payload(payload, required)`` — assert the resume contract keys are
13
+ present so a resume fails fast at SAVE, not 3 epochs into a relaunch (prior
14
+ bites: optimizer-state, QAT-alpha). Default contract = the training-resume set.
15
+
16
+ torch is lazy-imported (only on save/load), so importing ``blut_core`` stays cheap.
17
+ """
18
+ from __future__ import annotations
19
+
20
+ import hashlib
21
+ import os
22
+ import shutil
23
+ import tempfile
24
+ from pathlib import Path
25
+ from typing import Any, Dict, Iterable, Optional
26
+
27
+ __all__ = ["save", "load", "validate_payload", "RESUME_CONTRACT", "CheckpointError"]
28
+
29
+ # ADR 0044 P7 resume payload contract. A trainer that wants resumability MUST
30
+ # put these in the checkpoint; validate_payload(strict=True) enforces it.
31
+ RESUME_CONTRACT = (
32
+ "model", "optimizer", "lr_scheduler", "rng_state", "step",
33
+ )
34
+
35
+
36
+ class CheckpointError(RuntimeError):
37
+ """A checkpoint is missing, corrupt, or violates the resume contract."""
38
+
39
+
40
+ def _sha256(path: Path) -> str:
41
+ h = hashlib.sha256()
42
+ with path.open("rb") as f:
43
+ for chunk in iter(lambda: f.read(1 << 20), b""):
44
+ h.update(chunk)
45
+ return h.hexdigest()
46
+
47
+
48
+ def validate_payload(payload: Dict[str, Any], required: Iterable[str] = RESUME_CONTRACT,
49
+ *, strict: bool = True) -> list:
50
+ """Return the list of missing contract keys. ``strict`` → raise on any missing.
51
+ Use before save so a resume can't silently lack optimizer/RNG/step."""
52
+ if not isinstance(payload, dict):
53
+ raise CheckpointError(f"payload must be a dict, got {type(payload).__name__}")
54
+ missing = [k for k in required if k not in payload]
55
+ if missing and strict:
56
+ raise CheckpointError(
57
+ f"checkpoint payload missing resume-contract keys {missing}; "
58
+ f"present: {sorted(payload)}")
59
+ return missing
60
+
61
+
62
+ def _free_bytes(directory: Path) -> int:
63
+ return shutil.disk_usage(directory).free
64
+
65
+
66
+ def save(payload: Dict[str, Any], path, *, contract: Optional[Iterable[str]] = None,
67
+ keep_last_n: Optional[int] = None, min_free_mb: int = 256) -> Path:
68
+ """Atomically save ``payload`` to ``path`` with a sidecar SHA.
69
+
70
+ ``contract`` (default None = skip) → validate_payload(strict) first.
71
+ ``keep_last_n`` → after writing, keep only the newest N ``<stem>_ep*.pt``.
72
+ Raises ``CheckpointError`` on insufficient free space or a contract miss."""
73
+ import torch # lazy
74
+ path = Path(path)
75
+ path.parent.mkdir(parents=True, exist_ok=True)
76
+ if contract is not None:
77
+ validate_payload(payload, contract, strict=True)
78
+
79
+ # Free-space preflight: a disk-full torch.save silently yields garbage.
80
+ if _free_bytes(path.parent) < min_free_mb * (1 << 20):
81
+ raise CheckpointError(
82
+ f"insufficient free space in {path.parent} (<{min_free_mb} MiB) — refusing to save")
83
+
84
+ fd, tmp_name = tempfile.mkstemp(dir=str(path.parent), prefix=path.stem + ".",
85
+ suffix=path.suffix + ".tmp")
86
+ tmp = Path(tmp_name)
87
+ try:
88
+ os.close(fd) # torch.save writes by path
89
+ torch.save(payload, tmp)
90
+ with tmp.open("rb") as f:
91
+ os.fsync(f.fileno())
92
+ sha = _sha256(tmp)
93
+ os.replace(tmp, path)
94
+ # Sidecar: write to a tmp, fsync, then atomic rename — write_text() left
95
+ # the SHA in the page cache (no fsync), and a non-atomic overwrite could
96
+ # leave a .sha256 pointing at the PREVIOUS checkpoint after a crash.
97
+ sha_path = path.with_suffix(path.suffix + ".sha256")
98
+ sha_tmp = sha_path.with_suffix(sha_path.suffix + ".tmp")
99
+ payload_bytes = (sha + "\n").encode("utf-8")
100
+ try:
101
+ sfd = os.open(str(sha_tmp), os.O_WRONLY | os.O_CREAT | os.O_TRUNC, 0o644)
102
+ try:
103
+ n = os.write(sfd, payload_bytes)
104
+ if n != len(payload_bytes): # never short on a local FS at this size
105
+ raise OSError(f"short sidecar write: {n}/{len(payload_bytes)}")
106
+ os.fsync(sfd)
107
+ finally:
108
+ os.close(sfd)
109
+ os.replace(sha_tmp, sha_path)
110
+ except Exception:
111
+ try:
112
+ sha_tmp.unlink()
113
+ except OSError:
114
+ pass
115
+ raise
116
+ except Exception:
117
+ try:
118
+ tmp.unlink()
119
+ except Exception:
120
+ pass
121
+ raise
122
+
123
+ if keep_last_n is not None and keep_last_n > 0:
124
+ _rotate(path, keep_last_n)
125
+ return path
126
+
127
+
128
+ def _rotate(path: Path, keep_last_n: int) -> None:
129
+ # Rotates the EPOCH-SNAPSHOT siblings `{stem}_ep<N>{suffix}` (the trainer
130
+ # convention, e.g. snn_4state_best_ep10.pt), keeping the newest N. It does
131
+ # NOT touch the base `{stem}{suffix}` (the canonical/best checkpoint) — that
132
+ # is intentional: you keep best + the last N epoch snapshots.
133
+ sibs = sorted(path.parent.glob(f"{path.stem}_ep*{path.suffix}"),
134
+ key=lambda p: p.stat().st_mtime, reverse=True)
135
+ for old in sibs[keep_last_n:]:
136
+ for victim in (old, old.with_suffix(old.suffix + ".sha256")):
137
+ try:
138
+ victim.unlink()
139
+ except OSError:
140
+ pass
141
+
142
+
143
+ def load(path, *, validate: bool = True, map_location: Any = "cpu") -> Dict[str, Any]:
144
+ """Load a checkpoint, verifying integrity first. Raises ``CheckpointError``
145
+ on a missing/corrupt file rather than returning silent garbage."""
146
+ import torch # lazy
147
+ path = Path(path)
148
+ if not path.exists():
149
+ raise CheckpointError(f"no checkpoint: {path}")
150
+ if validate:
151
+ sidecar = path.with_suffix(path.suffix + ".sha256")
152
+ if sidecar.exists():
153
+ want = sidecar.read_text().strip()
154
+ got = _sha256(path)
155
+ if got != want:
156
+ raise CheckpointError(
157
+ f"checkpoint SHA mismatch for {path}: sidecar {want[:12]}… got {got[:12]}…")
158
+ else:
159
+ # No sidecar → integrity is UNVERIFIED (e.g. a crash between the
160
+ # checkpoint rename and the sidecar write; the sidecar is not part
161
+ # of the atomic unit). Surface it; don't silently accept-as-clean.
162
+ import sys as _sys
163
+ _sys.stderr.write(
164
+ f"[checkpoint] WARNING: no .sha256 sidecar for {path}; integrity unverified\n")
165
+ try:
166
+ # NOT weights_only=True: the resume payload carries optimizer/LR/RNG/
167
+ # scheduler objects (not plain tensors), which weights_only rejects.
168
+ # Trust model: checkpoints are own-produced under the job dir, not loaded
169
+ # from untrusted sources — the sidecar SHA guards corruption, not supply chain.
170
+ return torch.load(path, map_location=map_location)
171
+ except Exception as e:
172
+ raise CheckpointError(f"failed to load {path}: {type(e).__name__}: {e}") from e
blut_core/evaluator.py ADDED
@@ -0,0 +1,98 @@
1
+ """Generic ingredient-based evaluator for BLUT core cookbook.
2
+
3
+ Reads an eval config JSON (written by the EvaluateModel Rust stage),
4
+ builds the eval ingredient, and runs evaluation.
5
+
6
+ Usage:
7
+ python -m blut_core.evaluator --config eval_config.json --output eval_report.json
8
+ """
9
+ from __future__ import annotations
10
+
11
+ import argparse
12
+ import json
13
+ import sys
14
+ from pathlib import Path
15
+
16
+ sys.path.insert(0, str(Path(__file__).parent.parent))
17
+
18
+ from blut_core import build_ingredient
19
+
20
+
21
+ def load_config(config_path: str) -> dict:
22
+ with open(config_path) as f:
23
+ return json.load(f)
24
+
25
+
26
+ def main():
27
+ parser = argparse.ArgumentParser(description="BLUT core generic evaluator")
28
+ parser.add_argument("--config", required=True, help="Path to eval config JSON")
29
+ parser.add_argument("--output", required=True, help="Path to output eval report JSON")
30
+ args = parser.parse_args()
31
+
32
+ config = load_config(args.config)
33
+
34
+ checkpoint_path = config["checkpoint_path"]
35
+ dataset_path = config["dataset_path"]
36
+ eval_cfg = config["eval"]
37
+ batch_size = config.get("batch_size", 64)
38
+ device = config.get("device", "cuda")
39
+
40
+ # Build eval ingredient
41
+ print(f"[evaluator] building eval: {eval_cfg['kind']}:{eval_cfg['name']}")
42
+ evaluate_fn = build_ingredient(eval_cfg["kind"], eval_cfg["name"], eval_cfg.get("config"))
43
+
44
+ # Load checkpoint
45
+ import torch
46
+ print(f"[evaluator] loading checkpoint from {checkpoint_path}")
47
+ checkpoint_path = Path(checkpoint_path)
48
+
49
+ # Try to load as HuggingFace model
50
+ model = None
51
+ hf_dir = checkpoint_path / "hf"
52
+ if hf_dir.exists():
53
+ try:
54
+ from transformers import AutoModelForCausalLM
55
+ model = AutoModelForCausalLM.from_pretrained(str(hf_dir))
56
+ print(f"[evaluator] loaded HF model from {hf_dir}")
57
+ except Exception as e:
58
+ print(f"[evaluator] warning: could not load HF model: {e}")
59
+
60
+ if model is None:
61
+ # Try loading as torch checkpoint
62
+ pt_path = checkpoint_path / "model.pt"
63
+ if pt_path.exists():
64
+ print(f"[evaluator] found torch checkpoint at {pt_path}")
65
+ # Need a model architecture to load into — this is a limitation
66
+ # of the generic evaluator. Real usage should provide the model.
67
+ print("[evaluator] error: no model architecture available for torch checkpoint")
68
+ sys.exit(1)
69
+ else:
70
+ print(f"[evaluator] error: no checkpoint found at {checkpoint_path}")
71
+ sys.exit(1)
72
+
73
+ # Load dataset
74
+ print(f"[evaluator] loading dataset from {dataset_path}")
75
+ try:
76
+ from datasets import load_dataset
77
+ ds = load_dataset("json", data_files=str(dataset_path), split="train")
78
+ print(f"[evaluator] loaded {len(ds)} examples")
79
+ except Exception as e:
80
+ print(f"[evaluator] error: could not load dataset: {e}")
81
+ sys.exit(1)
82
+
83
+ # Run evaluation
84
+ print(f"[evaluator] running evaluation...")
85
+ if device != "cpu" and torch.cuda.is_available():
86
+ model = model.to(device)
87
+
88
+ metrics = evaluate_fn(model, ds, device=device)
89
+ print(f"[evaluator] metrics: {metrics}")
90
+
91
+ # Write output
92
+ with open(args.output, "w") as f:
93
+ json.dump(metrics, f, indent=2)
94
+ print(f"[evaluator] report written to {args.output}")
95
+
96
+
97
+ if __name__ == "__main__":
98
+ main()
@@ -0,0 +1,20 @@
1
+ """Generic ingredient specs — the core training primitives.
2
+
3
+ Importing this module registers all generic specs into the global registry.
4
+ Domain cookbooks import ``blut_core`` first (registering these), then register
5
+ their own domain-specific specs on top.
6
+ """
7
+
8
+ # Import all _specs modules to trigger their @register_ingredient decorators.
9
+ from blut_core.ingredients.data import _specs as _data_specs # noqa: F401
10
+ from blut_core.ingredients.model import _specs as _model_specs # noqa: F401
11
+ from blut_core.ingredients.optimizer import _specs as _optimizer_specs # noqa: F401
12
+ from blut_core.ingredients.scheduler import _specs as _scheduler_specs # noqa: F401
13
+ from blut_core.ingredients.loss import _specs as _loss_specs # noqa: F401
14
+ from blut_core.ingredients.step import _specs as _step_specs # noqa: F401
15
+ from blut_core.ingredients.ema import _specs as _ema_specs # noqa: F401
16
+ from blut_core.ingredients.checkpoint import _specs as _checkpoint_specs # noqa: F401
17
+ from blut_core.ingredients.eval import _specs as _eval_specs # noqa: F401
18
+ from blut_core.ingredients.sampler import _specs as _sampler_specs # noqa: F401
19
+ from blut_core.ingredients.logging import _specs as _logging_specs # noqa: F401
20
+ from blut_core.ingredients.forward import _specs as _forward_specs # noqa: F401
@@ -0,0 +1,2 @@
1
+ # Compatibility shim — all registrations happen in __init__.py already.
2
+ # This module exists because blut_core.__init__ imports it.
File without changes