Код: Выделить всё
import torch.distributed as dist
def make_hook(grad_fn, note=None):
if grad_fn is not None and grad_fn.name is not None:
def hook(*args, **kwargs):
print(f"[{dist.get_rank()}] {grad_fn.name()} with {len(args)} args "
f"and {len(kwargs)} kwargs [{note or '/'}]")
return hook
else:
return None
def register_hooks_on_grads(grad_fn, make_hook_fn):
if not grad_fn:
return
hook = make_hook_fn(grad_fn)
if hook:
grad_fn.register_hook(hook)
for fn, _ in grad_fn.next_functions:
if not fn:
continue
var = getattr(fn, "variable", None)
if var is None:
register_hooks_on_grads(fn, make_hook_fn)
x = torch.zeros(15, requires_grad=True)
y = x.exp()
z = y.sum()
register_hooks_on_grads(z.grad_fn, make_hook)
Я предполагаю, что первый аргумент, вероятно, содержит входные данные для оператора (или чего-то еще). был сохранен с помощью ctx.save_for_backward), и что второй аргумент содержит градиенты. Прав ли я в этом? Могу ли я просто повторить обратное вычисление с помощью grad_fn(*args) или есть что-то еще (например, состояние)?
К сожалению, я не нашел никакой документации по этот. Я благодарен за любые ссылки на соответствующую документацию.
Подробнее здесь: https://stackoverflow.com/questions/787 ... d-backward