Step the EMA model towards the current model.
(
ema_model: torch.nn.Module,
model: torch.nn.Module,
decay: float = 0.9999,
sharded: bool = True,
)
| 387 | # https://github.com/hpcaitech/Open-Sora/blob/main/opensora/utils/train_utils.py#L44 |
| 388 | @torch.no_grad() |
| 389 | def update_ema( |
| 390 | ema_model: torch.nn.Module, |
| 391 | model: torch.nn.Module, |
| 392 | decay: float = 0.9999, |
| 393 | sharded: bool = True, |
| 394 | ) -> None: |
| 395 | """Step the EMA model towards the current model.""" |
| 396 | ema_params = OrderedDict(ema_model.named_parameters()) |
| 397 | model_params = OrderedDict(model.named_parameters()) |
| 398 | |
| 399 | for name, param in model_params.items(): |
| 400 | if name == 'pos_embed': |
| 401 | continue |
| 402 | if not param.requires_grad: |
| 403 | continue |
| 404 | param_data = param.data |
| 405 | # assert param_data.dtype == torch.float32 |
| 406 | # TODO get float32 version of parameters from optimizer |
| 407 | ema_params[name].mul_(decay).add_( |
| 408 | param_data.to(torch.float32), alpha=1 - decay) |
| 409 | |
| 410 | |
| 411 | def linear_lr_warmpup(warmup_steps): |