Download clean/video/fakestormer/package_utils/misc.py from deepsafe/model-code: direct link, hf CLI and curl.
- Browser
- Download file 2.24 kB
-
https://huggingface.co/deepsafe/model-code/resolve/main/clean/video/fakestormer/package_utils/misc.py
- Command line
-
hf download hf://deepsafe/model-code/clean/video/fakestormer/package_utils/misc.py
-
curl -L -o misc.py https://huggingface.co/deepsafe/model-code/resolve/main/clean/video/fakestormer/package_utils/misc.py
2.24 kB
| # -*- coding:utf-8 -*- | |
| import torch | |
| from torch._six import inf | |
| class NativeScalerWithGradNormCount: | |
| state_dict_key = "amp_scaler" | |
| def __init__(self): | |
| self._scaler = torch.cuda.amp.GradScaler() | |
| def __call__( | |
| self, | |
| cfg, | |
| loss, | |
| optimizer, | |
| clip_grad=None, | |
| parameters=None, | |
| create_graph=False, | |
| update_grad=True, | |
| step=0, | |
| ): | |
| self._scaler.scale(loss).backward(create_graph=create_graph) | |
| if update_grad: | |
| if clip_grad is not None: | |
| assert parameters is not None | |
| self._scaler.unscale_( | |
| optimizer | |
| ) # unscale the gradients of optimizer's assigned params in-place | |
| norm = torch.nn.utils.clip_grad_norm_(parameters, clip_grad) | |
| else: | |
| self._scaler.unscale_(optimizer) | |
| norm = get_grad_norm_(parameters) | |
| if cfg.TRAIN.optimizer != "SAM": | |
| self._scaler.step(optimizer) | |
| else: | |
| if step == 0: | |
| optimizer.first_step(zero_grad=True) | |
| else: | |
| self._scaler = optimizer.second_step( | |
| zero_grad=False, scaler=self._scaler | |
| ) | |
| self._scaler.update() | |
| else: | |
| norm = None | |
| return norm | |
| def state_dict(self): | |
| return self._scaler.state_dict() | |
| def load_state_dict(self, state_dict): | |
| self._scaler.load_state_dict(state_dict) | |
| def get_grad_norm_(parameters, norm_type: float = 2.0) -> torch.Tensor: | |
| if isinstance(parameters, torch.Tensor): | |
| parameters = [parameters] | |
| parameters = [p for p in parameters if p.grad is not None] | |
| norm_type = float(norm_type) | |
| if len(parameters) == 0: | |
| return torch.tensor(0.0) | |
| device = parameters[0].grad.device | |
| if norm_type == inf: | |
| total_norm = max(p.grad.detach().abs().max().to(device) for p in parameters) | |
| else: | |
| total_norm = torch.norm( | |
| torch.stack( | |
| [torch.norm(p.grad.detach(), norm_type).to(device) for p in parameters] | |
| ), | |
| norm_type, | |
| ) | |
| return total_norm | |