| 514 | |
| 515 | class ImageLogger(Callback): |
| 516 | def __init__(self, batch_frequency, max_images, clamp=True, increase_log_steps=True, |
| 517 | rescale=True, disabled=False, log_on_batch_idx=False, log_first_step=False, |
| 518 | log_images_kwargs=None): |
| 519 | super().__init__() |
| 520 | self.rescale = rescale |
| 521 | self.batch_freq = batch_frequency |
| 522 | self.max_images = max_images |
| 523 | self.save_freq = 250 |
| 524 | self.logger_log_images = { |
| 525 | pl.loggers.TestTubeLogger: self._testtube, |
| 526 | } |
| 527 | self.log_steps = [2 ** n for n in range(int(np.log2(self.batch_freq)) + 1)] |
| 528 | if not increase_log_steps: |
| 529 | self.log_steps = [self.batch_freq] |
| 530 | self.clamp = clamp |
| 531 | self.disabled = disabled |
| 532 | self.log_on_batch_idx = log_on_batch_idx |
| 533 | self.log_images_kwargs = log_images_kwargs if log_images_kwargs else {} |
| 534 | self.log_first_step = log_first_step |
| 535 | |
| 536 | @rank_zero_only |
| 537 | def _testtube(self, pl_module, images, batch_idx, split): |