Spaces:
Paused
Paused
| from torch.utils.tensorboard import SummaryWriter | |
| from accelerate.tracking import GeneralTracker, on_main_process | |
| import os | |
| from typing import Union | |
| class TbTracker(GeneralTracker): | |
| name = "tensorboard" | |
| requires_logging_directory = True | |
| def __init__(self, run_name: str, logging_dir: Union[str, os.PathLike], | |
| **kwargs): | |
| super().__init__() | |
| self.run_name = run_name | |
| self.logging_dir = os.path.join(logging_dir, run_name) | |
| self.writer = SummaryWriter(self.logging_dir, **kwargs) | |
| def tracker(self): | |
| return self.writer | |
| def add_scalar(self, tag, scalar_value, **kwargs): | |
| self.writer.add_scalar(tag=tag, scalar_value=scalar_value, **kwargs) | |
| def add_text(self, tag, text_string, **kwargs): | |
| self.writer.add_text(tag=tag, text_string=text_string, **kwargs) | |
| def add_figure(self, tag, figure, **kwargs): | |
| self.writer.add_figure(tag=tag, figure=figure, **kwargs) | |