refactor: move signal_handler from parallel/ to top-level for broader reuse
This commit is contained in:
@@ -12,7 +12,7 @@ import torch
|
|||||||
import torch.distributed as dist
|
import torch.distributed as dist
|
||||||
import torch.multiprocessing as mp
|
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__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|||||||
@@ -5,7 +5,7 @@ import torch.distributed as dist
|
|||||||
|
|
||||||
from astrai.config import TrainConfig
|
from astrai.config import TrainConfig
|
||||||
from astrai.parallel.setup import spawn_parallel_fn
|
from astrai.parallel.setup import spawn_parallel_fn
|
||||||
from astrai.parallel.signal_handler import (
|
from astrai.signal_handler import (
|
||||||
register_signal_handlers,
|
register_signal_handlers,
|
||||||
unregister_signal_handlers,
|
unregister_signal_handlers,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -10,7 +10,7 @@ from torch.utils.data import Dataset
|
|||||||
|
|
||||||
from astrai.config import TrainConfig
|
from astrai.config import TrainConfig
|
||||||
from astrai.model.transformer import AutoRegressiveLM
|
from astrai.model.transformer import AutoRegressiveLM
|
||||||
from astrai.parallel.signal_handler import register_signal_handlers
|
from astrai.signal_handler import register_signal_handlers
|
||||||
from astrai.trainer import Trainer
|
from astrai.trainer import Trainer
|
||||||
from astrai.trainer.schedule import SchedulerFactory
|
from astrai.trainer.schedule import SchedulerFactory
|
||||||
from astrai.trainer.train_context import TrainContext
|
from astrai.trainer.train_context import TrainContext
|
||||||
|
|||||||
Reference in New Issue
Block a user