|
|
|
|
| from ultralytics.utils import LOGGER, SETTINGS, TESTS_RUNNING, colorstr
|
|
|
| try:
|
|
|
| from torch.utils.tensorboard import SummaryWriter
|
|
|
| assert not TESTS_RUNNING
|
| assert SETTINGS["tensorboard"] is True
|
| WRITER = None
|
| PREFIX = colorstr("TensorBoard: ")
|
|
|
|
|
| import warnings
|
| from copy import deepcopy
|
|
|
| from ultralytics.utils.torch_utils import de_parallel, torch
|
|
|
| except (ImportError, AssertionError, TypeError, AttributeError):
|
|
|
|
|
| SummaryWriter = None
|
|
|
|
|
| def _log_scalars(scalars: dict, step: int = 0) -> None:
|
| """
|
| Log scalar values to TensorBoard.
|
|
|
| Args:
|
| scalars (dict): Dictionary of scalar values to log to TensorBoard. Keys are scalar names and values are the
|
| corresponding scalar values.
|
| step (int): Global step value to record with the scalar values. Used for x-axis in TensorBoard graphs.
|
|
|
| Examples:
|
| >>> # Log training metrics
|
| >>> metrics = {"loss": 0.5, "accuracy": 0.95}
|
| >>> _log_scalars(metrics, step=100)
|
| """
|
| if WRITER:
|
| for k, v in scalars.items():
|
| WRITER.add_scalar(k, v, step)
|
|
|
|
|
| def _log_tensorboard_graph(trainer) -> None:
|
| """
|
| Log model graph to TensorBoard.
|
|
|
| This function attempts to visualize the model architecture in TensorBoard by tracing the model with a dummy input
|
| tensor. It first tries a simple method suitable for YOLO models, and if that fails, falls back to a more complex
|
| approach for models like RTDETR that may require special handling.
|
|
|
| Args:
|
| trainer (BaseTrainer): The trainer object containing the model to visualize. Must have attributes:
|
| - model: PyTorch model to visualize
|
| - args: Configuration arguments with 'imgsz' attribute
|
|
|
| Notes:
|
| This function requires TensorBoard integration to be enabled and the global WRITER to be initialized.
|
| It handles potential warnings from the PyTorch JIT tracer and attempts to gracefully handle different
|
| model architectures.
|
| """
|
|
|
| imgsz = trainer.args.imgsz
|
| imgsz = (imgsz, imgsz) if isinstance(imgsz, int) else imgsz
|
| p = next(trainer.model.parameters())
|
| im = torch.zeros((1, 3, *imgsz), device=p.device, dtype=p.dtype)
|
|
|
| with warnings.catch_warnings():
|
| warnings.simplefilter("ignore", category=UserWarning)
|
| warnings.simplefilter("ignore", category=torch.jit.TracerWarning)
|
|
|
|
|
| try:
|
| trainer.model.eval()
|
| WRITER.add_graph(torch.jit.trace(de_parallel(trainer.model), im, strict=False), [])
|
| LOGGER.info(f"{PREFIX}model graph visualization added ✅")
|
| return
|
|
|
| except Exception:
|
|
|
| try:
|
| model = deepcopy(de_parallel(trainer.model))
|
| model.eval()
|
| model = model.fuse(verbose=False)
|
| for m in model.modules():
|
| if hasattr(m, "export"):
|
| m.export = True
|
| m.format = "torchscript"
|
| model(im)
|
| WRITER.add_graph(torch.jit.trace(model, im, strict=False), [])
|
| LOGGER.info(f"{PREFIX}model graph visualization added ✅")
|
| except Exception as e:
|
| LOGGER.warning(f"{PREFIX}WARNING ⚠️ TensorBoard graph visualization failure {e}")
|
|
|
|
|
| def on_pretrain_routine_start(trainer) -> None:
|
| """Initialize TensorBoard logging with SummaryWriter."""
|
| if SummaryWriter:
|
| try:
|
| global WRITER
|
| WRITER = SummaryWriter(str(trainer.save_dir))
|
| LOGGER.info(f"{PREFIX}Start with 'tensorboard --logdir {trainer.save_dir}', view at http://localhost:6006/")
|
| except Exception as e:
|
| LOGGER.warning(f"{PREFIX}WARNING ⚠️ TensorBoard not initialized correctly, not logging this run. {e}")
|
|
|
|
|
| def on_train_start(trainer) -> None:
|
| """Log TensorBoard graph."""
|
| if WRITER:
|
| _log_tensorboard_graph(trainer)
|
|
|
|
|
| def on_train_epoch_end(trainer) -> None:
|
| """Logs scalar statistics at the end of a training epoch."""
|
| _log_scalars(trainer.label_loss_items(trainer.tloss, prefix="train"), trainer.epoch + 1)
|
| _log_scalars(trainer.lr, trainer.epoch + 1)
|
|
|
|
|
| def on_fit_epoch_end(trainer) -> None:
|
| """Logs epoch metrics at end of training epoch."""
|
| _log_scalars(trainer.metrics, trainer.epoch + 1)
|
|
|
|
|
| callbacks = (
|
| {
|
| "on_pretrain_routine_start": on_pretrain_routine_start,
|
| "on_train_start": on_train_start,
|
| "on_fit_epoch_end": on_fit_epoch_end,
|
| "on_train_epoch_end": on_train_epoch_end,
|
| }
|
| if SummaryWriter
|
| else {}
|
| )
|
|
|