Team Ai
Apppublic

dinhdat1110/diffusion-model

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes
ema.py76 linesDownload Raw Back to utils
1from pytorch_lightning.callbacks import Callback2from timm.utils.model import get_state_dict, unwrap_model3from timm.utils.model_ema import ModelEmaV24# Cell5 6 7class EMACallback(Callback):8    """9    Model Exponential Moving Average. Empirically it has been found that using the moving average10    of the trained parameters of a deep network is better than using its trained parameters directly.11 12    If `use_ema_weights`, then the ema parameters of the network is set after training end.13    """14 15    def __init__(self, decay=0.9999, use_ema_weights: bool = True):16        self.decay = decay17        self.ema = None18        self.use_ema_weights = use_ema_weights19 20    def on_fit_start(self, trainer, pl_module, *args):21        "Initialize `ModelEmaV2` from timm to keep a copy of the moving average of the weights"22        self.ema = ModelEmaV2(pl_module, decay=self.decay, device=None)23 24    def on_train_batch_end(25        self, trainer, pl_module, *args26    ):27        "Update the stored parameters using a moving average"28        # Update currently maintained parameters.29        self.ema.update(pl_module)30 31    def on_validation_epoch_start(self, trainer, pl_module, *args):32        "do validation using the stored parameters"33        # save original parameters before replacing with EMA version34        self.store(pl_module.parameters())35 36        # update the LightningModule with the EMA weights37        # ~ Copy EMA parameters to LightningModule38        self.copy_to(self.ema.module.parameters(), pl_module.parameters())39 40    def on_validation_end(self, trainer, pl_module, *args):41        "Restore original parameters to resume training later"42        self.restore(pl_module.parameters())43 44    def on_train_end(self, trainer, pl_module, *args):45        # update the LightningModule with the EMA weights46        if self.use_ema_weights:47            self.copy_to(self.ema.module.parameters(), pl_module.parameters())48            msg = "Model weights replaced with the EMA version."49 50    def on_save_checkpoint(self, trainer, pl_module, checkpoint, *args):51        if self.ema is not None:52            return {"state_dict_ema": get_state_dict(self.ema, unwrap_model)}53 54    def on_load_checkpoint(self, callback_state, *args):55        if self.ema is not None:56            self.ema.module.load_state_dict(callback_state["state_dict_ema"])57 58    def store(self, parameters):59        "Save the current parameters for restoring later."60        self.collected_params = [param.clone() for param in parameters]61 62    def restore(self, parameters):63        """64        Restore the parameters stored with the `store` method.65        Useful to validate the model with EMA parameters without affecting the66        original optimization process.67        """68        for c_param, param in zip(self.collected_params, parameters):69            param.data.copy_(c_param.data)70 71    def copy_to(self, shadow_parameters, parameters):72        "Copy current parameters into given collection of parameters."73        for s_param, param in zip(shadow_parameters, parameters):74            if param.requires_grad:75                param.data.copy_(s_param.data)76