Понимание и самоанализ torch.autograd.backwardPython

Программы на Python
Anonymous
Понимание и самоанализ torch.autograd.backward

Сообщение Anonymous »

Чтобы найти ошибку, я пытаюсь проанализировать обратный расчет в PyTorch. Следуя описанию механики Autograd факела, я добавил обратные перехваты к каждому параметру моей модели, а также перехваты на grad_fn каждой активации. Следующий фрагмент кода показывает, как я добавляю перехватчики в grad_fn:

Код: Выделить всё

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)
При запуске моей модели я заметил, что каждый вызов перехватчика получает два аргумента и не имеет аргументов с ключевым словом. В случае функции AddBackward первый аргумент представляет собой список из двух тензоров, второй аргумент — список из одного тензора. То же самое справедливо и для функции LinearWithGradAccumulationAndAsyncCommunicationBackward. В случае функции MeanBackward оба аргумента представляют собой списки с одним тензором каждый.
Я предполагаю, что первый аргумент, вероятно, содержит входные данные для оператора (или чего-то еще). был сохранен с помощью ctx.save_for_backward), и что второй аргумент содержит градиенты. Прав ли я в этом? Могу ли я просто повторить обратное вычисление с помощью grad_fn(*args) или есть что-то еще (например, состояние)?
К сожалению, я не нашел никакой документации по этот. Я благодарен за любые ссылки на соответствующую документацию.

Подробнее здесь: https://stackoverflow.com/questions/787 ... d-backward

Вернуться в «Python»