(self)
| 169 | copy(current_buffers.data, ma_buffers.data) |
| 170 | |
| 171 | def get_current_decay(self): |
| 172 | epoch = (self.step - self.update_after_step - 1).clamp(min = 0.) |
| 173 | value = 1 - (1 + epoch / self.inv_gamma) ** - self.power |
| 174 | |
| 175 | if epoch.item() <= 0: |
| 176 | return 0. |
| 177 | |
| 178 | return value.clamp(min = self.min_value, max = self.beta).item() |
| 179 | |
| 180 | def update(self): |
| 181 | step = self.step.item() |
no outgoing calls
no test coverage detected