| 16 | return False |
| 17 | |
| 18 | class LossScaler(object): |
| 19 | def __init__(self, scale=1.0, dynamic=False): |
| 20 | self._dynamic = dynamic |
| 21 | self._loss_scale = 2.**16 if self._dynamic else scale |
| 22 | self._max_loss_scale = 2.**24 |
| 23 | self._scale_seq_len = 2000 |
| 24 | self._unskipped = 0 |
| 25 | self._has_overflow = False |
| 26 | |
| 27 | @property |
| 28 | def loss_scale(self): |
| 29 | return self._loss_scale |
| 30 | |
| 31 | @property |
| 32 | def has_overflow(self): |
| 33 | return self._has_overflow |
| 34 | |
| 35 | def unscale_and_update(self, param_groups, scale): |
| 36 | if not self._dynamic: |
| 37 | for p in iter_params(param_groups): |
| 38 | if p.grad is not None: |
| 39 | p.grad.data.mul_(1. / scale) |
| 40 | return |
| 41 | |
| 42 | self._has_overflow = False |
| 43 | for p in iter_params(param_groups): |
| 44 | if p.grad is not None: |
| 45 | self._has_overflow = scale_check_overflow(p.grad.data, |
| 46 | 1. / scale) |
| 47 | if self._has_overflow: |
| 48 | break |
| 49 | |
| 50 | # if self._overflow_buf.any(): |
| 51 | if self._has_overflow: |
| 52 | should_skip = True |
| 53 | self._loss_scale /= 2. |
| 54 | self._unskipped = 0 |
| 55 | else: |
| 56 | should_skip = False |
| 57 | self._unskipped += 1 |
| 58 | |
| 59 | if self._unskipped == self._scale_seq_len: |
| 60 | self._loss_scale = min(self._max_loss_scale, self._loss_scale * 2.) |
| 61 | self._unskipped = 0 |
| 62 | |
| 63 | return should_skip |
| 64 | |
| 65 | def backward(self, loss): |
| 66 | scaled_loss = loss*self.loss_scale |
| 67 | scaled_loss.backward() |