diff --git a/scripts/shadow/hyperlexical/loop.py b/scripts/shadow/hyperlexical/loop.py index fba5fade..ba9d6013 100644 --- a/scripts/shadow/hyperlexical/loop.py +++ b/scripts/shadow/hyperlexical/loop.py @@ -4,6 +4,8 @@ import json import os +import time +from dataclasses import dataclass from pathlib import Path from .align import atom_token_index, offsets_from_tokenizer, pool_indices @@ -13,6 +15,7 @@ from .classify_metrics import ( NONE_LABEL, SELECT_METRIC_CLASSIFY, + SELECT_METRIC_ENV, macro_f1_nonnone, none_false_positive_rate, resolve_select_metric, @@ -86,6 +89,16 @@ UNBIND_EVERY_N_ENV = "HYPERLEX_UNBIND_EVERY_N" SAVE_BEST_UNBIND_ENV = "HYPERLEX_SAVE_BEST_UNBIND" INIT_FROM_ENV = "HYPERLEX_INIT_FROM" +EARLY_STOP_ENV = "HYPERLEX_EARLY_STOP" +EARLY_STOP_PATIENCE_ENV = "HYPERLEX_EARLY_STOP_PATIENCE" +EARLY_STOP_MIN_EPOCHS_ENV = "HYPERLEX_EARLY_STOP_MIN_EPOCHS" +STOP_REASON_MAX_EPOCHS = "max_epochs" +STOP_REASON_EARLY_STOPPING = "early_stopping" +# Observational seconds on epoch-progress.jsonl and the completion receipt. +# Python round() to 6 decimal places (microseconds). Not a selection input. +WALLCLOCK_SECONDS_DECIMALS = 6 +_EARLY_STOP_OFF = frozenset({"0", "false", "no", "off"}) +_EARLY_STOP_ON = frozenset({"1", "true", "yes", "on"}) INIT_EXPAND_VOCAB_ENV = "HYPERLEX_INIT_EXPAND_VOCAB" UNBIND_LOSS_WEIGHT_DEFAULT = 1.0 UNBIND_EVERY_N_DEFAULT = 1 @@ -104,6 +117,156 @@ def resolve_save_best_unbind(raw: str | None = None) -> bool: return str(raw).strip().lower() in {"1", "true", "yes", "on"} +@dataclass(frozen=True) +class EarlyStopConfig: + """Optional classify-metric early stop. Default-off. Does not change the epoch cap.""" + + enabled: bool + max_epochs: int + patience: int | None = None + minimum_epochs: int | None = None + select_metric: str = "" + + +def _parse_early_stop_int(env_name: str, raw: str | None) -> int: + if raw is None or not str(raw).strip(): + raise ValueError(f"{env_name} is required when {EARLY_STOP_ENV} is enabled") + text = str(raw).strip() + sign = "" + digits = text + if text[0] in "+-": + sign, digits = text[0], text[1:] + if sign == "+" and not digits: + raise ValueError(f"{env_name} must be an integer, got {raw!r}") + if not digits.isdigit(): + raise ValueError(f"{env_name} must be an integer, got {raw!r}") + return int(sign + digits) + + +def resolve_early_stop_config(*, max_epochs: int, select_metric: str) -> EarlyStopConfig: + """Read early-stop env before the optimizer is constructed. + + Unset or disabled (0/false/no/off) leaves the configured epoch schedule + unchanged. ``HYPERLEX_TRAIN_EPOCHS`` is only the max-epoch cap and never + turns this on. Enabled runs require ``HLX_SELECT_METRIC=classify_macro_f1_nonnone``, + patience >= 0, and minimum scored epochs in ``1..max_epochs``. + """ + raw = os.environ.get(EARLY_STOP_ENV) + token = "" if raw is None else str(raw).strip().lower() + if raw is None or token == "" or token in _EARLY_STOP_OFF: + return EarlyStopConfig( + enabled=False, + max_epochs=max_epochs, + select_metric=select_metric, + ) + if token not in _EARLY_STOP_ON: + raise ValueError( + f"{EARLY_STOP_ENV} must be unset, off, or on (1/true/yes/on), got {raw!r}" + ) + if select_metric != SELECT_METRIC_CLASSIFY: + raise ValueError( + f"{EARLY_STOP_ENV} requires {SELECT_METRIC_ENV}={SELECT_METRIC_CLASSIFY}; " + f"got {select_metric!r}" + ) + patience = _parse_early_stop_int( + EARLY_STOP_PATIENCE_ENV, os.environ.get(EARLY_STOP_PATIENCE_ENV) + ) + minimum_epochs = _parse_early_stop_int( + EARLY_STOP_MIN_EPOCHS_ENV, os.environ.get(EARLY_STOP_MIN_EPOCHS_ENV) + ) + if patience < 0: + raise ValueError(f"{EARLY_STOP_PATIENCE_ENV} must be >= 0, got {patience}") + if minimum_epochs < 1: + raise ValueError( + f"{EARLY_STOP_MIN_EPOCHS_ENV} must be >= 1 scored epoch, got {minimum_epochs}" + ) + if minimum_epochs > max_epochs: + raise ValueError( + f"{EARLY_STOP_MIN_EPOCHS_ENV}={minimum_epochs} exceeds " + f"HYPERLEX_TRAIN_EPOCHS={max_epochs}" + ) + return EarlyStopConfig( + enabled=True, + max_epochs=max_epochs, + patience=patience, + minimum_epochs=minimum_epochs, + select_metric=select_metric, + ) + + +def note_strict_improvement( + best_value: float, + best_epoch: int | None, + score: float, + epoch_index: int, +) -> tuple[float, int | None, bool]: + """Strict increase replaces the checkpoint. A tie keeps the earlier epoch.""" + if score > best_value: + return score, epoch_index, True + return best_value, best_epoch, False + + +def early_stop_break( + *, + enabled: bool, + epoch_index: int, + best_epoch: int | None, + epochs_scored: int, + minimum_epochs: int | None, + patience: int | None, + max_epochs: int, +) -> bool: + """Return true only when the loop should break before the epoch cap. + + Call this after the epoch is scored and any strict improvement is recorded. + Patience is ``epoch_index - best_epoch`` completed epochs after the best + (0-based indices). The best epoch itself does not consume patience. + Ties do not move ``best_epoch``, so they do consume patience. + Do not stop before ``minimum_epochs`` epochs have been scored. + When the patience condition lands on the final configured epoch, the max + epoch cap wins and this returns false so the loop records ``max_epochs``. + """ + if not enabled: + return False + if epoch_index + 1 >= max_epochs: + return False + if minimum_epochs is None or patience is None or best_epoch is None: + return False + if epochs_scored < minimum_epochs: + return False + return (epoch_index - best_epoch) >= patience + + +def round_seconds(value: float) -> float: + """Seconds rounded to 6 decimal places. Observational; not a selection input.""" + return round(float(value), WALLCLOCK_SECONDS_DECIMALS) + + +def monotonic_seconds() -> float: + """Monotonic clock in seconds. Tests replace this; production uses time.monotonic.""" + return time.monotonic() + + +def epoch_timing_fields( + *, + epoch_started: float, + epoch_ended: float, + training_started: float, +) -> dict[str, float]: + """Wall-clock fields for one epoch-progress.jsonl row. + + Units are seconds. Both values use ``round_seconds`` (6 decimal places). + ``epoch_wallclock_seconds`` is this epoch's scored body + (epoch_ended - epoch_started). ``training_elapsed_seconds`` is the time + from the start of the epoch loop to this epoch's end. Neither value is + read by checkpoint selection or early stopping. + """ + return { + "epoch_wallclock_seconds": round_seconds(epoch_ended - epoch_started), + "training_elapsed_seconds": round_seconds(epoch_ended - training_started), + } + + def resolve_init_from(raw: str | None = None) -> Path | None: """Optional warm-start dir with model.safetensors (or heads.pt). @@ -588,8 +751,9 @@ def run_loop( ), } trainable = [p for p in encoder.parameters() if p.requires_grad] + list(classify.parameters()) + list(role_head.parameters()) + list(filler_head.parameters()) - opt = AdamW(trainable, lr=float(os.environ.get("HYPERLEX_TRAIN_LR", "2e-5"))) epochs = int(os.environ.get("HYPERLEX_TRAIN_EPOCHS", "2")) + early_stop = resolve_early_stop_config(max_epochs=epochs, select_metric=select_metric) + opt = AdamW(trainable, lr=float(os.environ.get("HYPERLEX_TRAIN_LR", "2e-5"))) batch = int(os.environ.get("HYPERLEX_TRAIN_BATCH", "8")) unbind_loss_weight = resolve_unbind_loss_weight() unbind_every_n = resolve_unbind_every_n() @@ -683,6 +847,7 @@ def step_unbind(row) -> None: save_best_unbind = resolve_save_best_unbind() best_exact = float("-inf") best_macro = float("-inf") + best_epoch_index: int | None = None best_metrics: dict | None = None best_state: dict | None = None best_residual_records: list[dict] = [] @@ -754,7 +919,10 @@ def score(): layout["aligner"] = "char_span + offset_mapping" write_skeleton(out_dir, maps=maps) + training_started = monotonic_seconds() + stop_reason = STOP_REASON_MAX_EPOCHS for ep in range(epochs): + epoch_started = monotonic_seconds() phase_rows, phase_meta = select_unbind_for_epoch(unbind_tr, ep, curriculum) unbind_cycle = 0 classify_batch_i = 0 @@ -791,10 +959,15 @@ def score(): "non-none gold; macro-F1 is NOT_COMPUTABLE" ) score_now = float(metrics["classify_macro_f1_nonnone"]) - if score_now > best_macro: + best_macro, best_epoch_index, improved = note_strict_improvement( + best_macro, + best_epoch_index, + score_now, + ep, + ) + if improved: if device.type == "cuda": torch.cuda.synchronize() - best_macro = score_now best_metrics = dict(metrics) best_state = _build_weight_state( encoder, classify, role_head, filler_head, maps, layout @@ -889,6 +1062,13 @@ def score(): ) if seed_receipt is not None: progress["seed"] = seed_receipt["seed"] + progress.update( + epoch_timing_fields( + epoch_started=epoch_started, + epoch_ended=monotonic_seconds(), + training_started=training_started, + ) + ) with progress_path.open("a", encoding="utf-8") as pf: pf.write(json.dumps(progress, sort_keys=True) + "\n") pf.flush() @@ -897,7 +1077,19 @@ def score(): f"[epoch {ep}] unbind_exact={exact:.6f} best={best_exact if best_exact != float('-inf') else None} saved_best={saved_best}", flush=True, ) - + if early_stop_break( + enabled=early_stop.enabled, + epoch_index=ep, + best_epoch=best_epoch_index, + epochs_scored=ep + 1, + minimum_epochs=early_stop.minimum_epochs, + patience=early_stop.patience, + max_epochs=epochs, + ): + stop_reason = STOP_REASON_EARLY_STOPPING + break + + training_elapsed_seconds = round_seconds(monotonic_seconds() - training_started) final_state = _build_weight_state( encoder, classify, role_head, filler_head, maps, layout ) @@ -941,6 +1133,8 @@ def score(): "device": str(device), "cuda": bool(torch.cuda.is_available()), "epochs": epochs, + "stop_reason": stop_reason, + "training_elapsed_seconds": training_elapsed_seconds, **provenance(root), "n_train_classify": len(classify_tr), "task_accounting": task_accounting, diff --git a/specs/007-hyperlexical-model/AARON-SPARK-TRAIN.md b/specs/007-hyperlexical-model/AARON-SPARK-TRAIN.md index f5a276ff..3b9c1fe5 100644 --- a/specs/007-hyperlexical-model/AARON-SPARK-TRAIN.md +++ b/specs/007-hyperlexical-model/AARON-SPARK-TRAIN.md @@ -70,6 +70,21 @@ export HYPERLEX_OFFLINE=1 export TOKENIZERS_PARALLELISM=false export HYPERLEX_TRAIN_OUT="$HOME/.hyperlex/models/hyperlex-encoder-modernbert-base-seed" export HYPERLEX_TRAIN_EPOCHS=2 +# HYPERLEX_TRAIN_EPOCHS is the max-epoch cap. It does not enable early stopping. +# Early stop is optional and default-off. Turn it on only with the classify +# selection metric. Improvement is a strict increase; ties keep the earlier +# checkpoint. After the epoch is scored, stop when +# epoch_index - best_epoch >= patience and at least MIN_EPOCHS epochs have +# been scored (0-based epoch_index). The best checkpoint is still restored. +# epoch-progress.jsonl gains observational seconds rounded to 6 decimal places +# (epoch_wallclock_seconds, training_elapsed_seconds). They do not affect +# selection or stopping. A completed train-receipt.json records stop_reason +# (max_epochs or early_stopping) and training_elapsed_seconds. A failed loop +# still raises and does not write that completion receipt. +# export HLX_SELECT_METRIC=classify_macro_f1_nonnone +# export HYPERLEX_EARLY_STOP=1 +# export HYPERLEX_EARLY_STOP_PATIENCE=4 +# export HYPERLEX_EARLY_STOP_MIN_EPOCHS=4 export HYPERLEX_TRAIN_BATCH=8 export HYPERLEX_TRAIN_LR=2e-5 # optional recipe bump (default 2, clamp 1..min(encoder layers, 8)): diff --git a/tests/shadow/test_hyperlexical_early_stop.py b/tests/shadow/test_hyperlexical_early_stop.py new file mode 100644 index 00000000..b6557366 --- /dev/null +++ b/tests/shadow/test_hyperlexical_early_stop.py @@ -0,0 +1,462 @@ +"""Default-off early stop for the classify training loop. No sleeps.""" + +from __future__ import annotations + +import ast +import inspect +from pathlib import Path +import sys + +import pytest + +ROOT = Path(__file__).resolve().parents[2] +sys.path.insert(0, str(ROOT / "scripts" / "shadow")) +SRC = ROOT / "src" +if str(SRC) not in sys.path: + sys.path.insert(0, str(SRC)) + +from hyperlexical.classify_metrics import SELECT_METRIC_CLASSIFY, SELECT_METRIC_UNBIND +from hyperlexical.loop import ( + EARLY_STOP_ENV, + EARLY_STOP_MIN_EPOCHS_ENV, + EARLY_STOP_PATIENCE_ENV, + STOP_REASON_EARLY_STOPPING, + STOP_REASON_MAX_EPOCHS, + early_stop_break, + epoch_timing_fields, + monotonic_seconds, + note_strict_improvement, + resolve_early_stop_config, + round_seconds, +) +import hyperlexical.loop as loop + +# Best classify score at epoch 3. Later values never strictly exceed it. +CANDIDATE = [0.50, 0.40, 0.45, 0.70, 0.55, 0.60, 0.58, 0.69, 0.10, 0.11, 0.12, 0.13] + + +def _enable(monkeypatch, *, patience="4", minimum_epochs="4"): + monkeypatch.setenv(EARLY_STOP_ENV, "1") + monkeypatch.setenv(EARLY_STOP_PATIENCE_ENV, patience) + monkeypatch.setenv(EARLY_STOP_MIN_EPOCHS_ENV, minimum_epochs) + + +def _disable(monkeypatch): + monkeypatch.delenv(EARLY_STOP_ENV, raising=False) + monkeypatch.delenv(EARLY_STOP_PATIENCE_ENV, raising=False) + monkeypatch.delenv(EARLY_STOP_MIN_EPOCHS_ENV, raising=False) + + +def run_classify_schedule( + scores: list[float], + *, + max_epochs: int, + enabled: bool, + patience: int | None = None, + minimum_epochs: int | None = None, +) -> dict[str, object]: + """Same order as the trainer: score, strict improvement, then stop check.""" + best_value = float("-inf") + best_epoch = None + restored_epoch = None + stop_reason = STOP_REASON_MAX_EPOCHS + scored = 0 + last_epoch = None + for ep in range(max_epochs): + if ep >= len(scores): + raise AssertionError(f"missing score for epoch {ep}") + best_value, best_epoch, improved = note_strict_improvement( + best_value, best_epoch, scores[ep], ep + ) + if improved: + restored_epoch = ep + scored = ep + 1 + last_epoch = ep + if early_stop_break( + enabled=enabled, + epoch_index=ep, + best_epoch=best_epoch, + epochs_scored=scored, + minimum_epochs=minimum_epochs, + patience=patience, + max_epochs=max_epochs, + ): + stop_reason = STOP_REASON_EARLY_STOPPING + break + return { + "stop_reason": stop_reason, + "epochs_scored": scored, + "last_epoch": last_epoch, + "best_epoch": best_epoch, + "restored_epoch": restored_epoch, + "best_value": best_value, + } + + +def test_default_off_runs_the_full_schedule(monkeypatch): + _disable(monkeypatch) + monkeypatch.setenv("HYPERLEX_TRAIN_EPOCHS", "12") + monkeypatch.setenv(EARLY_STOP_PATIENCE_ENV, "-1") + cfg = resolve_early_stop_config(max_epochs=12, select_metric=SELECT_METRIC_CLASSIFY) + assert cfg.enabled is False + result = run_classify_schedule( + CANDIDATE, + max_epochs=12, + enabled=cfg.enabled, + patience=4, + minimum_epochs=4, + ) + assert result["epochs_scored"] == 12 + assert result["last_epoch"] == 11 + assert result["stop_reason"] == STOP_REASON_MAX_EPOCHS + assert result["best_epoch"] == 3 + + +def test_train_epochs_alone_does_not_enable_early_stop(monkeypatch): + _disable(monkeypatch) + monkeypatch.setenv("HYPERLEX_TRAIN_EPOCHS", "12") + cfg = resolve_early_stop_config(max_epochs=12, select_metric=SELECT_METRIC_UNBIND) + assert cfg.enabled is False + assert cfg.patience is None + assert cfg.minimum_epochs is None + + +def test_minimum_epochs_block_an_earlier_patience_fire(): + scores = [0.90, 0.10, 0.10, 0.10, 0.10, 0.10] + result = run_classify_schedule( + scores, + max_epochs=6, + enabled=True, + patience=1, + minimum_epochs=4, + ) + assert result["best_epoch"] == 0 + assert result["last_epoch"] == 3 + assert result["epochs_scored"] == 4 + assert result["stop_reason"] == STOP_REASON_EARLY_STOPPING + + +def test_best_at_epoch_3_stops_after_epoch_7_with_patience_4(): + result = run_classify_schedule( + CANDIDATE, + max_epochs=12, + enabled=True, + patience=4, + minimum_epochs=4, + ) + assert result["best_epoch"] == 3 + assert result["last_epoch"] == 7 + assert result["epochs_scored"] == 8 + assert result["stop_reason"] == STOP_REASON_EARLY_STOPPING + assert result["restored_epoch"] == 3 + + +def test_strict_improvement_resets_patience(): + scores = list(CANDIDATE) + scores[5] = 0.80 + result = run_classify_schedule( + scores, + max_epochs=12, + enabled=True, + patience=4, + minimum_epochs=4, + ) + assert result["best_epoch"] == 5 + assert result["last_epoch"] == 9 + assert result["epochs_scored"] == 10 + assert result["stop_reason"] == STOP_REASON_EARLY_STOPPING + + +def test_ties_keep_the_earlier_checkpoint_and_consume_patience(): + scores = list(CANDIDATE) + scores[7] = scores[3] + assert scores[7] == scores[3] + result = run_classify_schedule( + scores, + max_epochs=12, + enabled=True, + patience=4, + minimum_epochs=4, + ) + assert result["best_epoch"] == 3 + assert result["restored_epoch"] == 3 + assert result["last_epoch"] == 7 + assert result["epochs_scored"] == 8 + assert result["best_value"] == scores[3] + tied = [0.50, 0.50] + short = run_classify_schedule( + tied, + max_epochs=4, + enabled=True, + patience=1, + minimum_epochs=1, + ) + assert short["best_epoch"] == 0 + assert short["restored_epoch"] == 0 + assert short["last_epoch"] == 1 + assert short["stop_reason"] == STOP_REASON_EARLY_STOPPING + + +def test_improvement_on_the_would_be_stopping_epoch_prevents_stopping(): + scores = list(CANDIDATE) + [0.10, 0.10, 0.10, 0.10] + scores[7] = 0.71 + stopped_without = run_classify_schedule( + CANDIDATE, + max_epochs=16, + enabled=True, + patience=4, + minimum_epochs=4, + ) + assert stopped_without["last_epoch"] == 7 + result = run_classify_schedule( + scores, + max_epochs=16, + enabled=True, + patience=4, + minimum_epochs=4, + ) + assert result["best_epoch"] == 7 + assert result["restored_epoch"] == 7 + assert result["last_epoch"] == 11 + assert result["epochs_scored"] == 12 + assert result["stop_reason"] == STOP_REASON_EARLY_STOPPING + + +def test_max_epoch_cap_wins_when_reached_first(): + early = run_classify_schedule( + CANDIDATE, + max_epochs=6, + enabled=True, + patience=4, + minimum_epochs=4, + ) + assert early["epochs_scored"] == 6 + assert early["last_epoch"] == 5 + assert early["best_epoch"] == 3 + assert early["stop_reason"] == STOP_REASON_MAX_EPOCHS + # Patience would also be true on the final configured epoch. The cap wins. + on_the_cap = run_classify_schedule( + CANDIDATE, + max_epochs=8, + enabled=True, + patience=4, + minimum_epochs=4, + ) + assert on_the_cap["last_epoch"] == 7 + assert on_the_cap["epochs_scored"] == 8 + assert on_the_cap["stop_reason"] == STOP_REASON_MAX_EPOCHS + assert on_the_cap["restored_epoch"] == 3 + + +def test_best_checkpoint_restored_after_early_stop(): + result = run_classify_schedule( + CANDIDATE, + max_epochs=12, + enabled=True, + patience=4, + minimum_epochs=4, + ) + assert result["stop_reason"] == STOP_REASON_EARLY_STOPPING + assert result["restored_epoch"] == result["best_epoch"] == 3 + assert result["restored_epoch"] != result["last_epoch"] + src = Path(loop.__file__).read_text(encoding="utf-8") + tree = ast.parse(src) + fn = next(node for node in tree.body if isinstance(node, ast.FunctionDef) and node.name == "run_loop") + break_line = None + restore_line = None + for node in ast.walk(fn): + if isinstance(node, ast.Call) and getattr(node.func, "id", None) == "early_stop_break": + break_line = node.lineno + if isinstance(node, ast.BoolOp): + text = ast.get_source_segment(src, node) or "" + if "best_state is not None" in text: + restore_line = node.lineno + assert break_line is not None and restore_line is not None + assert break_line < restore_line + + +@pytest.mark.parametrize( + ("env", "metric", "match"), + [ + ({"patience": "-1", "minimum": "4"}, SELECT_METRIC_CLASSIFY, "PATIENCE"), + ({"patience": "4", "minimum": "0"}, SELECT_METRIC_CLASSIFY, "MIN_EPOCHS"), + ({"patience": "4", "minimum": "-2"}, SELECT_METRIC_CLASSIFY, "MIN_EPOCHS"), + ({"patience": "4", "minimum": "13"}, SELECT_METRIC_CLASSIFY, "exceeds"), + ({"patience": "4.5", "minimum": "4"}, SELECT_METRIC_CLASSIFY, "integer"), + ({"patience": "4", "minimum": "nope"}, SELECT_METRIC_CLASSIFY, "integer"), + ({"patience": None, "minimum": "4"}, SELECT_METRIC_CLASSIFY, "required"), + ({"patience": "4", "minimum": None}, SELECT_METRIC_CLASSIFY, "required"), + ({"patience": "4", "minimum": "4"}, SELECT_METRIC_UNBIND, "classify_macro_f1_nonnone"), + ], +) +def test_invalid_configuration_fails_before_optimizer(monkeypatch, env, metric, match): + monkeypatch.setenv(EARLY_STOP_ENV, "1") + if env["patience"] is None: + monkeypatch.delenv(EARLY_STOP_PATIENCE_ENV, raising=False) + else: + monkeypatch.setenv(EARLY_STOP_PATIENCE_ENV, env["patience"]) + if env["minimum"] is None: + monkeypatch.delenv(EARLY_STOP_MIN_EPOCHS_ENV, raising=False) + else: + monkeypatch.setenv(EARLY_STOP_MIN_EPOCHS_ENV, env["minimum"]) + constructed: list[object] = [] + + def adamw(*args, **kwargs): + constructed.append(args) + raise AssertionError("optimizer constructed") + + with pytest.raises(ValueError, match=match): + resolve_early_stop_config(max_epochs=12, select_metric=metric) + adamw([object()], lr=2e-5) + assert constructed == [] + + src = Path(loop.__file__).read_text(encoding="utf-8") + tree = ast.parse(src) + fn = next(node for node in tree.body if isinstance(node, ast.FunctionDef) and node.name == "run_loop") + resolve_line = None + adam_line = None + for node in ast.walk(fn): + if not isinstance(node, ast.Call): + continue + name = getattr(node.func, "id", None) or getattr(node.func, "attr", None) + if name == "resolve_early_stop_config": + resolve_line = node.lineno + if name == "AdamW": + adam_line = node.lineno + assert resolve_line is not None and adam_line is not None + assert resolve_line < adam_line + + +def test_unknown_early_stop_token_is_rejected(monkeypatch): + monkeypatch.setenv(EARLY_STOP_ENV, "maybe") + with pytest.raises(ValueError, match="HYPERLEX_EARLY_STOP"): + resolve_early_stop_config(max_epochs=12, select_metric=SELECT_METRIC_CLASSIFY) + + +def test_explicit_off_runs_full_schedule(monkeypatch): + monkeypatch.setenv(EARLY_STOP_ENV, "off") + monkeypatch.setenv(EARLY_STOP_PATIENCE_ENV, "1") + monkeypatch.setenv(EARLY_STOP_MIN_EPOCHS_ENV, "1") + cfg = resolve_early_stop_config(max_epochs=12, select_metric=SELECT_METRIC_CLASSIFY) + assert cfg.enabled is False + result = run_classify_schedule( + CANDIDATE, + max_epochs=cfg.max_epochs, + enabled=cfg.enabled, + ) + assert result["epochs_scored"] == 12 + assert result["stop_reason"] == STOP_REASON_MAX_EPOCHS + + +def test_wallclock_fields_are_nonnegative_and_cumulative(monkeypatch): + ticks = iter([0.0, 1.0, 1.25, 1.25, 2.50, 2.50, 2.75]) + + def clock() -> float: + return next(ticks) + + monkeypatch.setattr(loop, "monotonic_seconds", clock) + training_started = loop.monotonic_seconds() + rows = [] + for epoch in range(3): + epoch_started = loop.monotonic_seconds() + row = { + "epoch": epoch, + "unbind_exact": 0.0, + "saved_best": False, + "classify_acc": 0.5, + } + row.update( + epoch_timing_fields( + epoch_started=epoch_started, + epoch_ended=loop.monotonic_seconds(), + training_started=training_started, + ) + ) + rows.append(row) + elapsed = [row["training_elapsed_seconds"] for row in rows] + assert elapsed == sorted(elapsed) + assert all(later >= earlier for earlier, later in zip(elapsed, elapsed[1:])) + for row in rows: + assert row["epoch_wallclock_seconds"] >= 0 + assert row["training_elapsed_seconds"] >= 0 + assert "epoch" in row and "saved_best" in row + assert rows[0]["epoch_wallclock_seconds"] == 0.25 + assert rows[1]["epoch_wallclock_seconds"] == 1.25 + assert rows[2]["epoch_wallclock_seconds"] == 0.25 + assert rows[-1]["training_elapsed_seconds"] == 2.75 + assert round_seconds(1.23456789) == round(1.23456789, 6) + assert inspect.signature(epoch_timing_fields).parameters.keys() >= { + "epoch_started", + "epoch_ended", + "training_started", + } + + +def test_wallclock_values_do_not_affect_selection(monkeypatch): + def boom() -> float: + raise AssertionError("selection read the clock") + + monkeypatch.setattr(loop, "monotonic_seconds", boom) + first = run_classify_schedule( + CANDIDATE, + max_epochs=12, + enabled=True, + patience=4, + minimum_epochs=4, + ) + second = run_classify_schedule( + CANDIDATE, + max_epochs=12, + enabled=True, + patience=4, + minimum_epochs=4, + ) + assert first == second + assert first["best_epoch"] == 3 + assert "epoch_wallclock_seconds" not in inspect.signature(note_strict_improvement).parameters + assert "training_elapsed_seconds" not in inspect.signature(early_stop_break).parameters + stamped = epoch_timing_fields(epoch_started=0.0, epoch_ended=50.0, training_started=0.0) + assert stamped["training_elapsed_seconds"] == 50.0 + assert first["best_epoch"] == 3 + + +def test_disabled_selection_keeps_strict_increase_and_earlier_ties(): + scores = [0.20, 0.50, 0.50, 0.40, 0.80, 0.80] + best_value = float("-inf") + best_epoch = None + improved_flags = [] + for epoch, score in enumerate(scores): + best_value, best_epoch, improved = note_strict_improvement( + best_value, best_epoch, score, epoch + ) + improved_flags.append(improved) + assert improved_flags == [True, True, False, False, True, False] + assert best_epoch == 4 + result = run_classify_schedule( + scores, + max_epochs=len(scores), + enabled=False, + ) + assert result["epochs_scored"] == len(scores) + assert result["stop_reason"] == STOP_REASON_MAX_EPOCHS + assert result["best_epoch"] == 4 + assert result["restored_epoch"] == 4 + + +def test_completion_receipt_records_only_loop_stop_reasons(): + src = Path(loop.__file__).read_text(encoding="utf-8") + assert '"stop_reason": stop_reason' in src + assert '"training_elapsed_seconds": training_elapsed_seconds' in src + tree = ast.parse(src) + fn = next(node for node in tree.body if isinstance(node, ast.FunctionDef) and node.name == "run_loop") + assigned: list[str] = [] + for node in ast.walk(fn): + if not isinstance(node, ast.Assign): + continue + for target in node.targets: + if isinstance(target, ast.Name) and target.id == "stop_reason": + assigned.append(ast.get_source_segment(src, node.value) or "") + assert assigned == ["STOP_REASON_MAX_EPOCHS", "STOP_REASON_EARLY_STOPPING"] + assert STOP_REASON_MAX_EPOCHS == "max_epochs" + assert STOP_REASON_EARLY_STOPPING == "early_stopping" + assert "stop_reason" not in inspect.signature(monotonic_seconds).parameters