torch_em.trainer.default_trainer

  1from __future__ import annotations
  2
  3import os
  4import time
  5import inspect
  6import warnings
  7import contextlib
  8from tqdm import tqdm
  9from functools import partial
 10from datetime import datetime
 11from collections import OrderedDict
 12from importlib import import_module
 13from typing import Any, Callable, Dict, Optional, Union, Literal
 14
 15import numpy as np
 16
 17import torch
 18
 19from .wandb_logger import WandbLogger
 20from .tensorboard_logger import TensorboardLogger
 21from ..util import auto_compile, get_constructor_arguments, is_compiled
 22
 23
 24class DefaultTrainer:
 25    """Trainer class for training a segmentation network.
 26
 27    The trainer class implements the core logic for training a network with pytorch.
 28    It implements a training loop to run training and validation, which is started with `fit`.
 29    The checkpoints and logs from the training run will be saved in the current working directory,
 30    or in the directory specifified by `save_root`. Training can be continued from a checkpoint
 31    by passing it's location to the `load_from_checkpoint` argument of `fit`.
 32
 33    A pre-configured instance of the trainer can be obtained from `torch_em.default_segmentation_trainer`.
 34    Alternatively, the trainer class can also be instantiated as in this example:
 35    ```python
 36    import torch
 37    from torch_em.loss import DiceLoss
 38    from torch_em.model import UNet2d
 39    from torch_em.data.datasets.light_microscopy import get_dsb_loader
 40    from torch_em.trainer import DefaultTrainer
 41
 42    # The training data will be downloaded to this location.
 43    data_root = "/path/to/save/the/training/data"
 44    patch_shape = (256, 256)
 45
 46    # Create the model and optimizer.
 47    model = UNet2d(in_channels=1, out_channels=1)
 48    optimizer = torch.optim.AdamW(model.parameters())
 49
 50    trainer = DefaultTrainer(
 51        name="unet-training",
 52        train_loader=get_dsb_loader(path=data_root, patch_shape=patch_shape, split="train"),
 53        val_loader=get_dsb_loader(path=data_root, patch_shape=patch_shape, split="test"),
 54        model=model,
 55        loss=DiceLoss(),  # The loss function.
 56        optimizer=optimizer,
 57        metric=DiceLoss(),  # The metric. The trainer expects smaller values to represent better results.
 58        device="cuda",  # The device to use for training.
 59    )
 60    trainer.fit(iterations=int(2.5e4))  # Train for 25.000 iterations.
 61    ```
 62
 63    Args:
 64        name: The name of the checkpoint that will be created by the trainer.
 65        train_loader: The data loader containing the training data.
 66        val_loader: The data loader containing the validation data.
 67        model: The model to train.
 68        loss: The loss function for training.
 69        optimizer: The optimizer.
 70        metric: The metric for validation.
 71        device: The torch device to use for training. If None, will use a GPU if available.
 72        lr_scheduler: The learning rate scheduler.
 73        log_image_interval: The interval for saving images during logging, in training iterations.
 74        mixed_precision: Whether to train with mixed precision.
 75        early_stopping: The patience for early stopping in epochs. If None, early stopping will not be used.
 76        logger: The logger class. Will be instantiated for logging.
 77            By default uses `torch_em.training.tensorboard_logger.TensorboardLogger`.
 78        logger_kwargs: The keyword arguments for the logger class.
 79        id_: Unique identifier for the trainer. If None then `name` will be used.
 80        save_root: The root folder for saving the checkpoint and logs.
 81        compile_model: Whether to compile the model before training.
 82        rank: Rank argument for distributed training. See `torch_em.multi_gpu_training` for details.
 83        mixed_precision_dtype: The dtype for autocast in mixed precision training, 'float16' or 'bfloat16'.
 84            The default is 'float16' on the GPU and 'bfloat16' on the CPU. Use 'bfloat16' to avoid overflows.
 85    """
 86    def __init__(
 87        self,
 88        name: Optional[str],
 89        train_loader: torch.utils.data.DataLoader,
 90        val_loader: torch.utils.data.DataLoader,
 91        model: torch.nn.Module,
 92        loss: torch.nn.Module,
 93        optimizer: torch.optim.Optimizer,
 94        metric: Callable,
 95        device: Union[str, torch.device],
 96        lr_scheduler: Optional[torch.optim.lr_scheduler._LRScheduler] = None,
 97        log_image_interval: int = 100,
 98        mixed_precision: bool = True,
 99        early_stopping: Optional[int] = None,
100        logger=TensorboardLogger,
101        logger_kwargs: Optional[Dict[str, Any]] = None,
102        id_: Optional[str] = None,
103        save_root: Optional[str] = None,
104        compile_model: Optional[Union[bool, str]] = None,
105        rank: Optional[int] = None,
106        mixed_precision_dtype: Optional[str] = None,
107    ):
108        if name is None and not issubclass(logger, WandbLogger):
109            raise TypeError("Name cannot be None if not using the WandbLogger")
110
111        self._generate_name = name is None
112        self.name = name
113        self.id_ = id_ or name
114        self.train_loader = train_loader
115        self.val_loader = val_loader
116        self.model = model
117        self.loss = loss
118        self.optimizer = optimizer
119        self.metric = metric
120        self.device = torch.device(device)
121        self.lr_scheduler = lr_scheduler
122        self.log_image_interval = log_image_interval
123        self.save_root = save_root
124        self.compile_model = compile_model
125        self.rank = rank
126        self._device_type = "cpu" if self.device.type == "cpu" else "cuda"
127
128        self._iteration = 0
129        self._epoch = 0
130        self._best_epoch = 0
131
132        self.mixed_precision = mixed_precision
133        # These are the defaults of torch.autocast for each device type.
134        self.mixed_precision_dtype = mixed_precision_dtype or ("bfloat16" if self._device_type == "cpu" else "float16")
135        self.early_stopping = early_stopping
136        self.train_time = 0.0
137
138        if mixed_precision:
139            # Only float16 needs gradient scaling. bfloat16 has the same range as float32.
140            self.scaler = torch.GradScaler(self._device_type, enabled=self.mixed_precision_dtype == "float16")
141        else:
142            self.scaler = None
143
144        self.logger_class = logger
145        self.logger_kwargs = logger_kwargs
146        self.log_image_interval = log_image_interval
147
148    @property
149    def checkpoint_folder(self):
150        assert self.id_ is not None  # Because the logger may generate and set trainer.id on logger.__init__.
151        # Save_root enables saving the checkpoints somewhere else than in the local folder.
152        # This is handy for filesystems with limited space, where saving the checkpoints
153        # and log files can lead to running out of space.
154        save_root = getattr(self, "save_root", None)
155        return os.path.join("./checkpoints", self.id_) if save_root is None else\
156            os.path.join(save_root, "./checkpoints", self.id_)
157
158    @property
159    def iteration(self):
160        return self._iteration
161
162    @property
163    def epoch(self):
164        return self._epoch
165
166    class Deserializer:
167        """Determines how to deserialize the trainer kwargs from serialized 'init_data'.
168
169        Examples:
170            To extend the initialization process you can inherite from this Deserializer in an inherited Trainer class.
171            Note that `DefaultTrainer.Deserializer.load_generic()` covers most cases already.
172
173            This example adds `the_answer` kwarg, which requires 'calculations' upon initialization:
174            >>> class MyTrainer(DefaultTrainer):
175            >>>     def __init__(self, *args, the_answer: int, **kwargs):
176            >>>         super().__init__(*args, **kwargs)
177            >>>         self.the_answer = the_answer  # this allows the default Serializer to save the new kwarg,
178            >>>                                       # see DefaultTrainer.Serializer
179            >>>
180            >>>     class Deserializer(DefaultTrainer.Deserializer):
181            >>>         def load_the_answer(self):
182            >>>             generic_answer = self.init_data["the_answer"]
183            >>>             # (device dependent) special deserialization
184            >>>             if self.trainer_kwargs["device"].type == "cpu":  # accessing previously deserialized kwarg
185            >>>                 self.trainer_kwargs["the_answer"] = generic_answer + 1
186            >>>             else:
187            >>>                 self.trainer_kwargs["the_answer"] = generic_answer * 2
188
189        Args:
190            init_data: The initialization data of the trainer.
191            save_path: The path where the checkpoint was saved.
192            device: The device.
193        """
194
195        def __init__(self, init_data: Dict, save_path: str, device: Union[str, torch.device]):
196            self.init_data = init_data
197            self.save_path = save_path
198            # Populate with deserialized trainer kwargs during deserialization; possibly overwrite 'device'.
199            self.trainer_kwargs: Dict[str, Any] = dict(
200                device=torch.device(self.init_data["device"]) if device is None else torch.device(device)
201            )
202
203        def load(self, kwarg_name: str, optional):
204            """@private
205            """
206            # `optional` is True if self.trainer.__class__.__init__ specifies a default value for 'kwarg_name'
207            if kwarg_name == "device":
208                pass  # deserialized in __init__
209            elif kwarg_name.endswith("_loader"):
210                self.load_data_loader(kwarg_name, optional)
211            else:
212                load = getattr(self, f"load_{kwarg_name}", self.load_generic)
213                load(kwarg_name, optional=optional)
214
215        def load_data_loader(self, loader_name, optional) -> None:
216            """@private
217            """
218            ds = self.init_data.get(loader_name.replace("_loader", "_dataset"))
219            if ds is None and optional:
220                return
221
222            loader_kwargs = self.init_data[f"{loader_name}_kwargs"]
223            loader = torch.utils.data.DataLoader(ds, **loader_kwargs)
224            # monkey patch shuffle loader_name to the loader
225            loader.shuffle = loader_kwargs.get("shuffle", False)
226            self.trainer_kwargs[loader_name] = loader
227
228        def load_generic(
229            self,
230            kwarg_name: str,
231            *dynamic_args: Dict,
232            optional: bool,
233            only_class: bool = False,
234            dynamic_kwargs: Optional[Dict[str, Any]] = None,
235        ) -> None:
236            """@private
237            """
238            if kwarg_name in self.init_data:
239                self.trainer_kwargs[kwarg_name] = self.init_data[kwarg_name]
240                return
241
242            this_cls = self.init_data.get(f"{kwarg_name}_class", None)
243            if this_cls is None:
244                if optional:
245                    return
246                else:
247                    raise RuntimeError(f"Could not find init data for {kwarg_name} in {self.save_path}")
248
249            assert isinstance(this_cls, str), this_cls
250            assert "." in this_cls, this_cls
251            cls_p, cls_m = this_cls.rsplit(".", 1)
252            this_cls = getattr(import_module(cls_p), cls_m)
253            if only_class:
254                self.trainer_kwargs[kwarg_name] = this_cls
255            else:
256                self.trainer_kwargs[kwarg_name] = this_cls(
257                    *dynamic_args, **self.init_data.get(f"{kwarg_name}_kwargs", {}), **(dynamic_kwargs or {})
258                )
259
260        def load_name(self, kwarg_name: str, optional: bool):
261            """@private
262            """
263            self.trainer_kwargs[kwarg_name] = os.path.split(os.path.dirname(self.save_path))[1]
264
265        def load_optimizer(self, kwarg_name: str, optional: bool):
266            """@private
267            """
268            self.load_generic(kwarg_name, self.trainer_kwargs["model"].parameters(), optional=optional)
269
270        def load_lr_scheduler(self, kwarg_name: str, optional: bool):
271            """@private
272            """
273            self.load_generic(kwarg_name, self.trainer_kwargs["optimizer"], optional=optional)
274
275        # todo: remove and rename kwarg 'logger' to 'logger_class'
276        def load_logger(self, kwarg_name: str, optional: bool):
277            """@private
278            """
279            assert kwarg_name == "logger"
280            self.load_generic("logger", optional=optional, only_class=True)
281
282    @staticmethod
283    def _get_save_dict(save_path, device):
284        if not os.path.exists(save_path):
285            raise ValueError(f"Cannot find checkpoint {save_path}")
286        return torch.load(save_path, map_location=device, weights_only=False)
287
288    @classmethod
289    def from_checkpoint(
290        cls,
291        checkpoint_folder: Union[os.PathLike, str],
292        name: Literal["best", "latest"] = "best",
293        device: Optional[Union[str, torch.device]] = None,
294    ):
295        """@private
296        """
297        save_path = os.path.join(checkpoint_folder, f"{name}.pt")
298        # make sure the correct device is set if we don't have access to CUDA
299        if not torch.cuda.is_available():
300            device = "cpu"
301        save_dict = cls._get_save_dict(save_path, device)
302        deserializer = cls.Deserializer(save_dict["init"], save_path, device)
303
304        has_kwargs = False
305        deserialized = []
306        for name, parameter in inspect.signature(cls).parameters.items():
307            if name == "kwargs":
308                has_kwargs = True
309                continue
310            deserializer.load(name, optional=parameter.default is not inspect.Parameter.empty)
311            deserialized.append(name)
312
313        # to deserialze kwargs we can't rely on inspecting the signature, so we
314        # go through the remaning kwarg names in init data instead
315        if has_kwargs:
316            kwarg_names = list(set(deserializer.init_data.keys()) - set(deserialized))
317            for name in kwarg_names:
318                if name.endswith("_kwargs"):
319                    continue
320                elif name.endswith("_dataset"):
321                    deserializer.load(name.replace("dataset", "loader"), optional=False)
322                elif name.endswith("_class"):
323                    deserializer.load(name.replace("_class", ""), optional=False)
324                else:
325                    deserializer.load(name, optional=False)
326
327        trainer = cls(**deserializer.trainer_kwargs)
328        trainer._initialize(0, save_dict)
329        trainer._is_initialized = True
330        return trainer
331
332    class Serializer:
333        """Implements how to serialize trainer kwargs from a trainer instance.
334
335        Examples:
336            To extend the serialization process you can inherite from this Serializer in a derived Trainer class.
337            Note that the methods `dump_generic_builtin()`, `dump_generic_class()` and `dump_generic_instance()`
338            called by the `dump()` method when appropriate cover most cases already.
339
340            This example adds `the_answer` kwarg, which requires extra steps on dumping only because we don't keep a
341            'the_answer' attribute:
342            >>> class MyTrainer(DefaultTrainer):
343            >>>     def __init__(self, *args, the_answer: int, **kwargs):
344            >>>         super().__init__(*args, **kwargs)
345            >>>         # self.the_answer = the_answer  # this would allow the default Serializer to save the new kwarg,
346            >>>         # but let's make things more interesting...
347            >>>         self.the = the_answer // 10
348            >>>         self.answer = the_answer % 10
349            >>>
350            >>>     class Serializer(DefaultTrainer.Serializer):
351            >>>         trainer: MyTrainer
352            >>>         def dump_the_answer(self, kwarg_name: str) -> None:  # custom dump method for 'the_answer' kwarg
353            >>>             assert kwarg_name == "the_answer"
354            >>>             # populate self.init_data with the serialized data required by Deserializer
355            >>>             # to restore the trainer kwargs
356            >>>             self.init_data["the_answer"] = self.trainer.the * 10 + self.trainer.answer
357
358            This example with both Serializer and Deserializer adds `the_answer` kwarg,
359            while saving it in two separate entries 'the' and 'answer'
360            >>> class MyTrainer(DefaultTrainer):
361            >>>     def __init__(self, *args, the_answer: int, **kwargs):
362            >>>         super().__init__(*args, **kwargs)
363            >>>         self.the_answer = the_answer
364            >>>
365            >>>     class Serializer(DefaultTrainer.Serializer):
366            >>>         trainer: MyTrainer
367            >>>         def dump_the_answer(self, kwarg_name: str):
368            >>>             assert kwarg_name == "the_answer"
369            >>>             self.init_data.update({
370            >>>                 "the": self.trainer.the_answer // 10,
371            >>>                 "answer": self.trainer.the_answer % 10
372            >>>             })
373            >>>
374            >>>     class Deserializer(DefaultTrainer.Deserializer):
375            >>>         def load_the_answer(self, kwarg_name: str, optional: bool):
376            >>>             assert kwarg_name == "the_answer"
377            >>>             # 'optional' is True if MyTrainer.__init__ specifies a default value for 'kwarg_name'
378            >>>             self.trainer_kwargs[kwarg_name] = self.init_data["the"] * 10 + self.init_data["answer"]
379
380        Args:
381            trainer: The trainer instance.
382        """
383
384        def __init__(self, trainer: DefaultTrainer):
385            self.trainer = trainer
386            self.init_data = {}  # to be populated during serialization process
387
388        def dump(self, kwarg_name: str) -> None:
389            """@private
390            """
391            dumper = getattr(self, f"dump_{kwarg_name}", None)
392            if dumper is not None:
393                dumper(kwarg_name)
394            elif kwarg_name.endswith("_loader"):
395                self.dump_data_loader(kwarg_name)
396            elif kwarg_name.endswith("_class"):
397                self.dump_generic_class(kwarg_name)
398            elif not hasattr(self.trainer, kwarg_name):
399                raise AttributeError(
400                    f"{self.trainer.__class__} missing attribute '{kwarg_name}' "
401                    f"or special dump method {self.trainer.__class__}.Serializer.dump_{kwarg_name}()"
402                )
403            else:
404                assert hasattr(self.trainer, kwarg_name)
405                obj = getattr(self.trainer, kwarg_name)
406                if obj is None or type(obj) in (
407                    bool,
408                    bytearray,
409                    bytes,
410                    dict,
411                    float,
412                    frozenset,
413                    int,
414                    list,
415                    set,
416                    str,
417                    tuple,
418                ):
419                    self.dump_generic_builtin(kwarg_name)
420                else:
421                    self.dump_generic_instance(kwarg_name)
422
423        def dump_generic_builtin(self, kwarg_name: str) -> None:
424            """@private
425            """
426            assert hasattr(self.trainer, kwarg_name)
427            self.init_data[kwarg_name] = getattr(self.trainer, kwarg_name)
428
429        def dump_generic_class(self, kwarg_name: str) -> None:
430            """@private
431            """
432            assert hasattr(self.trainer, kwarg_name)
433            assert kwarg_name.endswith("_class")
434            obj = getattr(self.trainer, kwarg_name)
435            self.init_data[kwarg_name] = None if obj is None else f"{obj.__module__}.{obj.__name__}"
436
437        def dump_generic_instance(self, kwarg_name: str) -> None:
438            """@private
439            """
440            assert hasattr(self.trainer, kwarg_name)
441            instance = getattr(self.trainer, kwarg_name)
442            self.init_data.update(
443                {
444                    f"{kwarg_name}_class": f"{instance.__class__.__module__}.{instance.__class__.__name__}",
445                    f"{kwarg_name}_kwargs": get_constructor_arguments(instance),
446                }
447            )
448
449        def dump_device(self, kwarg_name: str):
450            """@private
451            """
452            assert hasattr(self.trainer, kwarg_name)
453            self.init_data[kwarg_name] = str(getattr(self.trainer, kwarg_name))
454
455        def dump_data_loader(self, kwarg_name: str) -> None:
456            """@private
457            """
458            assert hasattr(self.trainer, kwarg_name)
459            loader = getattr(self.trainer, kwarg_name)
460            if loader is None:
461                return
462            self.init_data.update(
463                {
464                    f"{kwarg_name.replace('_loader', '_dataset')}": loader.dataset,
465                    f"{kwarg_name}_kwargs": get_constructor_arguments(loader),
466                }
467            )
468
469        def dump_logger(self, kwarg_name: str):  # todo: remove and rename kwarg 'logger' to 'logger_class'
470            """@private
471            """
472            self.dump_generic_class(f"{kwarg_name}_class")
473
474        def dump_model(self, kwarg_name: str):
475            """@private
476            """
477            if is_compiled(self.trainer.model):
478                self.init_data.update(
479                    {"model_class": self.trainer._model_class, "model_kwargs": self.trainer._model_kwargs}
480                )
481            else:
482                self.dump_generic_instance("model")
483
484    def _build_init(self) -> Dict[str, Any]:
485        serializer = self.Serializer(self)
486        for name in inspect.signature(self.__class__).parameters:
487            # special rules to serialize kwargs
488            # if a trainer class inherits from DefaultTrainer and has **kwargs
489            # they need to be saved in self._kwargs
490            if name == "kwargs":
491                if not hasattr(self, "_kwargs"):
492                    msg = "The trainer class has **kwargs in its signature, but is missing the _kwargs attribute. " +\
493                          "Please add self._kwargs to its __init__ function"
494                    raise RuntimeError(msg)
495                kwargs = getattr(self, "_kwargs")
496                for kwarg_name in kwargs:
497                    serializer.dump(kwarg_name)
498                continue
499            serializer.dump(name)
500
501        return serializer.init_data
502
503    def _initialize(self, iterations, load_from_checkpoint, epochs=None):
504        assert self.train_loader is not None
505        assert self.val_loader is not None
506        assert self.model is not None
507        assert self.loss is not None
508        assert self.optimizer is not None
509        assert self.metric is not None
510        assert self.device is not None
511
512        if load_from_checkpoint is not None:
513            self.load_checkpoint(load_from_checkpoint)
514
515        if sum((iterations is not None, epochs is not None)) != 1:
516            raise ValueError(
517                "Exactly one of 'iterations' or 'epochs' has to be specified to initialize the trainer."
518                f"You have passed 'iterations'={iterations} and 'epochs'={epochs}"
519            )
520
521        if epochs is None:
522            epochs = int(np.ceil(float(iterations) / len(self.train_loader)))
523        else:
524            iterations = epochs * len(self.train_loader)
525
526        self.max_iteration = self._iteration + iterations
527        self.max_epoch = self._epoch + epochs
528
529        if not getattr(self, "_is_initialized", False):
530            # check if we compile the model (only supported by pytorch 2)
531            # to enable (de)serialization of compiled models, we keep track of the model class and kwargs
532            if is_compiled(self.model):
533                warnings.warn(
534                    "You have passed a compiled model to the trainer."
535                    "It will not be possible to (de)serialize the trainer with it."
536                    "If you want to be able to do this please pass the normal model."
537                    "It can be automatically compiled by setting 'compile_model' to True"
538                )
539            self._model_class = f"{self.model.__class__.__module__}.{self.model.__class__.__name__}"
540            self._model_kwargs = get_constructor_arguments(self.model)
541            self.model = auto_compile(self.model, self.compile_model)
542
543            self.model.to(self.device)
544            self.loss.to(self.device)
545
546            # this saves all the information that is necessary
547            # to fully load the trainer from the checkpoint
548            self.init_data = self._build_init()
549
550            if self.logger_class is None:
551                self.logger = None
552            else:
553                # may set self.name if self.name is None
554                save_root = getattr(self, "save_root", None)
555                try:
556                    self.logger = self.logger_class(self, save_root, **(self.logger_kwargs or {}))
557                except (PermissionError, RuntimeError):
558                    warnings.warn(
559                        f"The checkpoint folder at {self.checkpoint_folder} could not be created."
560                        "The most likely reason for this is that you copied the checkpoint somewhere else,"
561                        "so we skip this error to enable loading the model from this checkpoint."
562                    )
563
564            try:
565                os.makedirs(self.checkpoint_folder, exist_ok=True)
566            except PermissionError:
567                warnings.warn(
568                    f"The checkpoint folder at {self.checkpoint_folder} could not be created."
569                    "The most likely reason for this is that you copied the checkpoint somewhere else,"
570                    "so we skip this error to enable loading the model from this checkpoint."
571                )
572                pass
573
574        best_metric = np.inf
575        return best_metric
576
577    def save_checkpoint(self, name, current_metric, best_metric, train_time=0.0, **extra_save_dict):
578        """@private
579        """
580        save_path = os.path.join(self.checkpoint_folder, f"{name}.pt")
581        extra_init_dict = extra_save_dict.pop("init", {})
582        save_dict = {
583            "iteration": self._iteration,
584            "epoch": self._epoch,
585            "best_epoch": self._best_epoch,
586            "best_metric": best_metric,
587            "current_metric": current_metric,
588            "model_state": self.model.state_dict(),
589            "optimizer_state": self.optimizer.state_dict(),
590            "init": self.init_data | extra_init_dict,
591            "train_time": train_time,
592            "timestamp": datetime.now().strftime("%d-%m-%Y (%H:%M:%S)"),
593        }
594        save_dict.update(**extra_save_dict)
595        if self.scaler is not None:
596            save_dict.update({"scaler_state": self.scaler.state_dict()})
597        if self.lr_scheduler is not None:
598            save_dict.update({"scheduler_state": self.lr_scheduler.state_dict()})
599
600        rank = getattr(self, "rank", None)
601        if rank is None or rank == 0:
602            torch.save(save_dict, save_path)
603
604    def load_checkpoint(self, checkpoint="best"):
605        """@private
606        """
607        if isinstance(checkpoint, str):
608            save_path = os.path.join(self.checkpoint_folder, f"{checkpoint}.pt")
609            if not os.path.exists(save_path):
610                warnings.warn(f"Cannot load checkpoint. {save_path} does not exist.")
611                return
612            save_dict = torch.load(save_path, weights_only=False)
613        elif isinstance(checkpoint, dict):
614            save_dict = checkpoint
615        else:
616            raise RuntimeError
617
618        self._iteration = save_dict["iteration"]
619        # Checkpoints are saved at the end of an epoch, before the epoch counter is incremented.
620        # So the stored value is the index of the last finished epoch, and the next epoch has the index + 1.
621        self._epoch = save_dict["epoch"] + 1
622        self._best_epoch = save_dict["best_epoch"]
623        self.best_metric = save_dict["best_metric"]
624        self.current_metric = save_dict["current_metric"]
625        self.train_time = save_dict.get("train_time", 0.0)
626
627        model_state = save_dict["model_state"]
628        # to enable loading compiled models
629        compiled_prefix = "_orig_mod."
630        model_state = OrderedDict(
631            [(k[len(compiled_prefix):] if k.startswith(compiled_prefix) else k, v) for k, v in model_state.items()]
632        )
633        self.model.load_state_dict(model_state)
634        # we need to send the network to the device before loading the optimizer state!
635        self.model.to(self.device)
636
637        self.optimizer.load_state_dict(save_dict["optimizer_state"])
638        scaler_state = save_dict.get("scaler_state")
639        if self.scaler is not None and scaler_state:
640            self.scaler.load_state_dict(scaler_state)
641        if self.lr_scheduler is not None:
642            self.lr_scheduler.load_state_dict(save_dict["scheduler_state"])
643
644        return save_dict
645
646    def _verify_if_training_completed(self, checkpoint="latest"):
647        save_path = os.path.join(self.checkpoint_folder, f"{checkpoint}.pt")
648        save_dict = torch.load(save_path, weights_only=False) if os.path.exists(save_path) else None
649        if save_dict and self.max_iteration == save_dict.get("iteration"):
650            return True
651        return False
652
653    def fit(
654        self,
655        iterations: Optional[int] = None,
656        load_from_checkpoint: Optional[Union[os.PathLike, str]] = None,
657        epochs: Optional[int] = None,
658        save_every_kth_epoch: Optional[int] = None,
659        progress=None,
660        overwrite_training: bool = True,
661    ):
662        """Run neural network training.
663
664        Exactly one of 'iterations' or 'epochs' has to be passed.
665
666        Args:
667            iterations: How long to train, specified in iterations.
668            load_from_checkpoint: Path to a checkpoint from where training should be continued .
669            epochs: How long to train, specified in epochs.
670            save_every_kth_epoch: Save checkpoints after every kth epoch in a separate file.
671                The corresponding checkpoints will be saved with the naming scheme 'epoch-{epoch}.pt'.
672            progress: Optional progress bar for integration with external tools. Expected to follow the tqdm interface.
673            overwrite_training: Whether to overwrite existing checkpoints in the save directory.
674        """
675        best_metric = self._initialize(iterations, load_from_checkpoint, epochs)
676
677        if not overwrite_training:
678            if load_from_checkpoint is not None:
679                raise ValueError(
680                    "We do not support 'overwrite_training=False' and 'load_from_checkpoint' at the same time."
681                )
682
683            if self._verify_if_training_completed():
684                print(
685                    f"The model is trained for {self.max_iteration} iterations / {self.max_epoch} epochs "
686                    "and 'overwrite_training' is set to 'False'."
687                )
688                print(f"The checkpoints are located at '{os.path.abspath(self.checkpoint_folder)}'.")
689                return
690
691        print(
692            "Start fitting for",
693            self.max_iteration - self._iteration,
694            "iterations / ",
695            self.max_epoch - self._epoch,
696            "epochs",
697        )
698        print("with", len(self.train_loader), "iterations per epoch")
699
700        if self.mixed_precision:
701            train_epoch = self._train_epoch_mixed
702            validate = self._validate_mixed
703            print("Training with mixed precision")
704        else:
705            train_epoch = self._train_epoch
706            validate = self._validate
707            print("Training with single precision")
708
709        total_iterations = epochs * len(self.train_loader) if iterations is None else iterations
710        if progress is None:
711            progress = tqdm(total=total_iterations, desc=f"Epoch {self._epoch}", leave=True)
712        else:
713            progress.total = total_iterations
714            progress.set_description(f"Epoch {self._epoch}")
715
716        msg = "Epoch %i: average [s/it]: %f, current metric: %f, best metric: %f"
717        train_epochs = self.max_epoch - self._epoch
718        t_start = time.time()
719        for epoch in range(train_epochs):
720
721            # Ensure data is shuffled differently at each epoch.
722            try:
723                self.train_loader.sampler.set_epoch(epoch)
724            except AttributeError:
725                pass
726
727            # Run training and validation for this epoch
728            t_per_iter = train_epoch(progress)
729            current_metric = validate()
730
731            # perform all the post-epoch steps:
732
733            # apply the learning rate scheduler
734            if self.lr_scheduler is not None:
735                self.lr_scheduler.step(current_metric)
736
737            # how long did we train in total?
738            total_train_time = (time.time() - t_start) + self.train_time
739
740            # save this checkpoint as the new best checkpoint if
741            # it has the best overall validation metric
742            if current_metric < best_metric:
743                best_metric = current_metric
744                self._best_epoch = self._epoch
745                self.save_checkpoint("best", current_metric, best_metric, train_time=total_train_time)
746
747            # save this checkpoint as the latest checkpoint
748            self.save_checkpoint("latest", current_metric, best_metric, train_time=total_train_time)
749
750            # if we save after every k-th epoch then check if we need to save now
751            if save_every_kth_epoch is not None and (self._epoch + 1) % save_every_kth_epoch == 0:
752                self.save_checkpoint(
753                    f"epoch-{self._epoch + 1}", current_metric, best_metric, train_time=total_train_time
754                )
755
756            # if early stopping has been specified then check if the stopping condition is met
757            if self.early_stopping is not None:
758                epochs_since_best = self._epoch - self._best_epoch
759                if epochs_since_best > self.early_stopping:
760                    print("Stopping training because there has been no improvement for", self.early_stopping, "epochs")
761                    break
762
763            self._epoch += 1
764            progress.set_description(msg % (self._epoch, t_per_iter, current_metric, best_metric), refresh=True)
765
766        print(f"Finished training after {self._epoch} epochs / {self._iteration} iterations.")
767        print(f"The best epoch is number {self._best_epoch}.")
768
769        if self._generate_name:
770            self.name = None
771
772        # Update the train time
773        self.train_time = total_train_time
774
775        # TODO save the model to wandb if we have the wandb logger
776        if isinstance(self.logger, WandbLogger):
777            self.logger.get_wandb().finish()
778
779    def _backprop(self, loss):
780        loss.backward()
781        self.optimizer.step()
782
783    def _backprop_mixed(self, loss):
784        self.scaler.scale(loss).backward()
785        self.scaler.step(self.optimizer)
786        self.scaler.update()
787
788    def _train_epoch(self, progress):
789        return self._train_epoch_impl(progress, contextlib.nullcontext, self._backprop)
790
791    def _train_epoch_mixed(self, progress):
792        return self._train_epoch_impl(
793            progress,
794            partial(torch.autocast, device_type=self._device_type, dtype=getattr(torch, self.mixed_precision_dtype)),
795            self._backprop_mixed
796        )
797
798    def _forward_and_loss(self, x, y):
799        pred = self.model(x)
800        if self._iteration % self.log_image_interval == 0:
801            if pred.requires_grad:
802                pred.retain_grad()
803
804        loss = self.loss(pred, y)
805        return pred, loss
806
807    def _train_epoch_impl(self, progress, forward_context, backprop: Callable[[torch.Tensor], None]):
808        self.model.train()
809
810        n_iter = 0
811        t_per_iter = time.time()
812        for x, y in self.train_loader:
813            x, y = x.to(self.device, non_blocking=True), y.to(self.device, non_blocking=True)
814
815            self.optimizer.zero_grad()
816
817            with forward_context():
818                pred, loss = self._forward_and_loss(x, y)
819
820            backprop(loss)
821
822            lr = [pm["lr"] for pm in self.optimizer.param_groups][0]
823            if self.logger is not None:
824                self.logger.log_train(self._iteration, loss, lr, x, y, pred, log_gradients=True)
825
826            self._iteration += 1
827            n_iter += 1
828            if self._iteration >= self.max_iteration:
829                break
830            progress.update(1)
831
832        t_per_iter = (time.time() - t_per_iter) / n_iter
833        return t_per_iter
834
835    def _validate(self):
836        return self._validate_impl(contextlib.nullcontext)
837
838    def _validate_mixed(self):
839        return self._validate_impl(
840            partial(torch.autocast, device_type=self._device_type, dtype=getattr(torch, self.mixed_precision_dtype))
841        )
842
843    def _validate_impl(self, forward_context):
844        self.model.eval()
845
846        metric_val = 0.0
847        loss_val = 0.0
848
849        with torch.no_grad():
850            for x, y in self.val_loader:
851                x, y = x.to(self.device, non_blocking=True), y.to(self.device, non_blocking=True)
852                with forward_context():
853                    pred, loss = self._forward_and_loss(x, y)
854                    metric = self.metric(pred, y)
855
856                loss_val += loss.item()
857                metric_val += metric.item()
858
859        metric_val /= len(self.val_loader)
860        loss_val /= len(self.val_loader)
861        if self.logger is not None:
862            self.logger.log_validation(self._iteration, metric_val, loss_val, x, y, pred)
863        return metric_val
class DefaultTrainer:
 25class DefaultTrainer:
 26    """Trainer class for training a segmentation network.
 27
 28    The trainer class implements the core logic for training a network with pytorch.
 29    It implements a training loop to run training and validation, which is started with `fit`.
 30    The checkpoints and logs from the training run will be saved in the current working directory,
 31    or in the directory specifified by `save_root`. Training can be continued from a checkpoint
 32    by passing it's location to the `load_from_checkpoint` argument of `fit`.
 33
 34    A pre-configured instance of the trainer can be obtained from `torch_em.default_segmentation_trainer`.
 35    Alternatively, the trainer class can also be instantiated as in this example:
 36    ```python
 37    import torch
 38    from torch_em.loss import DiceLoss
 39    from torch_em.model import UNet2d
 40    from torch_em.data.datasets.light_microscopy import get_dsb_loader
 41    from torch_em.trainer import DefaultTrainer
 42
 43    # The training data will be downloaded to this location.
 44    data_root = "/path/to/save/the/training/data"
 45    patch_shape = (256, 256)
 46
 47    # Create the model and optimizer.
 48    model = UNet2d(in_channels=1, out_channels=1)
 49    optimizer = torch.optim.AdamW(model.parameters())
 50
 51    trainer = DefaultTrainer(
 52        name="unet-training",
 53        train_loader=get_dsb_loader(path=data_root, patch_shape=patch_shape, split="train"),
 54        val_loader=get_dsb_loader(path=data_root, patch_shape=patch_shape, split="test"),
 55        model=model,
 56        loss=DiceLoss(),  # The loss function.
 57        optimizer=optimizer,
 58        metric=DiceLoss(),  # The metric. The trainer expects smaller values to represent better results.
 59        device="cuda",  # The device to use for training.
 60    )
 61    trainer.fit(iterations=int(2.5e4))  # Train for 25.000 iterations.
 62    ```
 63
 64    Args:
 65        name: The name of the checkpoint that will be created by the trainer.
 66        train_loader: The data loader containing the training data.
 67        val_loader: The data loader containing the validation data.
 68        model: The model to train.
 69        loss: The loss function for training.
 70        optimizer: The optimizer.
 71        metric: The metric for validation.
 72        device: The torch device to use for training. If None, will use a GPU if available.
 73        lr_scheduler: The learning rate scheduler.
 74        log_image_interval: The interval for saving images during logging, in training iterations.
 75        mixed_precision: Whether to train with mixed precision.
 76        early_stopping: The patience for early stopping in epochs. If None, early stopping will not be used.
 77        logger: The logger class. Will be instantiated for logging.
 78            By default uses `torch_em.training.tensorboard_logger.TensorboardLogger`.
 79        logger_kwargs: The keyword arguments for the logger class.
 80        id_: Unique identifier for the trainer. If None then `name` will be used.
 81        save_root: The root folder for saving the checkpoint and logs.
 82        compile_model: Whether to compile the model before training.
 83        rank: Rank argument for distributed training. See `torch_em.multi_gpu_training` for details.
 84        mixed_precision_dtype: The dtype for autocast in mixed precision training, 'float16' or 'bfloat16'.
 85            The default is 'float16' on the GPU and 'bfloat16' on the CPU. Use 'bfloat16' to avoid overflows.
 86    """
 87    def __init__(
 88        self,
 89        name: Optional[str],
 90        train_loader: torch.utils.data.DataLoader,
 91        val_loader: torch.utils.data.DataLoader,
 92        model: torch.nn.Module,
 93        loss: torch.nn.Module,
 94        optimizer: torch.optim.Optimizer,
 95        metric: Callable,
 96        device: Union[str, torch.device],
 97        lr_scheduler: Optional[torch.optim.lr_scheduler._LRScheduler] = None,
 98        log_image_interval: int = 100,
 99        mixed_precision: bool = True,
100        early_stopping: Optional[int] = None,
101        logger=TensorboardLogger,
102        logger_kwargs: Optional[Dict[str, Any]] = None,
103        id_: Optional[str] = None,
104        save_root: Optional[str] = None,
105        compile_model: Optional[Union[bool, str]] = None,
106        rank: Optional[int] = None,
107        mixed_precision_dtype: Optional[str] = None,
108    ):
109        if name is None and not issubclass(logger, WandbLogger):
110            raise TypeError("Name cannot be None if not using the WandbLogger")
111
112        self._generate_name = name is None
113        self.name = name
114        self.id_ = id_ or name
115        self.train_loader = train_loader
116        self.val_loader = val_loader
117        self.model = model
118        self.loss = loss
119        self.optimizer = optimizer
120        self.metric = metric
121        self.device = torch.device(device)
122        self.lr_scheduler = lr_scheduler
123        self.log_image_interval = log_image_interval
124        self.save_root = save_root
125        self.compile_model = compile_model
126        self.rank = rank
127        self._device_type = "cpu" if self.device.type == "cpu" else "cuda"
128
129        self._iteration = 0
130        self._epoch = 0
131        self._best_epoch = 0
132
133        self.mixed_precision = mixed_precision
134        # These are the defaults of torch.autocast for each device type.
135        self.mixed_precision_dtype = mixed_precision_dtype or ("bfloat16" if self._device_type == "cpu" else "float16")
136        self.early_stopping = early_stopping
137        self.train_time = 0.0
138
139        if mixed_precision:
140            # Only float16 needs gradient scaling. bfloat16 has the same range as float32.
141            self.scaler = torch.GradScaler(self._device_type, enabled=self.mixed_precision_dtype == "float16")
142        else:
143            self.scaler = None
144
145        self.logger_class = logger
146        self.logger_kwargs = logger_kwargs
147        self.log_image_interval = log_image_interval
148
149    @property
150    def checkpoint_folder(self):
151        assert self.id_ is not None  # Because the logger may generate and set trainer.id on logger.__init__.
152        # Save_root enables saving the checkpoints somewhere else than in the local folder.
153        # This is handy for filesystems with limited space, where saving the checkpoints
154        # and log files can lead to running out of space.
155        save_root = getattr(self, "save_root", None)
156        return os.path.join("./checkpoints", self.id_) if save_root is None else\
157            os.path.join(save_root, "./checkpoints", self.id_)
158
159    @property
160    def iteration(self):
161        return self._iteration
162
163    @property
164    def epoch(self):
165        return self._epoch
166
167    class Deserializer:
168        """Determines how to deserialize the trainer kwargs from serialized 'init_data'.
169
170        Examples:
171            To extend the initialization process you can inherite from this Deserializer in an inherited Trainer class.
172            Note that `DefaultTrainer.Deserializer.load_generic()` covers most cases already.
173
174            This example adds `the_answer` kwarg, which requires 'calculations' upon initialization:
175            >>> class MyTrainer(DefaultTrainer):
176            >>>     def __init__(self, *args, the_answer: int, **kwargs):
177            >>>         super().__init__(*args, **kwargs)
178            >>>         self.the_answer = the_answer  # this allows the default Serializer to save the new kwarg,
179            >>>                                       # see DefaultTrainer.Serializer
180            >>>
181            >>>     class Deserializer(DefaultTrainer.Deserializer):
182            >>>         def load_the_answer(self):
183            >>>             generic_answer = self.init_data["the_answer"]
184            >>>             # (device dependent) special deserialization
185            >>>             if self.trainer_kwargs["device"].type == "cpu":  # accessing previously deserialized kwarg
186            >>>                 self.trainer_kwargs["the_answer"] = generic_answer + 1
187            >>>             else:
188            >>>                 self.trainer_kwargs["the_answer"] = generic_answer * 2
189
190        Args:
191            init_data: The initialization data of the trainer.
192            save_path: The path where the checkpoint was saved.
193            device: The device.
194        """
195
196        def __init__(self, init_data: Dict, save_path: str, device: Union[str, torch.device]):
197            self.init_data = init_data
198            self.save_path = save_path
199            # Populate with deserialized trainer kwargs during deserialization; possibly overwrite 'device'.
200            self.trainer_kwargs: Dict[str, Any] = dict(
201                device=torch.device(self.init_data["device"]) if device is None else torch.device(device)
202            )
203
204        def load(self, kwarg_name: str, optional):
205            """@private
206            """
207            # `optional` is True if self.trainer.__class__.__init__ specifies a default value for 'kwarg_name'
208            if kwarg_name == "device":
209                pass  # deserialized in __init__
210            elif kwarg_name.endswith("_loader"):
211                self.load_data_loader(kwarg_name, optional)
212            else:
213                load = getattr(self, f"load_{kwarg_name}", self.load_generic)
214                load(kwarg_name, optional=optional)
215
216        def load_data_loader(self, loader_name, optional) -> None:
217            """@private
218            """
219            ds = self.init_data.get(loader_name.replace("_loader", "_dataset"))
220            if ds is None and optional:
221                return
222
223            loader_kwargs = self.init_data[f"{loader_name}_kwargs"]
224            loader = torch.utils.data.DataLoader(ds, **loader_kwargs)
225            # monkey patch shuffle loader_name to the loader
226            loader.shuffle = loader_kwargs.get("shuffle", False)
227            self.trainer_kwargs[loader_name] = loader
228
229        def load_generic(
230            self,
231            kwarg_name: str,
232            *dynamic_args: Dict,
233            optional: bool,
234            only_class: bool = False,
235            dynamic_kwargs: Optional[Dict[str, Any]] = None,
236        ) -> None:
237            """@private
238            """
239            if kwarg_name in self.init_data:
240                self.trainer_kwargs[kwarg_name] = self.init_data[kwarg_name]
241                return
242
243            this_cls = self.init_data.get(f"{kwarg_name}_class", None)
244            if this_cls is None:
245                if optional:
246                    return
247                else:
248                    raise RuntimeError(f"Could not find init data for {kwarg_name} in {self.save_path}")
249
250            assert isinstance(this_cls, str), this_cls
251            assert "." in this_cls, this_cls
252            cls_p, cls_m = this_cls.rsplit(".", 1)
253            this_cls = getattr(import_module(cls_p), cls_m)
254            if only_class:
255                self.trainer_kwargs[kwarg_name] = this_cls
256            else:
257                self.trainer_kwargs[kwarg_name] = this_cls(
258                    *dynamic_args, **self.init_data.get(f"{kwarg_name}_kwargs", {}), **(dynamic_kwargs or {})
259                )
260
261        def load_name(self, kwarg_name: str, optional: bool):
262            """@private
263            """
264            self.trainer_kwargs[kwarg_name] = os.path.split(os.path.dirname(self.save_path))[1]
265
266        def load_optimizer(self, kwarg_name: str, optional: bool):
267            """@private
268            """
269            self.load_generic(kwarg_name, self.trainer_kwargs["model"].parameters(), optional=optional)
270
271        def load_lr_scheduler(self, kwarg_name: str, optional: bool):
272            """@private
273            """
274            self.load_generic(kwarg_name, self.trainer_kwargs["optimizer"], optional=optional)
275
276        # todo: remove and rename kwarg 'logger' to 'logger_class'
277        def load_logger(self, kwarg_name: str, optional: bool):
278            """@private
279            """
280            assert kwarg_name == "logger"
281            self.load_generic("logger", optional=optional, only_class=True)
282
283    @staticmethod
284    def _get_save_dict(save_path, device):
285        if not os.path.exists(save_path):
286            raise ValueError(f"Cannot find checkpoint {save_path}")
287        return torch.load(save_path, map_location=device, weights_only=False)
288
289    @classmethod
290    def from_checkpoint(
291        cls,
292        checkpoint_folder: Union[os.PathLike, str],
293        name: Literal["best", "latest"] = "best",
294        device: Optional[Union[str, torch.device]] = None,
295    ):
296        """@private
297        """
298        save_path = os.path.join(checkpoint_folder, f"{name}.pt")
299        # make sure the correct device is set if we don't have access to CUDA
300        if not torch.cuda.is_available():
301            device = "cpu"
302        save_dict = cls._get_save_dict(save_path, device)
303        deserializer = cls.Deserializer(save_dict["init"], save_path, device)
304
305        has_kwargs = False
306        deserialized = []
307        for name, parameter in inspect.signature(cls).parameters.items():
308            if name == "kwargs":
309                has_kwargs = True
310                continue
311            deserializer.load(name, optional=parameter.default is not inspect.Parameter.empty)
312            deserialized.append(name)
313
314        # to deserialze kwargs we can't rely on inspecting the signature, so we
315        # go through the remaning kwarg names in init data instead
316        if has_kwargs:
317            kwarg_names = list(set(deserializer.init_data.keys()) - set(deserialized))
318            for name in kwarg_names:
319                if name.endswith("_kwargs"):
320                    continue
321                elif name.endswith("_dataset"):
322                    deserializer.load(name.replace("dataset", "loader"), optional=False)
323                elif name.endswith("_class"):
324                    deserializer.load(name.replace("_class", ""), optional=False)
325                else:
326                    deserializer.load(name, optional=False)
327
328        trainer = cls(**deserializer.trainer_kwargs)
329        trainer._initialize(0, save_dict)
330        trainer._is_initialized = True
331        return trainer
332
333    class Serializer:
334        """Implements how to serialize trainer kwargs from a trainer instance.
335
336        Examples:
337            To extend the serialization process you can inherite from this Serializer in a derived Trainer class.
338            Note that the methods `dump_generic_builtin()`, `dump_generic_class()` and `dump_generic_instance()`
339            called by the `dump()` method when appropriate cover most cases already.
340
341            This example adds `the_answer` kwarg, which requires extra steps on dumping only because we don't keep a
342            'the_answer' attribute:
343            >>> class MyTrainer(DefaultTrainer):
344            >>>     def __init__(self, *args, the_answer: int, **kwargs):
345            >>>         super().__init__(*args, **kwargs)
346            >>>         # self.the_answer = the_answer  # this would allow the default Serializer to save the new kwarg,
347            >>>         # but let's make things more interesting...
348            >>>         self.the = the_answer // 10
349            >>>         self.answer = the_answer % 10
350            >>>
351            >>>     class Serializer(DefaultTrainer.Serializer):
352            >>>         trainer: MyTrainer
353            >>>         def dump_the_answer(self, kwarg_name: str) -> None:  # custom dump method for 'the_answer' kwarg
354            >>>             assert kwarg_name == "the_answer"
355            >>>             # populate self.init_data with the serialized data required by Deserializer
356            >>>             # to restore the trainer kwargs
357            >>>             self.init_data["the_answer"] = self.trainer.the * 10 + self.trainer.answer
358
359            This example with both Serializer and Deserializer adds `the_answer` kwarg,
360            while saving it in two separate entries 'the' and 'answer'
361            >>> class MyTrainer(DefaultTrainer):
362            >>>     def __init__(self, *args, the_answer: int, **kwargs):
363            >>>         super().__init__(*args, **kwargs)
364            >>>         self.the_answer = the_answer
365            >>>
366            >>>     class Serializer(DefaultTrainer.Serializer):
367            >>>         trainer: MyTrainer
368            >>>         def dump_the_answer(self, kwarg_name: str):
369            >>>             assert kwarg_name == "the_answer"
370            >>>             self.init_data.update({
371            >>>                 "the": self.trainer.the_answer // 10,
372            >>>                 "answer": self.trainer.the_answer % 10
373            >>>             })
374            >>>
375            >>>     class Deserializer(DefaultTrainer.Deserializer):
376            >>>         def load_the_answer(self, kwarg_name: str, optional: bool):
377            >>>             assert kwarg_name == "the_answer"
378            >>>             # 'optional' is True if MyTrainer.__init__ specifies a default value for 'kwarg_name'
379            >>>             self.trainer_kwargs[kwarg_name] = self.init_data["the"] * 10 + self.init_data["answer"]
380
381        Args:
382            trainer: The trainer instance.
383        """
384
385        def __init__(self, trainer: DefaultTrainer):
386            self.trainer = trainer
387            self.init_data = {}  # to be populated during serialization process
388
389        def dump(self, kwarg_name: str) -> None:
390            """@private
391            """
392            dumper = getattr(self, f"dump_{kwarg_name}", None)
393            if dumper is not None:
394                dumper(kwarg_name)
395            elif kwarg_name.endswith("_loader"):
396                self.dump_data_loader(kwarg_name)
397            elif kwarg_name.endswith("_class"):
398                self.dump_generic_class(kwarg_name)
399            elif not hasattr(self.trainer, kwarg_name):
400                raise AttributeError(
401                    f"{self.trainer.__class__} missing attribute '{kwarg_name}' "
402                    f"or special dump method {self.trainer.__class__}.Serializer.dump_{kwarg_name}()"
403                )
404            else:
405                assert hasattr(self.trainer, kwarg_name)
406                obj = getattr(self.trainer, kwarg_name)
407                if obj is None or type(obj) in (
408                    bool,
409                    bytearray,
410                    bytes,
411                    dict,
412                    float,
413                    frozenset,
414                    int,
415                    list,
416                    set,
417                    str,
418                    tuple,
419                ):
420                    self.dump_generic_builtin(kwarg_name)
421                else:
422                    self.dump_generic_instance(kwarg_name)
423
424        def dump_generic_builtin(self, kwarg_name: str) -> None:
425            """@private
426            """
427            assert hasattr(self.trainer, kwarg_name)
428            self.init_data[kwarg_name] = getattr(self.trainer, kwarg_name)
429
430        def dump_generic_class(self, kwarg_name: str) -> None:
431            """@private
432            """
433            assert hasattr(self.trainer, kwarg_name)
434            assert kwarg_name.endswith("_class")
435            obj = getattr(self.trainer, kwarg_name)
436            self.init_data[kwarg_name] = None if obj is None else f"{obj.__module__}.{obj.__name__}"
437
438        def dump_generic_instance(self, kwarg_name: str) -> None:
439            """@private
440            """
441            assert hasattr(self.trainer, kwarg_name)
442            instance = getattr(self.trainer, kwarg_name)
443            self.init_data.update(
444                {
445                    f"{kwarg_name}_class": f"{instance.__class__.__module__}.{instance.__class__.__name__}",
446                    f"{kwarg_name}_kwargs": get_constructor_arguments(instance),
447                }
448            )
449
450        def dump_device(self, kwarg_name: str):
451            """@private
452            """
453            assert hasattr(self.trainer, kwarg_name)
454            self.init_data[kwarg_name] = str(getattr(self.trainer, kwarg_name))
455
456        def dump_data_loader(self, kwarg_name: str) -> None:
457            """@private
458            """
459            assert hasattr(self.trainer, kwarg_name)
460            loader = getattr(self.trainer, kwarg_name)
461            if loader is None:
462                return
463            self.init_data.update(
464                {
465                    f"{kwarg_name.replace('_loader', '_dataset')}": loader.dataset,
466                    f"{kwarg_name}_kwargs": get_constructor_arguments(loader),
467                }
468            )
469
470        def dump_logger(self, kwarg_name: str):  # todo: remove and rename kwarg 'logger' to 'logger_class'
471            """@private
472            """
473            self.dump_generic_class(f"{kwarg_name}_class")
474
475        def dump_model(self, kwarg_name: str):
476            """@private
477            """
478            if is_compiled(self.trainer.model):
479                self.init_data.update(
480                    {"model_class": self.trainer._model_class, "model_kwargs": self.trainer._model_kwargs}
481                )
482            else:
483                self.dump_generic_instance("model")
484
485    def _build_init(self) -> Dict[str, Any]:
486        serializer = self.Serializer(self)
487        for name in inspect.signature(self.__class__).parameters:
488            # special rules to serialize kwargs
489            # if a trainer class inherits from DefaultTrainer and has **kwargs
490            # they need to be saved in self._kwargs
491            if name == "kwargs":
492                if not hasattr(self, "_kwargs"):
493                    msg = "The trainer class has **kwargs in its signature, but is missing the _kwargs attribute. " +\
494                          "Please add self._kwargs to its __init__ function"
495                    raise RuntimeError(msg)
496                kwargs = getattr(self, "_kwargs")
497                for kwarg_name in kwargs:
498                    serializer.dump(kwarg_name)
499                continue
500            serializer.dump(name)
501
502        return serializer.init_data
503
504    def _initialize(self, iterations, load_from_checkpoint, epochs=None):
505        assert self.train_loader is not None
506        assert self.val_loader is not None
507        assert self.model is not None
508        assert self.loss is not None
509        assert self.optimizer is not None
510        assert self.metric is not None
511        assert self.device is not None
512
513        if load_from_checkpoint is not None:
514            self.load_checkpoint(load_from_checkpoint)
515
516        if sum((iterations is not None, epochs is not None)) != 1:
517            raise ValueError(
518                "Exactly one of 'iterations' or 'epochs' has to be specified to initialize the trainer."
519                f"You have passed 'iterations'={iterations} and 'epochs'={epochs}"
520            )
521
522        if epochs is None:
523            epochs = int(np.ceil(float(iterations) / len(self.train_loader)))
524        else:
525            iterations = epochs * len(self.train_loader)
526
527        self.max_iteration = self._iteration + iterations
528        self.max_epoch = self._epoch + epochs
529
530        if not getattr(self, "_is_initialized", False):
531            # check if we compile the model (only supported by pytorch 2)
532            # to enable (de)serialization of compiled models, we keep track of the model class and kwargs
533            if is_compiled(self.model):
534                warnings.warn(
535                    "You have passed a compiled model to the trainer."
536                    "It will not be possible to (de)serialize the trainer with it."
537                    "If you want to be able to do this please pass the normal model."
538                    "It can be automatically compiled by setting 'compile_model' to True"
539                )
540            self._model_class = f"{self.model.__class__.__module__}.{self.model.__class__.__name__}"
541            self._model_kwargs = get_constructor_arguments(self.model)
542            self.model = auto_compile(self.model, self.compile_model)
543
544            self.model.to(self.device)
545            self.loss.to(self.device)
546
547            # this saves all the information that is necessary
548            # to fully load the trainer from the checkpoint
549            self.init_data = self._build_init()
550
551            if self.logger_class is None:
552                self.logger = None
553            else:
554                # may set self.name if self.name is None
555                save_root = getattr(self, "save_root", None)
556                try:
557                    self.logger = self.logger_class(self, save_root, **(self.logger_kwargs or {}))
558                except (PermissionError, RuntimeError):
559                    warnings.warn(
560                        f"The checkpoint folder at {self.checkpoint_folder} could not be created."
561                        "The most likely reason for this is that you copied the checkpoint somewhere else,"
562                        "so we skip this error to enable loading the model from this checkpoint."
563                    )
564
565            try:
566                os.makedirs(self.checkpoint_folder, exist_ok=True)
567            except PermissionError:
568                warnings.warn(
569                    f"The checkpoint folder at {self.checkpoint_folder} could not be created."
570                    "The most likely reason for this is that you copied the checkpoint somewhere else,"
571                    "so we skip this error to enable loading the model from this checkpoint."
572                )
573                pass
574
575        best_metric = np.inf
576        return best_metric
577
578    def save_checkpoint(self, name, current_metric, best_metric, train_time=0.0, **extra_save_dict):
579        """@private
580        """
581        save_path = os.path.join(self.checkpoint_folder, f"{name}.pt")
582        extra_init_dict = extra_save_dict.pop("init", {})
583        save_dict = {
584            "iteration": self._iteration,
585            "epoch": self._epoch,
586            "best_epoch": self._best_epoch,
587            "best_metric": best_metric,
588            "current_metric": current_metric,
589            "model_state": self.model.state_dict(),
590            "optimizer_state": self.optimizer.state_dict(),
591            "init": self.init_data | extra_init_dict,
592            "train_time": train_time,
593            "timestamp": datetime.now().strftime("%d-%m-%Y (%H:%M:%S)"),
594        }
595        save_dict.update(**extra_save_dict)
596        if self.scaler is not None:
597            save_dict.update({"scaler_state": self.scaler.state_dict()})
598        if self.lr_scheduler is not None:
599            save_dict.update({"scheduler_state": self.lr_scheduler.state_dict()})
600
601        rank = getattr(self, "rank", None)
602        if rank is None or rank == 0:
603            torch.save(save_dict, save_path)
604
605    def load_checkpoint(self, checkpoint="best"):
606        """@private
607        """
608        if isinstance(checkpoint, str):
609            save_path = os.path.join(self.checkpoint_folder, f"{checkpoint}.pt")
610            if not os.path.exists(save_path):
611                warnings.warn(f"Cannot load checkpoint. {save_path} does not exist.")
612                return
613            save_dict = torch.load(save_path, weights_only=False)
614        elif isinstance(checkpoint, dict):
615            save_dict = checkpoint
616        else:
617            raise RuntimeError
618
619        self._iteration = save_dict["iteration"]
620        # Checkpoints are saved at the end of an epoch, before the epoch counter is incremented.
621        # So the stored value is the index of the last finished epoch, and the next epoch has the index + 1.
622        self._epoch = save_dict["epoch"] + 1
623        self._best_epoch = save_dict["best_epoch"]
624        self.best_metric = save_dict["best_metric"]
625        self.current_metric = save_dict["current_metric"]
626        self.train_time = save_dict.get("train_time", 0.0)
627
628        model_state = save_dict["model_state"]
629        # to enable loading compiled models
630        compiled_prefix = "_orig_mod."
631        model_state = OrderedDict(
632            [(k[len(compiled_prefix):] if k.startswith(compiled_prefix) else k, v) for k, v in model_state.items()]
633        )
634        self.model.load_state_dict(model_state)
635        # we need to send the network to the device before loading the optimizer state!
636        self.model.to(self.device)
637
638        self.optimizer.load_state_dict(save_dict["optimizer_state"])
639        scaler_state = save_dict.get("scaler_state")
640        if self.scaler is not None and scaler_state:
641            self.scaler.load_state_dict(scaler_state)
642        if self.lr_scheduler is not None:
643            self.lr_scheduler.load_state_dict(save_dict["scheduler_state"])
644
645        return save_dict
646
647    def _verify_if_training_completed(self, checkpoint="latest"):
648        save_path = os.path.join(self.checkpoint_folder, f"{checkpoint}.pt")
649        save_dict = torch.load(save_path, weights_only=False) if os.path.exists(save_path) else None
650        if save_dict and self.max_iteration == save_dict.get("iteration"):
651            return True
652        return False
653
654    def fit(
655        self,
656        iterations: Optional[int] = None,
657        load_from_checkpoint: Optional[Union[os.PathLike, str]] = None,
658        epochs: Optional[int] = None,
659        save_every_kth_epoch: Optional[int] = None,
660        progress=None,
661        overwrite_training: bool = True,
662    ):
663        """Run neural network training.
664
665        Exactly one of 'iterations' or 'epochs' has to be passed.
666
667        Args:
668            iterations: How long to train, specified in iterations.
669            load_from_checkpoint: Path to a checkpoint from where training should be continued .
670            epochs: How long to train, specified in epochs.
671            save_every_kth_epoch: Save checkpoints after every kth epoch in a separate file.
672                The corresponding checkpoints will be saved with the naming scheme 'epoch-{epoch}.pt'.
673            progress: Optional progress bar for integration with external tools. Expected to follow the tqdm interface.
674            overwrite_training: Whether to overwrite existing checkpoints in the save directory.
675        """
676        best_metric = self._initialize(iterations, load_from_checkpoint, epochs)
677
678        if not overwrite_training:
679            if load_from_checkpoint is not None:
680                raise ValueError(
681                    "We do not support 'overwrite_training=False' and 'load_from_checkpoint' at the same time."
682                )
683
684            if self._verify_if_training_completed():
685                print(
686                    f"The model is trained for {self.max_iteration} iterations / {self.max_epoch} epochs "
687                    "and 'overwrite_training' is set to 'False'."
688                )
689                print(f"The checkpoints are located at '{os.path.abspath(self.checkpoint_folder)}'.")
690                return
691
692        print(
693            "Start fitting for",
694            self.max_iteration - self._iteration,
695            "iterations / ",
696            self.max_epoch - self._epoch,
697            "epochs",
698        )
699        print("with", len(self.train_loader), "iterations per epoch")
700
701        if self.mixed_precision:
702            train_epoch = self._train_epoch_mixed
703            validate = self._validate_mixed
704            print("Training with mixed precision")
705        else:
706            train_epoch = self._train_epoch
707            validate = self._validate
708            print("Training with single precision")
709
710        total_iterations = epochs * len(self.train_loader) if iterations is None else iterations
711        if progress is None:
712            progress = tqdm(total=total_iterations, desc=f"Epoch {self._epoch}", leave=True)
713        else:
714            progress.total = total_iterations
715            progress.set_description(f"Epoch {self._epoch}")
716
717        msg = "Epoch %i: average [s/it]: %f, current metric: %f, best metric: %f"
718        train_epochs = self.max_epoch - self._epoch
719        t_start = time.time()
720        for epoch in range(train_epochs):
721
722            # Ensure data is shuffled differently at each epoch.
723            try:
724                self.train_loader.sampler.set_epoch(epoch)
725            except AttributeError:
726                pass
727
728            # Run training and validation for this epoch
729            t_per_iter = train_epoch(progress)
730            current_metric = validate()
731
732            # perform all the post-epoch steps:
733
734            # apply the learning rate scheduler
735            if self.lr_scheduler is not None:
736                self.lr_scheduler.step(current_metric)
737
738            # how long did we train in total?
739            total_train_time = (time.time() - t_start) + self.train_time
740
741            # save this checkpoint as the new best checkpoint if
742            # it has the best overall validation metric
743            if current_metric < best_metric:
744                best_metric = current_metric
745                self._best_epoch = self._epoch
746                self.save_checkpoint("best", current_metric, best_metric, train_time=total_train_time)
747
748            # save this checkpoint as the latest checkpoint
749            self.save_checkpoint("latest", current_metric, best_metric, train_time=total_train_time)
750
751            # if we save after every k-th epoch then check if we need to save now
752            if save_every_kth_epoch is not None and (self._epoch + 1) % save_every_kth_epoch == 0:
753                self.save_checkpoint(
754                    f"epoch-{self._epoch + 1}", current_metric, best_metric, train_time=total_train_time
755                )
756
757            # if early stopping has been specified then check if the stopping condition is met
758            if self.early_stopping is not None:
759                epochs_since_best = self._epoch - self._best_epoch
760                if epochs_since_best > self.early_stopping:
761                    print("Stopping training because there has been no improvement for", self.early_stopping, "epochs")
762                    break
763
764            self._epoch += 1
765            progress.set_description(msg % (self._epoch, t_per_iter, current_metric, best_metric), refresh=True)
766
767        print(f"Finished training after {self._epoch} epochs / {self._iteration} iterations.")
768        print(f"The best epoch is number {self._best_epoch}.")
769
770        if self._generate_name:
771            self.name = None
772
773        # Update the train time
774        self.train_time = total_train_time
775
776        # TODO save the model to wandb if we have the wandb logger
777        if isinstance(self.logger, WandbLogger):
778            self.logger.get_wandb().finish()
779
780    def _backprop(self, loss):
781        loss.backward()
782        self.optimizer.step()
783
784    def _backprop_mixed(self, loss):
785        self.scaler.scale(loss).backward()
786        self.scaler.step(self.optimizer)
787        self.scaler.update()
788
789    def _train_epoch(self, progress):
790        return self._train_epoch_impl(progress, contextlib.nullcontext, self._backprop)
791
792    def _train_epoch_mixed(self, progress):
793        return self._train_epoch_impl(
794            progress,
795            partial(torch.autocast, device_type=self._device_type, dtype=getattr(torch, self.mixed_precision_dtype)),
796            self._backprop_mixed
797        )
798
799    def _forward_and_loss(self, x, y):
800        pred = self.model(x)
801        if self._iteration % self.log_image_interval == 0:
802            if pred.requires_grad:
803                pred.retain_grad()
804
805        loss = self.loss(pred, y)
806        return pred, loss
807
808    def _train_epoch_impl(self, progress, forward_context, backprop: Callable[[torch.Tensor], None]):
809        self.model.train()
810
811        n_iter = 0
812        t_per_iter = time.time()
813        for x, y in self.train_loader:
814            x, y = x.to(self.device, non_blocking=True), y.to(self.device, non_blocking=True)
815
816            self.optimizer.zero_grad()
817
818            with forward_context():
819                pred, loss = self._forward_and_loss(x, y)
820
821            backprop(loss)
822
823            lr = [pm["lr"] for pm in self.optimizer.param_groups][0]
824            if self.logger is not None:
825                self.logger.log_train(self._iteration, loss, lr, x, y, pred, log_gradients=True)
826
827            self._iteration += 1
828            n_iter += 1
829            if self._iteration >= self.max_iteration:
830                break
831            progress.update(1)
832
833        t_per_iter = (time.time() - t_per_iter) / n_iter
834        return t_per_iter
835
836    def _validate(self):
837        return self._validate_impl(contextlib.nullcontext)
838
839    def _validate_mixed(self):
840        return self._validate_impl(
841            partial(torch.autocast, device_type=self._device_type, dtype=getattr(torch, self.mixed_precision_dtype))
842        )
843
844    def _validate_impl(self, forward_context):
845        self.model.eval()
846
847        metric_val = 0.0
848        loss_val = 0.0
849
850        with torch.no_grad():
851            for x, y in self.val_loader:
852                x, y = x.to(self.device, non_blocking=True), y.to(self.device, non_blocking=True)
853                with forward_context():
854                    pred, loss = self._forward_and_loss(x, y)
855                    metric = self.metric(pred, y)
856
857                loss_val += loss.item()
858                metric_val += metric.item()
859
860        metric_val /= len(self.val_loader)
861        loss_val /= len(self.val_loader)
862        if self.logger is not None:
863            self.logger.log_validation(self._iteration, metric_val, loss_val, x, y, pred)
864        return metric_val

