trainloop 0.10.0__tar.gz → 0.11.0__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- {trainloop-0.10.0 → trainloop-0.11.0}/PKG-INFO +1 -1
- {trainloop-0.10.0 → trainloop-0.11.0}/pyproject.toml +1 -1
- {trainloop-0.10.0 → trainloop-0.11.0}/pyproject.toml.orig +1 -1
- {trainloop-0.10.0 → trainloop-0.11.0}/src/trainloop/hooks.py +11 -7
- {trainloop-0.10.0 → trainloop-0.11.0}/README.md +0 -0
- {trainloop-0.10.0 → trainloop-0.11.0}/src/trainloop/__init__.py +0 -0
- {trainloop-0.10.0 → trainloop-0.11.0}/src/trainloop/py.typed +0 -0
- {trainloop-0.10.0 → trainloop-0.11.0}/src/trainloop/trainer.py +0 -0
- {trainloop-0.10.0 → trainloop-0.11.0}/src/trainloop/utils.py +0 -0
|
@@ -590,14 +590,18 @@ class EMAHook(BaseHook):
|
|
|
590
590
|
Args:
|
|
591
591
|
decay: EMA decay rate.
|
|
592
592
|
use_buffers: Whether to include model buffers in the EMA.
|
|
593
|
+
name: Name used for the EMA state in the trainer state dict.
|
|
593
594
|
"""
|
|
594
595
|
|
|
595
|
-
def __init__(
|
|
596
|
+
def __init__(
|
|
597
|
+
self, decay: float = 0.999, use_buffers: bool = False, name: str = "ema"
|
|
598
|
+
):
|
|
596
599
|
self.decay = decay
|
|
597
600
|
self.use_buffers = use_buffers
|
|
601
|
+
self.name = name
|
|
598
602
|
|
|
599
603
|
def on_before_train(self, trainer: BaseTrainer):
|
|
600
|
-
trainer.logger.info("=> Creating EMA model ...")
|
|
604
|
+
trainer.logger.info(f"=> Creating EMA model {self.name!r} ...")
|
|
601
605
|
# Note that AveragedModel does not seem to support FSDP. It will crash here.
|
|
602
606
|
self.ema_model = AveragedModel(
|
|
603
607
|
trainer.model,
|
|
@@ -610,10 +614,8 @@ class EMAHook(BaseHook):
|
|
|
610
614
|
self.ema_model.update_parameters(trainer.model)
|
|
611
615
|
|
|
612
616
|
def on_load_state_dict(self, trainer: BaseTrainer, state_dict: dict):
|
|
613
|
-
trainer.logger.info("=> Loading EMA model state ...")
|
|
614
|
-
incompatible_keys = set_model_state_dict(
|
|
615
|
-
self.ema_model, state_dict["ema_model"]
|
|
616
|
-
)
|
|
617
|
+
trainer.logger.info(f"=> Loading EMA model {self.name!r} state ...")
|
|
618
|
+
incompatible_keys = set_model_state_dict(self.ema_model, state_dict[self.name])
|
|
617
619
|
# This currently doesn't do anything because strict=True is implicit.
|
|
618
620
|
log_state_dict_incompatible_keys(
|
|
619
621
|
trainer.logger,
|
|
@@ -622,8 +624,10 @@ class EMAHook(BaseHook):
|
|
|
622
624
|
)
|
|
623
625
|
|
|
624
626
|
def on_state_dict(self, trainer: BaseTrainer, state_dict: dict):
|
|
627
|
+
if self.name in state_dict:
|
|
628
|
+
raise ValueError(f"State dict key {self.name!r} already exists")
|
|
625
629
|
# Note: sadly, we need to keep the AveragedModel wrapper, to save its n_averaged buffer
|
|
626
|
-
state_dict[
|
|
630
|
+
state_dict[self.name] = get_model_state_dict(self.ema_model)
|
|
627
631
|
|
|
628
632
|
|
|
629
633
|
class WandbHook(BaseHook):
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|