Anonymous-123/ImageNet-Editing
1
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 