Trainer class for training a segmentation network.

The trainer class implements the core logic for training a network with pytorch. It implements a training loop to run training and validation, which is started with fit. The checkpoints and logs from the training run will be saved in the current working directory, or in the directory specifified by save_root. Training can be continued from a checkpoint by passing it's location to the load_from_checkpoint argument of fit.

A pre-configured instance of the trainer can be obtained from torch_em.default_segmentation_trainer. Alternatively, the trainer class can also be instantiated as in this example:

import torch
from torch_em.loss import DiceLoss
from torch_em.model import UNet2d
from torch_em.data.datasets.light_microscopy import get_dsb_loader
from torch_em.trainer import DefaultTrainer

# The training data will be downloaded to this location.
data_root = "/path/to/save/the/training/data"
patch_shape = (256, 256)

# Create the model and optimizer.
model = UNet2d(in_channels=1, out_channels=1)
optimizer = torch.optim.AdamW(model.parameters())

trainer = DefaultTrainer(
    name="unet-training",
    train_loader=get_dsb_loader(path=data_root, patch_shape=patch_shape, split="train"),
    val_loader=get_dsb_loader(path=data_root, patch_shape=patch_shape, split="test"),
    model=model,
    loss=DiceLoss(),  # The loss function.
    optimizer=optimizer,
    metric=DiceLoss(),  # The metric. The trainer expects smaller values to represent better results.
    device="cuda",  # The device to use for training.
)
trainer.fit(iterations=int(2.5e4))  # Train for 25.000 iterations.
Arguments:
  • name: The name of the checkpoint that will be created by the trainer.
  • train_loader: The data loader containing the training data.
  • val_loader: The data loader containing the validation data.
  • model: The model to train.
  • loss: The loss function for training.
  • optimizer: The optimizer.
  • metric: The metric for validation.
  • device: The torch device to use for training. If None, will use a GPU if available.
  • lr_scheduler: The learning rate scheduler.
  • log_image_interval: The interval for saving images during logging, in training iterations.
  • mixed_precision: Whether to train with mixed precision.
  • early_stopping: The patience for early stopping in epochs. If None, early stopping will not be used.
  • logger: The logger class. Will be instantiated for logging. By default uses torch_em.training.tensorboard_logger.TensorboardLogger.
  • logger_kwargs: The keyword arguments for the logger class.
  • id_: Unique identifier for the trainer. If None then name will be used.
  • save_root: The root folder for saving the checkpoint and logs.
  • compile_model: Whether to compile the model before training.
  • rank: Rank argument for distributed training. See torch_em.multi_gpu_training for details.
  • mixed_precision_dtype: The dtype for autocast in mixed precision training, 'float16' or 'bfloat16'. The default is 'float16' on the GPU and 'bfloat16' on the CPU. Use 'bfloat16' to avoid overflows.
