Dynamatrix/DiffBIR-OpenXLab
0
1from typing import Dict, Any2import os3 4import numpy as np5import pytorch_lightning as pl6from pytorch_lightning.callbacks import ModelCheckpoint7from pytorch_lightning.utilities.types import STEP_OUTPUT8import torch9import torchvision10from PIL import Image11from pytorch_lightning.callbacks import Callback12from pytorch_lightning.utilities.distributed import rank_zero_only13 14from .mixins import ImageLoggerMixin15 16 17__all__ = [18 "ModelCheckpoint",19 "ImageLogger"20]21 22class ImageLogger(Callback):23 """24 Log images during training or validating.25 26 TODO: Support validating.27 """28 29 def __init__(30 self,31 log_every_n_steps: int=2000,32 max_images_each_step: int=4,33 log_images_kwargs: Dict[str, Any]=None34 ) -> "ImageLogger":35 super().__init__()36 self.log_every_n_steps = log_every_n_steps37 self.max_images_each_step = max_images_each_step38 self.log_images_kwargs = log_images_kwargs or dict()39 40 def on_fit_start(self, trainer: pl.Trainer, pl_module: pl.LightningModule) -> None:41 assert isinstance(pl_module, ImageLoggerMixin)42 43 @rank_zero_only44 def on_train_batch_end(45 self, trainer: pl.Trainer, pl_module: pl.LightningModule, outputs: STEP_OUTPUT,46 batch: Any, batch_idx: int, dataloader_idx: int47 ) -> None:48 if pl_module.global_step % self.log_every_n_steps == 0:49 is_train = pl_module.training50 if is_train:51 pl_module.freeze()52 53 with torch.no_grad():54 # returned images should be: nchw, rgb, [0, 1]55 images: Dict[str, torch.Tensor] = pl_module.log_images(batch, **self.log_images_kwargs)56 57 # save images58 save_dir = os.path.join(pl_module.logger.save_dir, "image_log", "train")59 os.makedirs(save_dir, exist_ok=True)60 for image_key in images:61 image = images[image_key].detach().cpu()62 N = min(self.max_images_each_step, len(image))63 grid = torchvision.utils.make_grid(image[:N], nrow=4)64 # chw -> hwc (hw if gray)65 grid = grid.transpose(0, 1).transpose(1, 2).squeeze(-1).numpy()66 grid = (grid * 255).clip(0, 255).astype(np.uint8)67 filename = "{}_step-{:06}_e-{:06}_b-{:06}.png".format(68 image_key, pl_module.global_step, pl_module.current_epoch, batch_idx69 )70 path = os.path.join(save_dir, filename)71 Image.fromarray(grid).save(path)72 73 if is_train:74 pl_module.unfreeze()75 