Files
AstrAI/tests/trainer/test_early_stopping.py
T
ViperEkura 5ba21f4eb3 refactor: eliminate test duplication via shared helpers
- Add tests/helpers.py with shared config, dataset, tokenizer, executor, and assertion helpers
- Replace 15 copies of device one-liner with session-scoped fixture
- Collapse 5 near-identical Dataset subclasses into RandomTokenDataset
- Remove duplicate _make_config/_make_model/_make_frozen and FakeTokenizer/FakeExecutor definitions
- Make test_callbacks and test_early_stopping use existing train_config_factory
- Replace 6 duplicate meta.json read blocks with load_shard_meta
- Fix mkdtemp leaks in test_lora.py with TemporaryDirectory
2026-07-27 22:34:53 +08:00

37 lines
1.0 KiB
Python

import os
from astrai.trainer.trainer import Trainer
from tests.helpers import load_checkpoint_meta
def test_early_stopping_simulation(
base_test_env, early_stopping_dataset, train_config_factory, device
):
"""Simulate early stopping behavior"""
train_config = train_config_factory(
model_fn=lambda: base_test_env["model"],
dataset=early_stopping_dataset,
test_dir=base_test_env["test_dir"],
device=device,
n_epoch=2,
ckpt_interval=1,
grad_accum_steps=2,
)
trainer = Trainer(train_config)
try:
trainer.train()
except Exception:
pass
# Resume from latest checkpoint
load_dir = os.path.join(base_test_env["test_dir"], "epoch_0_step_1")
trainer = Trainer(train_config)
trainer.train(param_path=load_dir, resume=True)
# Verify checkpoint was saved at expected step
load_dir = os.path.join(base_test_env["test_dir"], "epoch_1_step_5")
meta = load_checkpoint_meta(load_dir)
assert meta["consumed_samples"] == 20