DefaultTrainer( name: Optional[str], train_loader: torch.utils.data.dataloader.DataLoader, val_loader: torch.utils.data.dataloader.DataLoader, model: torch.nn.modules.module.Module, loss: torch.nn.modules.module.Module, optimizer: torch.optim.optimizer.Optimizer, metric: Callable, device: Union[str, torch.device], lr_scheduler: Optional[torch.optim.lr_scheduler._LRScheduler] = None, log_image_interval: int = 100, mixed_precision: bool = True, early_stopping: Optional[int] = None, logger=<class 'torch_em.trainer.tensorboard_logger.TensorboardLogger'>, logger_kwargs: Optional[Dict[str, Any]] = None, id_: Optional[str] = None, save_root: Optional[str] = None, compile_model: Union[bool, str, NoneType] = None, rank: Optional[int] = None, mixed_precision_dtype: Optional[str] = None)
 87    def __init__(
 88        self,
 89        name: Optional[str],
 90        train_loader: torch.utils.data.DataLoader,
 91        val_loader: torch.utils.data.DataLoader,
 92        model: torch.nn.Module,
 93        loss: torch.nn.Module,
 94        optimizer: torch.optim.Optimizer,
 95        metric: Callable,
 96        device: Union[str, torch.device],
 97        lr_scheduler: Optional[torch.optim.lr_scheduler._LRScheduler] = None,
 98        log_image_interval: int = 100,
 99        mixed_precision: bool = True,
100        early_stopping: Optional[int] = None,
101        logger=TensorboardLogger,
102        logger_kwargs: Optional[Dict[str, Any]] = None,
103        id_: Optional[str] = None,
104        save_root: Optional[str] = None,
105        compile_model: Optional[Union[bool, str]] = None,
106        rank: Optional[int] = None,
107        mixed_precision_dtype: Optional[str] = None,
108    ):
109        if name is None and not issubclass(logger, WandbLogger):
110            raise TypeError("Name cannot be None if not using the WandbLogger")
111
112        self._generate_name = name is None
113        self.name = name
114        self.id_ = id_ or name
115        self.train_loader = train_loader
116        self.val_loader = val_loader
117        self.model = model
118        self.loss = loss
119        self.optimizer = optimizer
120        self.metric = metric
121        self.device = torch.device(device)
122        self.lr_scheduler = lr_scheduler
123        self.log_image_interval = log_image_interval
124        self.save_root = save_root
125        self.compile_model = compile_model
126        self.rank = rank
127        self._device_type = "cpu" if self.device.type == "cpu" else "cuda"
128
129        self._iteration = 0
130        self._epoch = 0
131        self._best_epoch = 0
132
133        self.mixed_precision = mixed_precision
134        # These are the defaults of torch.autocast for each device type.
135        self.mixed_precision_dtype = mixed_precision_dtype or ("bfloat16" if self._device_type == "cpu" else "float16")
136        self.early_stopping = early_stopping
137        self.train_time = 0.0
138
139        if mixed_precision:
140            # Only float16 needs gradient scaling. bfloat16 has the same range as float32.
141            self.scaler = torch.GradScaler(self._device_type, enabled=self.mixed_precision_dtype == "float16")
142        else:
143            self.scaler = None
144
145        self.logger_class = logger
146        self.logger_kwargs = logger_kwargs
147        self.log_image_interval = log_image_interval
name
id_
train_loader
val_loader
model
loss
optimizer
metric
device
lr_scheduler
log_image_interval
save_root
compile_model
rank
mixed_precision
mixed_precision_dtype
early_stopping
train_time
logger_class
logger_kwargs
checkpoint_folder
149    @property
150    def checkpoint_folder(self):
151        assert self.id_ is not None  # Because the logger may generate and set trainer.id on logger.__init__.
152        # Save_root enables saving the checkpoints somewhere else than in the local folder.
153        # This is handy for filesystems with limited space, where saving the checkpoints
154        # and log files can lead to running out of space.
155        save_root = getattr(self, "save_root", None)
156        return os.path.join("./checkpoints", self.id_) if save_root is None else\
157            os.path.join(save_root, "./checkpoints", self.id_)
iteration
159    @property
160    def iteration(self):
161        return self._iteration
epoch
163    @property
164    def epoch(self):
165        return self._epoch
def fit( self, iterations: Optional[int] = None, load_from_checkpoint: Union[str, os.PathLike, NoneType] = None, epochs: Optional[int] = None, save_every_kth_epoch: Optional[int] = None, progress=None, overwrite_training: bool = True):
654    def fit(
655        self,
656        iterations: Optional[int] = None,
657        load_from_checkpoint: Optional[Union[os.PathLike, str]] = None,
658        epochs: Optional[int] = None,
659        save_every_kth_epoch: Optional[int] = None,
660        progress=None,
661        overwrite_training: bool = True,
662    ):
663        """Run neural network training.
664
665        Exactly one of 'iterations' or 'epochs' has to be passed.
666
667        Args:
668            iterations: How long to train, specified in iterations.
669            load_from_checkpoint: Path to a checkpoint from where training should be continued .
670            epochs: How long to train, specified in epochs.
671            save_every_kth_epoch: Save checkpoints after every kth epoch in a separate file.
672                The corresponding checkpoints will be saved with the naming scheme 'epoch-{epoch}.pt'.
673            progress: Optional progress bar for integration with external tools. Expected to follow the tqdm interface.
674            overwrite_training: Whether to overwrite existing checkpoints in the save directory.
675        """
676        best_metric = self._initialize(iterations, load_from_checkpoint, epochs)
677
678        if not overwrite_training:
679            if load_from_checkpoint is not None:
680                raise ValueError(
681                    "We do not support 'overwrite_training=False' and 'load_from_checkpoint' at the same time."
682                )
683
684            if self._verify_if_training_completed():
685                print(
686                    f"The model is trained for {self.max_iteration} iterations / {self.max_epoch} epochs "
687                    "and 'overwrite_training' is set to 'False'."
688                )
689                print(f"The checkpoints are located at '{os.path.abspath(self.checkpoint_folder)}'.")
690                return
691
692        print(
693            "Start fitting for",
694            self.max_iteration - self._iteration,
695            "iterations / ",
696            self.max_epoch - self._epoch,
697            "epochs",
698        )
699        print("with", len(self.train_loader), "iterations per epoch")
700
701        if self.mixed_precision:
702            train_epoch = self._train_epoch_mixed
703            validate = self._validate_mixed
704            print("Training with mixed precision")
705        else:
706            train_epoch = self._train_epoch
707            validate = self._validate
708            print("Training with single precision")
709
710        total_iterations = epochs * len(self.train_loader) if iterations is None else iterations
711        if progress is None:
712            progress = tqdm(total=total_iterations, desc=f"Epoch {self._epoch}", leave=True)
713        else:
714            progress.total = total_iterations
715            progress.set_description(f"Epoch {self._epoch}")
716
717        msg = "Epoch %i: average [s/it]: %f, current metric: %f, best metric: %f"
718        train_epochs = self.max_epoch - self._epoch
719        t_start = time.time()
720        for epoch in range(train_epochs):
721
722            # Ensure data is shuffled differently at each epoch.
723            try:
724                self.train_loader.sampler.set_epoch(epoch)
725            except AttributeError:
726                pass
727
728            # Run training and validation for this epoch
729            t_per_iter = train_epoch(progress)
730            current_metric = validate()
731
732            # perform all the post-epoch steps:
733
734            # apply the learning rate scheduler
735            if self.lr_scheduler is not None:
736                self.lr_scheduler.step(current_metric)
737
738            # how long did we train in total?
739            total_train_time = (time.time() - t_start) + self.train_time
740
741            # save this checkpoint as the new best checkpoint if
742            # it has the best overall validation metric
743            if current_metric < best_metric:
744                best_metric = current_metric
745                self._best_epoch = self._epoch
746                self.save_checkpoint("best", current_metric, best_metric, train_time=total_train_time)
747
748            # save this checkpoint as the latest checkpoint
749            self.save_checkpoint("latest", current_metric, best_metric, train_time=total_train_time)
750
751            # if we save after every k-th epoch then check if we need to save now
752            if save_every_kth_epoch is not None and (self._epoch + 1) % save_every_kth_epoch == 0:
753                self.save_checkpoint(
754                    f"epoch-{self._epoch + 1}", current_metric, best_metric, train_time=total_train_time
755                )
756
757            # if early stopping has been specified then check if the stopping condition is met
758            if self.early_stopping is not None:
759                epochs_since_best = self._epoch - self._best_epoch
760                if epochs_since_best > self.early_stopping:
761                    print("Stopping training because there has been no improvement for", self.early_stopping, "epochs")
762                    break
763
764            self._epoch += 1
765            progress.set_description(msg % (self._epoch, t_per_iter, current_metric, best_metric), refresh=True)
766
767        print(f"Finished training after {self._epoch} epochs / {self._iteration} iterations.")
768        print(f"The best epoch is number {self._best_epoch}.")
769
770        if self._generate_name:
771            self.name = None
772
773        # Update the train time
774        self.train_time = total_train_time
775
776        # TODO save the model to wandb if we have the wandb logger
777        if isinstance(self.logger, WandbLogger):
778            self.logger.get_wandb().finish()

