feat: 优化工厂模式的实现
This commit is contained in:
@@ -1,12 +1,8 @@
|
||||
from astrai.trainer.schedule import BaseScheduler, SchedulerFactory
|
||||
from astrai.trainer.strategy import BaseStrategy, StrategyFactory
|
||||
from astrai.trainer.train_callback import (
|
||||
CheckpointCallback,
|
||||
GradientClippingCallback,
|
||||
MetricLoggerCallback,
|
||||
ProgressBarCallback,
|
||||
SchedulerCallback,
|
||||
TrainCallback,
|
||||
CallbackFactory,
|
||||
)
|
||||
from astrai.trainer.trainer import Trainer
|
||||
|
||||
@@ -19,11 +15,7 @@ __all__ = [
|
||||
# Scheduler factory
|
||||
"SchedulerFactory",
|
||||
"BaseScheduler",
|
||||
# Callbacks
|
||||
# Callback factory
|
||||
"TrainCallback",
|
||||
"GradientClippingCallback",
|
||||
"SchedulerCallback",
|
||||
"CheckpointCallback",
|
||||
"ProgressBarCallback",
|
||||
"MetricLoggerCallback",
|
||||
"CallbackFactory",
|
||||
]
|
||||
|
||||
@@ -6,7 +6,7 @@ from typing import Any, Dict, List, Type
|
||||
|
||||
from torch.optim.lr_scheduler import LRScheduler
|
||||
|
||||
from astrai.core.factory import BaseFactory
|
||||
from astrai.factory import BaseFactory
|
||||
|
||||
|
||||
class BaseScheduler(LRScheduler, ABC):
|
||||
@@ -41,8 +41,6 @@ class SchedulerFactory(BaseFactory["BaseScheduler"]):
|
||||
scheduler = SchedulerFactory.create("custom", optimizer, **kwargs)
|
||||
"""
|
||||
|
||||
_registry: Dict[str, Type[BaseScheduler]] = {}
|
||||
|
||||
@classmethod
|
||||
def _validate_component(cls, scheduler_cls: Type[BaseScheduler]) -> None:
|
||||
"""Validate that the scheduler class inherits from BaseScheduler."""
|
||||
|
||||
@@ -10,7 +10,7 @@ import torch.nn.functional as F
|
||||
from torch import Tensor
|
||||
from torch.nn.parallel import DistributedDataParallel as DDP
|
||||
|
||||
from astrai.core.factory import BaseFactory
|
||||
from astrai.factory import BaseFactory
|
||||
|
||||
|
||||
def unwrap_model(model: nn.Module) -> nn.Module:
|
||||
@@ -122,8 +122,6 @@ class StrategyFactory(BaseFactory["BaseStrategy"]):
|
||||
strategy = StrategyFactory.create("custom", model, device)
|
||||
"""
|
||||
|
||||
_registry: Dict[str, type] = {}
|
||||
|
||||
@classmethod
|
||||
def _validate_component(cls, strategy_cls: type) -> None:
|
||||
"""Validate that the strategy class inherits from BaseStrategy."""
|
||||
|
||||
@@ -2,7 +2,7 @@ import json
|
||||
import os
|
||||
import time
|
||||
from pathlib import Path
|
||||
from typing import Callable, List, Optional, Protocol
|
||||
from typing import Callable, List, Optional, Protocol, runtime_checkable
|
||||
|
||||
import torch.nn as nn
|
||||
from torch.nn.utils import clip_grad_norm_
|
||||
@@ -21,8 +21,10 @@ from astrai.trainer.metric_util import (
|
||||
ctx_get_lr,
|
||||
)
|
||||
from astrai.trainer.train_context import TrainContext
|
||||
from astrai.factory import BaseFactory
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class TrainCallback(Protocol):
|
||||
"""
|
||||
Callback interface for trainer.
|
||||
@@ -56,6 +58,25 @@ class TrainCallback(Protocol):
|
||||
"""Called when an error occurs during training."""
|
||||
|
||||
|
||||
class CallbackFactory(BaseFactory[TrainCallback]):
|
||||
"""Factory for registering and creating training callbacks.
|
||||
|
||||
Example:
|
||||
@CallbackFactory.register("my_callback")
|
||||
class MyCallback(TrainCallback):
|
||||
...
|
||||
|
||||
callback = CallbackFactory.create("my_callback", **kwargs)
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def _validate_component(cls, callback_cls: type) -> None:
|
||||
"""Validate that the callback class inherits from TrainCallback."""
|
||||
if not issubclass(callback_cls, TrainCallback):
|
||||
raise TypeError(f"{callback_cls.__name__} must inherit from TrainCallback")
|
||||
|
||||
|
||||
@CallbackFactory.register("gradient_clipping")
|
||||
class GradientClippingCallback(TrainCallback):
|
||||
"""
|
||||
Gradient clipping callback for trainer.
|
||||
@@ -69,6 +90,7 @@ class GradientClippingCallback(TrainCallback):
|
||||
clip_grad_norm_(context.model.parameters(), self.max_grad_norm)
|
||||
|
||||
|
||||
@CallbackFactory.register("scheduler")
|
||||
class SchedulerCallback(TrainCallback):
|
||||
"""
|
||||
Scheduler callback for trainer.
|
||||
@@ -87,6 +109,7 @@ class SchedulerCallback(TrainCallback):
|
||||
context.scheduler.step()
|
||||
|
||||
|
||||
@CallbackFactory.register("checkpoint")
|
||||
class CheckpointCallback(TrainCallback):
|
||||
"""
|
||||
Checkpoint callback for trainer.
|
||||
@@ -135,6 +158,7 @@ class CheckpointCallback(TrainCallback):
|
||||
self._save_checkpoint(context)
|
||||
|
||||
|
||||
@CallbackFactory.register("progress_bar")
|
||||
class ProgressBarCallback(TrainCallback):
|
||||
"""
|
||||
Progress bar callback for trainer.
|
||||
@@ -169,6 +193,7 @@ class ProgressBarCallback(TrainCallback):
|
||||
self.progress_bar.close()
|
||||
|
||||
|
||||
@CallbackFactory.register("metric_logger")
|
||||
class MetricLoggerCallback(TrainCallback):
|
||||
def __init__(
|
||||
self,
|
||||
|
||||
@@ -5,12 +5,8 @@ from astrai.config import TrainConfig
|
||||
from astrai.data.serialization import Checkpoint
|
||||
from astrai.parallel.setup import spawn_parallel_fn
|
||||
from astrai.trainer.train_callback import (
|
||||
CheckpointCallback,
|
||||
GradientClippingCallback,
|
||||
MetricLoggerCallback,
|
||||
ProgressBarCallback,
|
||||
SchedulerCallback,
|
||||
TrainCallback,
|
||||
CallbackFactory,
|
||||
)
|
||||
from astrai.trainer.train_context import TrainContext, TrainContextBuilder
|
||||
|
||||
@@ -28,13 +24,13 @@ class Trainer:
|
||||
)
|
||||
|
||||
def _get_default_callbacks(self) -> List[TrainCallback]:
|
||||
train_config = self.train_config
|
||||
cfg = self.train_config
|
||||
return [
|
||||
ProgressBarCallback(train_config.n_epoch),
|
||||
CheckpointCallback(train_config.ckpt_dir, train_config.ckpt_interval),
|
||||
MetricLoggerCallback(train_config.ckpt_dir, train_config.ckpt_interval),
|
||||
GradientClippingCallback(train_config.max_grad_norm),
|
||||
SchedulerCallback(),
|
||||
CallbackFactory.create("progress_bar", cfg.n_epoch),
|
||||
CallbackFactory.create("checkpoint", cfg.ckpt_dir, cfg.ckpt_interval),
|
||||
CallbackFactory.create("metric_logger", cfg.ckpt_dir, cfg.ckpt_interval),
|
||||
CallbackFactory.create("gradient_clipping", cfg.max_grad_norm),
|
||||
CallbackFactory.create("scheduler"),
|
||||
]
|
||||
|
||||
def _build_context(self, checkpoint: Optional[Checkpoint]) -> TrainContext:
|
||||
|
||||
Reference in New Issue
Block a user