Bases: ABC
Saige Vision2 Engine 학습 API의 베이스 인터페이스 입니다.
기본적인 규약만 정해져있으며 각 태스크에 맞게 abstractmethod들을 구현하고, error_handler를 씌워서 노출시키면 됩니다.
About Init
init 메소드 작성 시 util.reproducibility.store_config 데코레이터를 활용하면 _config를 쉽게 저장할 수 있습니다.
Example:
class Trainer(BaseTrainer):
@store_config(attr="_config") # 인스턴스 생성 시 self._config 변수에 생성 파라미터들 저장됨
def init(self, param1, param2):
...
About Checkpoint
체크포인트를 저장/로드 하는 인터페이스는 save_checkpoint, load_checkpoint로 이미 구현되어 있으며,
태스크에 맞게 state_dict, load_state_dict를 구현하면 해당 함수를 이용해 상태를 저장/로드 하는 방식입니다.
_config
instance-attribute
version
instance-attribute
to_device
abstractmethod
to_device(device: Union[device, int, str]) -> None
Source code in SaigeToolkit/learning/trainer.py
| @abstractmethod
def to_device(self, device: Union[torch.device, int, str]) -> None:
pass
|
train_one_step
abstractmethod
Source code in SaigeToolkit/learning/trainer.py
| @abstractmethod
def train_one_step(self) -> dict:
pass
|
validation_enabled
abstractmethod
validation_enabled() -> bool
Source code in SaigeToolkit/learning/trainer.py
| @abstractmethod
def validation_enabled(self) -> bool:
pass
|
validate
abstractmethod
Source code in SaigeToolkit/learning/trainer.py
| @abstractmethod
def validate(self) -> dict:
pass
|
state_dict
abstractmethod
state_dict(**options) -> dict
Source code in SaigeToolkit/learning/trainer.py
| @abstractmethod
def state_dict(self, **options) -> dict:
pass
|
load_state_dict
abstractmethod
load_state_dict(state_dict: dict, version: Optional[Version] = None, **options) -> None
Source code in SaigeToolkit/learning/trainer.py
| @abstractmethod
def load_state_dict(self, state_dict: dict, version: Optional[Version] = None, **options) -> None:
pass
|
save_checkpoint
save_checkpoint(checkpoint_path: str, metadata: Optional[dict] = None, password: Optional[str] = None, **options)
Source code in SaigeToolkit/learning/trainer.py
| def save_checkpoint(
self,
checkpoint_path: str,
metadata: Optional[dict] = None,
password: Optional[str] = None,
**options,
):
CheckpointHandler.save(
checkpoint_path=checkpoint_path,
password=password,
version=str(self.version),
config=self.config,
state_dict=self.state_dict(**options),
metadata=metadata,
)
logger.info(f"{type(self).__name__} checkpoint saved: {checkpoint_path}")
|
load_checkpoint
load_checkpoint(checkpoint_path: str, password: Optional[str] = None, **options) -> Dict
Source code in SaigeToolkit/learning/trainer.py
| def load_checkpoint(
self,
checkpoint_path: str,
password: Optional[str] = None,
**options,
) -> Dict:
checkpoint = CheckpointHandler.load(
checkpoint_path=checkpoint_path,
password=password,
)
version = Version.from_string(checkpoint.get("version", "0.0.0"))
self.load_state_dict(checkpoint.get("state_dict"), version=version, **options)
logger.info(f"{type(self).__name__} loaded from {checkpoint_path}")
return checkpoint.pop("metadata", {})
|
build
classmethod
Source code in SaigeToolkit/learning/trainer.py
| @classmethod
def build(cls, config: dict):
init_cuda_when_available()
return cls(**config)
|