Run neural network training.

Exactly one of 'iterations' or 'epochs' has to be passed.

Arguments:
  • iterations: How long to train, specified in iterations.
  • load_from_checkpoint: Path to a checkpoint from where training should be continued .
  • epochs: How long to train, specified in epochs.
  • save_every_kth_epoch: Save checkpoints after every kth epoch in a separate file. The corresponding checkpoints will be saved with the naming scheme 'epoch-{epoch}.pt'.
  • progress: Optional progress bar for integration with external tools. Expected to follow the tqdm interface.
  • overwrite_training: Whether to overwrite existing checkpoints in the save directory.
class DefaultTrainer.Deserializer:
167    class Deserializer:
168        """Determines how to deserialize the trainer kwargs from serialized 'init_data'.
169
170        Examples:
171            To extend the initialization process you can inherite from this Deserializer in an inherited Trainer class.
172            Note that `DefaultTrainer.Deserializer.load_generic()` covers most cases already.
173
174            This example adds `the_answer` kwarg, which requires 'calculations' upon initialization:
175            >>> class MyTrainer(DefaultTrainer):
176            >>>     def __init__(self, *args, the_answer: int, **kwargs):
177            >>>         super().__init__(*args, **kwargs)
178            >>>         self.the_answer = the_answer  # this allows the default Serializer to save the new kwarg,
179            >>>                                       # see DefaultTrainer.Serializer
180            >>>
181            >>>     class Deserializer(DefaultTrainer.Deserializer):
182            >>>         def load_the_answer(self):
183            >>>             generic_answer = self.init_data["the_answer"]
184            >>>             # (device dependent) special deserialization
185            >>>             if self.trainer_kwargs["device"].type == "cpu":  # accessing previously deserialized kwarg
186            >>>                 self.trainer_kwargs["the_answer"] = generic_answer + 1
187            >>>             else:
188            >>>                 self.trainer_kwargs["the_answer"] = generic_answer * 2
189
190        Args:
191            init_data: The initialization data of the trainer.
192            save_path: The path where the checkpoint was saved.
193            device: The device.
194        """
195
196        def __init__(self, init_data: Dict, save_path: str, device: Union[str, torch.device]):
197            self.init_data = init_data
198            self.save_path = save_path
199            # Populate with deserialized trainer kwargs during deserialization; possibly overwrite 'device'.
200            self.trainer_kwargs: Dict[str, Any] = dict(
201                device=torch.device(self.init_data["device"]) if device is None else torch.device(device)
202            )
203
204        def load(self, kwarg_name: str, optional):
205            """@private
206            """
207            # `optional` is True if self.trainer.__class__.__init__ specifies a default value for 'kwarg_name'
208            if kwarg_name == "device":
209                pass  # deserialized in __init__
210            elif kwarg_name.endswith("_loader"):
211                self.load_data_loader(kwarg_name, optional)
212            else:
213                load = getattr(self, f"load_{kwarg_name}", self.load_generic)
214                load(kwarg_name, optional=optional)
215
216        def load_data_loader(self, loader_name, optional) -> None:
217            """@private
218            """
219            ds = self.init_data.get(loader_name.replace("_loader", "_dataset"))
220            if ds is None and optional:
221                return
222
223            loader_kwargs = self.init_data[f"{loader_name}_kwargs"]
224            loader = torch.utils.data.DataLoader(ds, **loader_kwargs)
225            # monkey patch shuffle loader_name to the loader
226            loader.shuffle = loader_kwargs.get("shuffle", False)
227            self.trainer_kwargs[loader_name] = loader
228
229        def load_generic(
230            self,
231            kwarg_name: str,
232            *dynamic_args: Dict,
233            optional: bool,
234            only_class: bool = False,
235            dynamic_kwargs: Optional[Dict[str, Any]] = None,
236        ) -> None:
237            """@private
238            """
239            if kwarg_name in self.init_data:
240                self.trainer_kwargs[kwarg_name] = self.init_data[kwarg_name]
241                return
242
243            this_cls = self.init_data.get(f"{kwarg_name}_class", None)
244            if this_cls is None:
245                if optional:
246                    return
247                else:
248                    raise RuntimeError(f"Could not find init data for {kwarg_name} in {self.save_path}")
249
250            assert isinstance(this_cls, str), this_cls
251            assert "." in this_cls, this_cls
252            cls_p, cls_m = this_cls.rsplit(".", 1)
253            this_cls = getattr(import_module(cls_p), cls_m)
254            if only_class:
255                self.trainer_kwargs[kwarg_name] = this_cls
256            else:
257                self.trainer_kwargs[kwarg_name] = this_cls(
258                    *dynamic_args, **self.init_data.get(f"{kwarg_name}_kwargs", {}), **(dynamic_kwargs or {})
259                )
260
261        def load_name(self, kwarg_name: str, optional: bool):
262            """@private
263            """
264            self.trainer_kwargs[kwarg_name] = os.path.split(os.path.dirname(self.save_path))[1]
265
266        def load_optimizer(self, kwarg_name: str, optional: bool):
267            """@private
268            """
269            self.load_generic(kwarg_name, self.trainer_kwargs["model"].parameters(), optional=optional)
270
271        def load_lr_scheduler(self, kwarg_name: str, optional: bool):
272            """@private
273            """
274            self.load_generic(kwarg_name, self.trainer_kwargs["optimizer"], optional=optional)
275
276        # todo: remove and rename kwarg 'logger' to 'logger_class'
277        def load_logger(self, kwarg_name: str, optional: bool):
278            """@private
279            """
280            assert kwarg_name == "logger"
281            self.load_generic("logger", optional=optional, only_class=True)

