CoolFace
Apppublic

Anonymous-123/ImageNet-Editing

sourceHugging Facecreativeml-openrail-mupdated 4y agoView on Hugging Face
1likes
fp16_util.py237 linesDownload Raw Back to guided_diffusion
1"""2Helpers to train with 16-bit precision.3"""4 5import numpy as np6import torch as th7import torch.nn as nn8from torch._utils import _flatten_dense_tensors, _unflatten_dense_tensors9 10from . import logger11 12INITIAL_LOG_LOSS_SCALE = 20.013 14 15def convert_module_to_f16(l):16    """17    Convert primitive modules to float16.18    """19    if isinstance(l, (nn.Conv1d, nn.Conv2d, nn.Conv3d)):20        l.weight.data = l.weight.data.half()21        if l.bias is not None:22            l.bias.data = l.bias.data.half()23 24 25def convert_module_to_f32(l):26    """27    Convert primitive modules to float32, undoing convert_module_to_f16().28    """29    if isinstance(l, (nn.Conv1d, nn.Conv2d, nn.Conv3d)):30        l.weight.data = l.weight.data.float()31        if l.bias is not None:32            l.bias.data = l.bias.data.float()33 34 35def make_master_params(param_groups_and_shapes):36    """37    Copy model parameters into a (differently-shaped) list of full-precision38    parameters.39    """40    master_params = []41    for param_group, shape in param_groups_and_shapes:42        master_param = nn.Parameter(43            _flatten_dense_tensors(44                [param.detach().float() for (_, param) in param_group]45            ).view(shape)46        )47        master_param.requires_grad = True48        master_params.append(master_param)49    return master_params50 51 52def model_grads_to_master_grads(param_groups_and_shapes, master_params):53    """54    Copy the gradients from the model parameters into the master parameters55    from make_master_params().56    """57    for master_param, (param_group, shape) in zip(58        master_params, param_groups_and_shapes59    ):60        master_param.grad = _flatten_dense_tensors(61            [param_grad_or_zeros(param) for (_, param) in param_group]62        ).view(shape)63 64 65def master_params_to_model_params(param_groups_and_shapes, master_params):66    """67    Copy the master parameter data back into the model parameters.68    """69    # Without copying to a list, if a generator is passed, this will70    # silently not copy any parameters.71    for master_param, (param_group, _) in zip(master_params, param_groups_and_shapes):72        for (_, param), unflat_master_param in zip(73            param_group, unflatten_master_params(param_group, master_param.view(-1))74        ):75            param.detach().copy_(unflat_master_param)76 77 78def unflatten_master_params(param_group, master_param):79    return _unflatten_dense_tensors(master_param, [param for (_, param) in param_group])80 81 82def get_param_groups_and_shapes(named_model_params):83    named_model_params = list(named_model_params)84    scalar_vector_named_params = (85        [(n, p) for (n, p) in named_model_params if p.ndim <= 1],86        (-1),87    )88    matrix_named_params = (89        [(n, p) for (n, p) in named_model_params if p.ndim > 1],90        (1, -1),91    )92    return [scalar_vector_named_params, matrix_named_params]93 94 95def master_params_to_state_dict(96    model, param_groups_and_shapes, master_params, use_fp1697):98    if use_fp16:99        state_dict = model.state_dict()100        for master_param, (param_group, _) in zip(101            master_params, param_groups_and_shapes102        ):103            for (name, _), unflat_master_param in zip(104                param_group, unflatten_master_params(param_group, master_param.view(-1))105            ):106                assert name in state_dict107                state_dict[name] = unflat_master_param108    else:109        state_dict = model.state_dict()110        for i, (name, _value) in enumerate(model.named_parameters()):111            assert name in state_dict112            state_dict[name] = master_params[i]113    return state_dict114 115 116def state_dict_to_master_params(model, state_dict, use_fp16):117    if use_fp16:118        named_model_params = [119            (name, state_dict[name]) for name, _ in model.named_parameters()120        ]121        param_groups_and_shapes = get_param_groups_and_shapes(named_model_params)122        master_params = make_master_params(param_groups_and_shapes)123    else:124        master_params = [state_dict[name] for name, _ in model.named_parameters()]125    return master_params126 127 128def zero_master_grads(master_params):129    for param in master_params:130        param.grad = None131 132 133def zero_grad(model_params):134    for param in model_params:135        # Taken from https://pytorch.org/docs/stable/_modules/torch/optim/optimizer.html#Optimizer.add_param_group136        if param.grad is not None:137            param.grad.detach_()138            param.grad.zero_()139 140 141def param_grad_or_zeros(param):142    if param.grad is not None:143        return param.grad.data.detach()144    else:145        return th.zeros_like(param)146 147 148class MixedPrecisionTrainer:149    def __init__(150        self,151        *,152        model,153        use_fp16=False,154        fp16_scale_growth=1e-3,155        initial_lg_loss_scale=INITIAL_LOG_LOSS_SCALE,156    ):157        self.model = model158        self.use_fp16 = use_fp16159        self.fp16_scale_growth = fp16_scale_growth160 161        self.model_params = list(self.model.parameters())162        self.master_params = self.model_params163        self.param_groups_and_shapes = None164        self.lg_loss_scale = initial_lg_loss_scale165 166        if self.use_fp16:167            self.param_groups_and_shapes = get_param_groups_and_shapes(168                self.model.named_parameters()169            )170            self.master_params = make_master_params(self.param_groups_and_shapes)171            self.model.convert_to_fp16()172 173    def zero_grad(self):174        zero_grad(self.model_params)175 176    def backward(self, loss: th.Tensor):177        if self.use_fp16:178            loss_scale = 2 ** self.lg_loss_scale179            (loss * loss_scale).backward()180        else:181            loss.backward()182 183    def optimize(self, opt: th.optim.Optimizer):184        if self.use_fp16:185            return self._optimize_fp16(opt)186        else:187            return self._optimize_normal(opt)188 189    def _optimize_fp16(self, opt: th.optim.Optimizer):190        logger.logkv_mean("lg_loss_scale", self.lg_loss_scale)191        model_grads_to_master_grads(self.param_groups_and_shapes, self.master_params)192        grad_norm, param_norm = self._compute_norms(grad_scale=2 ** self.lg_loss_scale)193        if check_overflow(grad_norm):194            self.lg_loss_scale -= 1195            logger.log(f"Found NaN, decreased lg_loss_scale to {self.lg_loss_scale}")196            zero_master_grads(self.master_params)197            return False198 199        logger.logkv_mean("grad_norm", grad_norm)200        logger.logkv_mean("param_norm", param_norm)201 202        self.master_params[0].grad.mul_(1.0 / (2 ** self.lg_loss_scale))203        opt.step()204        zero_master_grads(self.master_params)205        master_params_to_model_params(self.param_groups_and_shapes, self.master_params)206        self.lg_loss_scale += self.fp16_scale_growth207        return True208 209    def _optimize_normal(self, opt: th.optim.Optimizer):210        grad_norm, param_norm = self._compute_norms()211        logger.logkv_mean("grad_norm", grad_norm)212        logger.logkv_mean("param_norm", param_norm)213        opt.step()214        return True215 216    def _compute_norms(self, grad_scale=1.0):217        grad_norm = 0.0218        param_norm = 0.0219        for p in self.master_params:220            with th.no_grad():221                param_norm += th.norm(p, p=2, dtype=th.float32).item() ** 2222                if p.grad is not None:223                    grad_norm += th.norm(p.grad, p=2, dtype=th.float32).item() ** 2224        return np.sqrt(grad_norm) / grad_scale, np.sqrt(param_norm)225 226    def master_params_to_state_dict(self, master_params):227        return master_params_to_state_dict(228            self.model, self.param_groups_and_shapes, master_params, self.use_fp16229        )230 231    def state_dict_to_master_params(self, state_dict):232        return state_dict_to_master_params(self.model, state_dict, self.use_fp16)233 234 235def check_overflow(value):236    return (value == float("inf")) or (value == -float("inf")) or (value != value)237