feat: auto-checkpoint on SIGTERM/SIGINT with DDP support

- Register SIGTERM/SIGINT handlers in training loop, set stop flag on signal
- Check stop_requested at each epoch/batch boundary, break and call on_error to save checkpoint
- LocalStrategy parent forwards signal to child processes via terminate(), waits up to 600s for graceful exit
- TrainContext gains threading.Event-based stop_requested/request_stop
- Tests verify SIGTERM/SIGINT trigger checkpoint save with exit code 0, works on both CPU and GPU
This commit is contained in:
2026-07-25 20:40:54 +08:00
parent 8ab5631446
commit ceadc34ea9
5 changed files with 296 additions and 3 deletions
+177
View File
@@ -0,0 +1,177 @@
import json
import multiprocessing as mp
import os
import signal
import time
import pytest
import torch
import torch.optim as optim
from torch.utils.data import Dataset
from astrai.config import TrainConfig
from astrai.config.model_config import AutoRegressiveLMConfig
from astrai.model.transformer import AutoRegressiveLM
from astrai.parallel.signal_handler import register_signal_handlers
from astrai.trainer import Trainer
from astrai.trainer.schedule import SchedulerFactory
from astrai.trainer.train_context import TrainContext
class _PicklableDataset(Dataset):
def __init__(self, length=200, max_length=64, vocab_size=1000):
self.length = length
self.max_length = max_length
self.vocab_size = vocab_size
def __len__(self):
return self.length
def __getitem__(self, idx):
return {
"input_ids": torch.randint(0, self.vocab_size, (self.max_length,)),
"target_ids": torch.randint(0, self.vocab_size, (self.max_length,)),
}
def _build_model():
config = AutoRegressiveLMConfig(
vocab_size=1000,
hidden_size=8,
num_attention_heads=2,
num_key_value_heads=1,
intermediate_size=16,
max_position_embeddings=64,
num_hidden_layers=2,
rms_norm_eps=1e-5,
)
device = "cuda" if torch.cuda.is_available() else "cpu"
return AutoRegressiveLM(config).to(device=device)
class _ReadyCallback:
def __init__(self, ready_file):
self._ready_file = ready_file
def on_train_begin(self, context):
with open(self._ready_file, "w") as f:
f.write("ready")
f.flush()
os.fsync(f.fileno())
def _inner_run(batch_per_device, ckpt_interval, ckpt_dir, log_dir, ready_file):
dataset = _PicklableDataset()
def model_fn():
return _build_model()
def optimizer_fn(m):
return optim.AdamW(m.parameters(), lr=0.001)
def scheduler_fn(optim):
return SchedulerFactory.create(
"cosine", optim, warmup_steps=10, lr_decay_steps=10, min_rate=0.05
)
train_config = TrainConfig(
strategy="seq",
model_fn=model_fn,
dataset=dataset,
optimizer_fn=optimizer_fn,
scheduler_fn=scheduler_fn,
ckpt_dir=ckpt_dir,
log_dir=log_dir,
n_epoch=1,
batch_per_device=batch_per_device,
ckpt_interval=ckpt_interval,
grad_accum_steps=1,
random_seed=42,
device_type="cuda" if torch.cuda.is_available() else "cpu",
)
trainer = Trainer(train_config)
trainer.callbacks.insert(0, _ReadyCallback(ready_file))
trainer.train()
def _spawn_train_and_signal(ckpt_dir, sig, timeout=120):
log_dir = os.path.join(ckpt_dir, "logs")
ready_file = os.path.join(ckpt_dir, "ready.txt")
ctx = mp.get_context("spawn")
p = ctx.Process(
target=_inner_run,
args=(2, 1000, ckpt_dir, log_dir, ready_file),
)
p.start()
deadline = time.time() + 30
while time.time() < deadline:
if os.path.exists(ready_file):
with open(ready_file) as f:
if f.read().strip() == "ready":
break
if not p.is_alive():
break
time.sleep(0.5)
assert p.is_alive(), "Training process died before becoming ready"
os.kill(p.pid, sig)
p.join(timeout=timeout)
if p.is_alive():
p.kill()
p.join(timeout=5)
return p.exitcode
def test_context_stop_flag():
ctx = TrainContext()
assert not ctx.stop_requested
ctx.request_stop()
assert ctx.stop_requested
def test_register_signal_handlers():
ctx = TrainContext()
register_signal_handlers(ctx)
assert not ctx.stop_requested
os.kill(os.getpid(), signal.SIGTERM)
assert ctx.stop_requested
def test_sigterm_triggers_checkpoint_save(base_test_env):
exitcode = _spawn_train_and_signal(base_test_env["test_dir"], signal.SIGTERM)
assert exitcode == 0, f"Training process exited with code {exitcode} (expected 0)"
ckpt_dir = base_test_env["test_dir"]
meta_files = []
for root, dirs, files in os.walk(ckpt_dir):
for f in files:
if f == "meta.json":
meta_files.append(os.path.join(root, f))
assert len(meta_files) > 0, f"No checkpoint meta.json found in {ckpt_dir}"
with open(meta_files[-1]) as f:
meta = json.load(f)
assert "consumed_samples" in meta
assert meta["consumed_samples"] >= 0
@pytest.mark.slow
def test_sigint_triggers_checkpoint_save(base_test_env):
exitcode = _spawn_train_and_signal(base_test_env["test_dir"], signal.SIGINT)
assert exitcode == 0, f"Training process exited with code {exitcode} (expected 0)"
ckpt_dir = base_test_env["test_dir"]
meta_files = []
for root, dirs, files in os.walk(ckpt_dir):
for f in files:
if f == "meta.json":
meta_files.append(os.path.join(root, f))
assert len(meta_files) > 0, f"No checkpoint meta.json found in {ckpt_dir}"