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.
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.3
2
2
  Name: trainloop
3
- Version: 0.10.0
3
+ Version: 0.11.0
4
4
  Summary: Minimal PyTorch training loop with hooks and checkpointing.
5
5
  Author: Karim Knaebel
6
6
  Author-email: Karim Knaebel <contact@knaebel.dev>
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "trainloop"
3
- version = "0.10.0"
3
+ version = "0.11.0"
4
4
  description = "Minimal PyTorch training loop with hooks and checkpointing."
5
5
  readme = "README.md"
6
6
  requires-python = ">=3.10"
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "trainloop"
3
- version = "0.10.0"
3
+ version = "0.11.0"
4
4
  description = "Minimal PyTorch training loop with hooks and checkpointing."
5
5
  readme = "README.md"
6
6
  authors = [
@@ -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__(self, decay: float = 0.999, use_buffers: bool = False):
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["ema_model"] = get_model_state_dict(self.ema_model)
630
+ state_dict[self.name] = get_model_state_dict(self.ema_model)
627
631
 
628
632
 
629
633
  class WandbHook(BaseHook):
File without changes