| 141 | ) |
| 142 | |
| 143 | class FullModelGradientClippingOptimizer(optim): |
| 144 | def step(self, closure=None): |
| 145 | all_params = itertools.chain(*[x["params"] for x in self.param_groups]) |
| 146 | torch.nn.utils.clip_grad_norm_(all_params, clip_norm_val) |
| 147 | super().step(closure=closure) |
| 148 | |
| 149 | return FullModelGradientClippingOptimizer if enable else optim |
| 150 |
nothing calls this directly
no outgoing calls
no test coverage detected