refactor: move signal_handler from parallel/ to top-level for broader reuse

This commit is contained in:
2026-07-28 10:36:17 +08:00
parent 9f7cf50c56
commit 39f84f3b4c
4 changed files with 3 additions and 3 deletions
+1 -1
View File
@@ -12,7 +12,7 @@ import torch
import torch.distributed as dist
import torch.multiprocessing as mp
from astrai.parallel.signal_handler import install_early_signal_handlers
from astrai.signal_handler import install_early_signal_handlers
logger = logging.getLogger(__name__)
-53
View File
@@ -1,53 +0,0 @@
import logging
import os
import signal
import threading
logger = logging.getLogger(__name__)
_early_stop = threading.Event()
_active_context = None
def _early_handler(signum: int, frame):
sig = signal.Signals(signum)
logger.warning(
"Received %s (pid=%d), requesting graceful training stop...",
sig.name,
os.getpid(),
)
_early_stop.set()
if _active_context is not None:
_active_context.request_stop()
def install_early_signal_handlers():
for sig in (signal.SIGTERM, signal.SIGINT):
signal.signal(sig, _early_handler)
_unblock_signals()
def _unblock_signals():
try:
mask = signal.pthread_sigmask(signal.SIG_BLOCK, set())
blocked = {signal.SIGTERM, signal.SIGINT} & mask
if blocked:
signal.pthread_sigmask(signal.SIG_UNBLOCK, blocked)
except (AttributeError, OSError):
pass
def register_signal_handlers(context):
global _active_context
_active_context = context
for sig in (signal.SIGTERM, signal.SIGINT):
signal.signal(sig, _early_handler)
if _early_stop.is_set():
context.request_stop()
logger.warning("Signal was received during initialization, stopping...")
def unregister_signal_handlers():
global _active_context
_active_context = None
_early_stop.clear()