| 356 | |
| 357 | class SpeedMonitorCallback(Callback): |
| 358 | def __init__( |
| 359 | self, length_fn: Callable[[Any], int], batch_size: int, **kwargs: Any |
| 360 | ) -> None: |
| 361 | super().__init__() |
| 362 | self.speed_monitor: Optional[SpeedMonitorBase] = None |
| 363 | self.speed_monitor_kwargs = kwargs |
| 364 | self.length_fn = length_fn |
| 365 | self.batch_size = batch_size |
| 366 | self.eval_t0: int = 0 |
| 367 | self.train_t0: int = 0 |
| 368 | self.total_lengths: int = 0 |
| 369 | |
| 370 | def setup(self, trainer: Trainer, pl_module: LightningModule, stage: str) -> None: |
| 371 | if self.speed_monitor is not None: |