MCPcopy Create free account
hub / github.com/DPS2022/diffusion-posterior-sampling / MixedPrecisionTrainer

Class MixedPrecisionTrainer

guided_diffusion/fp16_util.py:146–230  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

144
145
146class MixedPrecisionTrainer:
147 def __init__(
148 self,
149 *,
150 model,
151 use_fp16=False,
152 fp16_scale_growth=1e-3,
153 initial_lg_loss_scale=INITIAL_LOG_LOSS_SCALE,
154 ):
155 self.model = model
156 self.use_fp16 = use_fp16
157 self.fp16_scale_growth = fp16_scale_growth
158
159 self.model_params = list(self.model.parameters())
160 self.master_params = self.model_params
161 self.param_groups_and_shapes = None
162 self.lg_loss_scale = initial_lg_loss_scale
163
164 if self.use_fp16:
165 self.param_groups_and_shapes = get_param_groups_and_shapes(
166 self.model.named_parameters()
167 )
168 self.master_params = make_master_params(self.param_groups_and_shapes)
169 self.model.convert_to_fp16()
170
171 def zero_grad(self):
172 zero_grad(self.model_params)
173
174 def backward(self, loss: th.Tensor):
175 if self.use_fp16:
176 loss_scale = 2 ** self.lg_loss_scale
177 (loss * loss_scale).backward()
178 else:
179 loss.backward()
180
181 def optimize(self, opt: th.optim.Optimizer):
182 if self.use_fp16:
183 return self._optimize_fp16(opt)
184 else:
185 return self._optimize_normal(opt)
186
187 def _optimize_fp16(self, opt: th.optim.Optimizer):
188 logger.logkv_mean("lg_loss_scale", self.lg_loss_scale)
189 model_grads_to_master_grads(self.param_groups_and_shapes, self.master_params)
190 grad_norm, param_norm = self._compute_norms(grad_scale=2 ** self.lg_loss_scale)
191 if check_overflow(grad_norm):
192 self.lg_loss_scale -= 1
193 logger.log(f"Found NaN, decreased lg_loss_scale to {self.lg_loss_scale}")
194 zero_master_grads(self.master_params)
195 return False
196
197 logger.logkv_mean("grad_norm", grad_norm)
198 logger.logkv_mean("param_norm", param_norm)
199
200 self.master_params[0].grad.mul_(1.0 / (2 ** self.lg_loss_scale))
201 opt.step()
202 zero_master_grads(self.master_params)
203 master_params_to_model_params(self.param_groups_and_shapes, self.master_params)

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected