| 491 | |
| 492 | class ImageLogger(Callback): |
| 493 | def __init__(self, batch_frequency, max_images, clamp=True, increase_log_steps=True, |
| 494 | rescale=True, disabled=False, log_on_batch_idx=False, log_first_step=False, |
| 495 | log_images_kwargs=None): |
| 496 | super().__init__() |
| 497 | self.rescale = rescale |
| 498 | self.batch_freq = batch_frequency |
| 499 | self.max_images = max_images |
| 500 | self.save_freq = 250 |
| 501 | self.logger_log_images = { |
| 502 | pl.loggers.TestTubeLogger: self._testtube, |
| 503 | } |
| 504 | self.log_steps = [2 ** n for n in range(int(np.log2(self.batch_freq)) + 1)] |
| 505 | if not increase_log_steps: |
| 506 | self.log_steps = [self.batch_freq] |
| 507 | self.clamp = clamp |
| 508 | self.disabled = disabled |
| 509 | self.log_on_batch_idx = log_on_batch_idx |
| 510 | self.log_images_kwargs = log_images_kwargs if log_images_kwargs else {} |
| 511 | self.log_first_step = log_first_step |
| 512 | |
| 513 | @rank_zero_only |
| 514 | def _testtube(self, pl_module, images, batch_idx, split): |