MCPcopy Create free account
hub / github.com/MotrixLab/ViMoGen / update_ema

Function update_ema

trainer/base_trainer.py:389–408  ·  view source on GitHub ↗

Step the EMA model towards the current model.

(
    ema_model: torch.nn.Module,
    model: torch.nn.Module,
    decay: float = 0.9999,
    sharded: bool = True,
)

Source from the content-addressed store, hash-verified

387# https://github.com/hpcaitech/Open-Sora/blob/main/opensora/utils/train_utils.py#L44
388@torch.no_grad()
389def 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
411def linear_lr_warmpup(warmup_steps):

Callers 1

mainFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected