MCPcopy Create free account
hub / github.com/Alpha-VLLM/LLaMA2-Accessory / update_ema

Function update_ema

Large-DiT-ImageNet/train.py:77–87  ·  view source on GitHub ↗

Step the EMA model towards the current model.

(ema_model, model, decay=0.9999)

Source from the content-addressed store, hash-verified

75
76@torch.no_grad()
77def update_ema(ema_model, model, decay=0.9999):
78 """
79 Step the EMA model towards the current model.
80 """
81 ema_params = OrderedDict(ema_model.named_parameters())
82 model_params = OrderedDict(model.named_parameters())
83 assert set(ema_params.keys()) == set(model_params.keys())
84
85 for name, param in model_params.items():
86 # TODO: Consider applying only to params that require_grad to avoid small numerical changes of pos_embed
87 ema_params[name].mul_(decay).add_(param.data, alpha=1 - decay)
88
89
90def cleanup():

Callers 1

mainFunction · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected