Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
202 changes: 198 additions & 4 deletions scripts/shadow/hyperlexical/loop.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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,
Expand Down Expand Up @@ -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
Expand All @@ -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).

Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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] = []
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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()
Expand All @@ -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
)
Expand Down Expand Up @@ -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,
Expand Down
15 changes: 15 additions & 0 deletions specs/007-hyperlexical-model/AARON-SPARK-TRAIN.md
Original file line number Diff line number Diff line change
Expand Up @@ -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)):
Expand Down
Loading
Loading