Skip to content

trainer

learning.trainer

logger module-attribute

logger = getLogger('SaigeResearch')

BaseTrainer

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

_config: dict

version instance-attribute

version: Version

config property

config: dict

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

train_one_step() -> dict
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

validate() -> dict
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

build(config: dict)
Source code in SaigeToolkit/learning/trainer.py
@classmethod
def build(cls, config: dict):
    init_cuda_when_available()
    return cls(**config)