Spaces:
Running on Zero
Running on Zero
Download grn/models/ema.py from hanjian/GRN: direct link, hf CLI and curl.
- Browser
- Download file 716 Bytes
-
https://huggingface.co/spaces/hanjian/GRN/resolve/main/grn/models/ema.py
- Command line
-
hf download hf://spaces/hanjian/GRN/grn/models/ema.py
-
curl -L -o ema.py https://huggingface.co/spaces/hanjian/GRN/resolve/main/grn/models/ema.py
716 Bytes
| import copy | |
| import torch | |
| from collections import OrderedDict | |
| def get_ema_model(model): | |
| ema_model = copy.deepcopy(model) | |
| ema_model.eval() | |
| for param in ema_model.parameters(): | |
| param.requires_grad = False | |
| return ema_model | |
| def update_ema(ema_model, model, decay): | |
| """ | |
| Step the EMA model towards the current model. | |
| """ | |
| ema_params = OrderedDict(ema_model.named_parameters()) | |
| model_params = OrderedDict(model.named_parameters()) | |
| for name, param in model_params.items(): | |
| # TODO: Consider applying only to params that require_grad to avoid small numerical changes of pos_embed | |
| ema_params[name].mul_(decay).add_(param.data, alpha=1 - decay) | |