Determines how to deserialize the trainer kwargs from serialized 'init_data'.

Examples:

To extend the initialization process you can inherite from this Deserializer in an inherited Trainer class. Note that DefaultTrainer.Deserializer.load_generic() covers most cases already.

This example adds the_answer kwarg, which requires 'calculations' upon initialization:

>>> class MyTrainer(DefaultTrainer):
>>>     def __init__(self, *args, the_answer: int, **kwargs):
>>>         super().__init__(*args, **kwargs)
>>>         self.the_answer = the_answer  # this allows the default Serializer to save the new kwarg,
>>>                                       # see DefaultTrainer.Serializer
>>>
>>>     class Deserializer(DefaultTrainer.Deserializer):
>>>         def load_the_answer(self):
>>>             generic_answer = self.init_data["the_answer"]
>>>             # (device dependent) special deserialization
>>>             if self.trainer_kwargs["device"].type == "cpu":  # accessing previously deserialized kwarg
>>>                 self.trainer_kwargs["the_answer"] = generic_answer + 1
>>>             else:
>>>                 self.trainer_kwargs["the_answer"] = generic_answer * 2
Arguments:
  • init_data: The initialization data of the trainer.
  • save_path: The path where the checkpoint was saved.
  • device: The device.
DefaultTrainer.Deserializer(init_data: Dict, save_path: str, device: Union[str, torch.device])
196        def __init__(self, init_data: Dict, save_path: str, device: Union[str, torch.device]):
197            self.init_data = init_data
198            self.save_path = save_path
199            # Populate with deserialized trainer kwargs during deserialization; possibly overwrite 'device'.
200            self.trainer_kwargs: Dict[str, Any] = dict(
201                device=torch.device(self.init_data["device"]) if device is None else torch.device(device)
202            )
init_data
save_path
trainer_kwargs: Dict[str, Any]
class DefaultTrainer.Serializer:
333    class Serializer:
334        """Implements how to serialize trainer kwargs from a trainer instance.
335
336        Examples:
337            To extend the serialization process you can inherite from this Serializer in a derived Trainer class.
338            Note that the methods `dump_generic_builtin()`, `dump_generic_class()` and `dump_generic_instance()`
339            called by the `dump()` method when appropriate cover most cases already.
340
341            This example adds `the_answer` kwarg, which requires extra steps on dumping only because we don't keep a
342            'the_answer' attribute:
343            >>> class MyTrainer(DefaultTrainer):
344            >>>     def __init__(self, *args, the_answer: int, **kwargs):
345            >>>         super().__init__(*args, **kwargs)
346            >>>         # self.the_answer = the_answer  # this would allow the default Serializer to save the new kwarg,
347            >>>         # but let's make things more interesting...
348            >>>         self.the = the_answer // 10
349            >>>         self.answer = the_answer % 10
350            >>>
351            >>>     class Serializer(DefaultTrainer.Serializer):
352            >>>         trainer: MyTrainer
353            >>>         def dump_the_answer(self, kwarg_name: str) -> None:  # custom dump method for 'the_answer' kwarg
354            >>>             assert kwarg_name == "the_answer"
355            >>>             # populate self.init_data with the serialized data required by Deserializer
356            >>>             # to restore the trainer kwargs
357            >>>             self.init_data["the_answer"] = self.trainer.the * 10 + self.trainer.answer
358
359            This example with both Serializer and Deserializer adds `the_answer` kwarg,
360            while saving it in two separate entries 'the' and 'answer'
361            >>> class MyTrainer(DefaultTrainer):
362            >>>     def __init__(self, *args, the_answer: int, **kwargs):
363            >>>         super().__init__(*args, **kwargs)
364            >>>         self.the_answer = the_answer
365            >>>
366            >>>     class Serializer(DefaultTrainer.Serializer):
367            >>>         trainer: MyTrainer
368            >>>         def dump_the_answer(self, kwarg_name: str):
369            >>>             assert kwarg_name == "the_answer"
370            >>>             self.init_data.update({
371            >>>                 "the": self.trainer.the_answer // 10,
372            >>>                 "answer": self.trainer.the_answer % 10
373            >>>             })
374            >>>
375            >>>     class Deserializer(DefaultTrainer.Deserializer):
376            >>>         def load_the_answer(self, kwarg_name: str, optional: bool):
377            >>>             assert kwarg_name == "the_answer"
378            >>>             # 'optional' is True if MyTrainer.__init__ specifies a default value for 'kwarg_name'
379            >>>             self.trainer_kwargs[kwarg_name] = self.init_data["the"] * 10 + self.init_data["answer"]
380
381        Args:
382            trainer: The trainer instance.
383        """
384
385        def __init__(self, trainer: DefaultTrainer):
386            self.trainer = trainer
387            self.init_data = {}  # to be populated during serialization process
388
389        def dump(self, kwarg_name: str) -> None:
390            """@private
391            """
392            dumper = getattr(self, f"dump_{kwarg_name}", None)
393            if dumper is not None:
394                dumper(kwarg_name)
395            elif kwarg_name.endswith("_loader"):
396                self.dump_data_loader(kwarg_name)
397            elif kwarg_name.endswith("_class"):
398                self.dump_generic_class(kwarg_name)
399            elif not hasattr(self.trainer, kwarg_name):
400                raise AttributeError(
401                    f"{self.trainer.__class__} missing attribute '{kwarg_name}' "
402                    f"or special dump method {self.trainer.__class__}.Serializer.dump_{kwarg_name}()"
403                )
404            else:
405                assert hasattr(self.trainer, kwarg_name)
406                obj = getattr(self.trainer, kwarg_name)
407                if obj is None or type(obj) in (
408                    bool,
409                    bytearray,
410                    bytes,
411                    dict,
412                    float,
413                    frozenset,
414                    int,
415                    list,
416                    set,
417                    str,
418                    tuple,
419                ):
420                    self.dump_generic_builtin(kwarg_name)
421                else:
422                    self.dump_generic_instance(kwarg_name)
423
424        def dump_generic_builtin(self, kwarg_name: str) -> None:
425            """@private
426            """
427            assert hasattr(self.trainer, kwarg_name)
428            self.init_data[kwarg_name] = getattr(self.trainer, kwarg_name)
429
430        def dump_generic_class(self, kwarg_name: str) -> None:
431            """@private
432            """
433            assert hasattr(self.trainer, kwarg_name)
434            assert kwarg_name.endswith("_class")
435            obj = getattr(self.trainer, kwarg_name)
436            self.init_data[kwarg_name] = None if obj is None else f"{obj.__module__}.{obj.__name__}"
437
438        def dump_generic_instance(self, kwarg_name: str) -> None:
439            """@private
440            """
441            assert hasattr(self.trainer, kwarg_name)
442            instance = getattr(self.trainer, kwarg_name)
443            self.init_data.update(
444                {
445                    f"{kwarg_name}_class": f"{instance.__class__.__module__}.{instance.__class__.__name__}",
446                    f"{kwarg_name}_kwargs": get_constructor_arguments(instance),
447                }
448            )
449
450        def dump_device(self, kwarg_name: str):
451            """@private
452            """
453            assert hasattr(self.trainer, kwarg_name)
454            self.init_data[kwarg_name] = str(getattr(self.trainer, kwarg_name))
455
456        def dump_data_loader(self, kwarg_name: str) -> None:
457            """@private
458            """
459            assert hasattr(self.trainer, kwarg_name)
460            loader = getattr(self.trainer, kwarg_name)
461            if loader is None:
462                return
463            self.init_data.update(
464                {
465                    f"{kwarg_name.replace('_loader', '_dataset')}": loader.dataset,
466                    f"{kwarg_name}_kwargs": get_constructor_arguments(loader),
467                }
468            )
469
470        def dump_logger(self, kwarg_name: str):  # todo: remove and rename kwarg 'logger' to 'logger_class'
471            """@private
472            """
473            self.dump_generic_class(f"{kwarg_name}_class")
474
475        def dump_model(self, kwarg_name: str):
476            """@private
477            """
478            if is_compiled(self.trainer.model):
479                self.init_data.update(
480                    {"model_class": self.trainer._model_class, "model_kwargs": self.trainer._model_kwargs}
481                )
482            else:
483                self.dump_generic_instance("model")

Implements how to serialize trainer kwargs from a trainer instance.

Examples:

To extend the serialization process you can inherite from this Serializer in a derived Trainer class. Note that the methods dump_generic_builtin(), dump_generic_class() and dump_generic_instance() called by the dump() method when appropriate cover most cases already.

This example adds the_answer kwarg, which requires extra steps on dumping only because we don't keep a 'the_answer' attribute:

>>> class MyTrainer(DefaultTrainer):
>>>     def __init__(self, *args, the_answer: int, **kwargs):
>>>         super().__init__(*args, **kwargs)
>>>         # self.the_answer = the_answer  # this would allow the default Serializer to save the new kwarg,
>>>         # but let's make things more interesting...
>>>         self.the = the_answer // 10
>>>         self.answer = the_answer % 10
>>>
>>>     class Serializer(DefaultTrainer.Serializer):
>>>         trainer: MyTrainer
>>>         def dump_the_answer(self, kwarg_name: str) -> None:  # custom dump method for 'the_answer' kwarg
>>>             assert kwarg_name == "the_answer"
>>>             # populate self.init_data with the serialized data required by Deserializer
>>>             # to restore the trainer kwargs
>>>             self.init_data["the_answer"] = self.trainer.the * 10 + self.trainer.answer

This example with both Serializer and Deserializer adds the_answer kwarg, while saving it in two separate entries 'the' and 'answer'

