(model_dest: nn.Module, model_src: nn.Module, rate)
| 11 | |
| 12 | |
| 13 | def ema(model_dest: nn.Module, model_src: nn.Module, rate): |
| 14 | param_dict_src = dict(model_src.named_parameters()) |
| 15 | for p_name, p_dest in model_dest.named_parameters(): |
| 16 | p_src = param_dict_src[p_name] |
| 17 | assert p_src is not p_dest |
| 18 | p_dest.data.mul_(rate).add_((1 - rate) * p_src.data) |
| 19 | |
| 20 | |
| 21 | class TrainState(object): |