Team Ai
Apppublic

Dynamatrix/DiffBIR-OpenXLab

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
0likes
callbacks.py75 linesDownload Raw Back to model
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