# this replaces default optimizer.zero_grad() and verifies reuse # of same gradient tensors def zero_grads(model): for p in model.parameters(): if p.grad is not None: p.grad.zero_() # uses dynamo utility to mark each of the gradient tensors as static def mark_static_address(optimizer): for group in optimizer.param_groups: for p in group['params']: if p.grad is not None: torch._dynamo.mark_static_address(p.grad)