>>> class MyTrainer(DefaultTrainer):
>>>     def __init__(self, *args, the_answer: int, **kwargs):
>>>         super().__init__(*args, **kwargs)
>>>         self.the_answer = the_answer
>>>
>>>     class Serializer(DefaultTrainer.Serializer):
>>>         trainer: MyTrainer
>>>         def dump_the_answer(self, kwarg_name: str):
>>>             assert kwarg_name == "the_answer"
>>>             self.init_data.update({
>>>                 "the": self.trainer.the_answer // 10,
>>>                 "answer": self.trainer.the_answer % 10
>>>             })
>>>
>>>     class Deserializer(DefaultTrainer.Deserializer):
>>>         def load_the_answer(self, kwarg_name: str, optional: bool):
>>>             assert kwarg_name == "the_answer"
>>>             # 'optional' is True if MyTrainer.__init__ specifies a default value for 'kwarg_name'
>>>             self.trainer_kwargs[kwarg_name] = self.init_data["the"] * 10 + self.init_data["answer"]
Arguments:
  • trainer: The trainer instance.
DefaultTrainer.Serializer(trainer: DefaultTrainer)
385        def __init__(self, trainer: DefaultTrainer):
386            self.trainer = trainer
387            self.init_data = {}  # to be populated during serialization process
trainer
init_data