dinhdat1110/diffusion-model
0
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 