Метод register_hook
Метод register_hook класса Tensor позволяет зарегистрировать функцию-хук,
которая будет вызвана при вычислении градиента для данного тензора.
Хук получает градиент в качестве аргумента и может модифицировать его или выполнять
дополнительные действия, такие как логирование или отладка.
Метод возвращает объект-дескриптор, который можно использовать для удаления хука.
Синтаксис
tensor.register_hook(hook_fn)
Параметр hook_fn - это функция, принимающая один аргумент (градиент)
и возвращающая либо изменённый градиент, либо None.
Пример
Давайте создадим простой граф вычислений и зарегистрируем хук для тензора, чтобы отследить значение градиента:
import torch
t = torch.tensor([2.0, 3.0], requires_grad=True)
res = t * 2
res = res.sum()
def hook_fn(grad):
print(f"Gradient: {grad}")
return grad
t.register_hook(hook_fn)
res.backward()
Результат выполнения кода:
Gradient: tensor([2., 2.])
Пример
Хук может модифицировать градиент, умножая его на константу или выполняя другие преобразования. В этом примере мы удвоим градиент:
import torch
t = torch.tensor([1.0, 2.0, 3.0], requires_grad=True)
res = t ** 2
res = res.sum()
def modify_grad(grad):
return grad * 2
t.register_hook(modify_grad)
res.backward()
print(t.grad)
Результат выполнения кода:
tensor([4., 8., 12.])
Пример
Хуки полезны для отладки, позволяя проверять градиенты промежуточных тензоров в глубоких сетях. Зарегистрируем хук для тензора в модели:
import torch
model = torch.nn.Linear(2, 1)
x = torch.tensor([[1.0, 2.0]], requires_grad=True)
y = model(x)
loss = y.sum()
def debug_hook(grad):
print(f"Intermediate gradient: {grad}")
return grad
x.register_hook(debug_hook)
loss.backward()
Результат выполнения кода (примерный вывод):
Intermediate gradient: tensor([[0.0389, 0.0541]])
Пример
Хук можно удалить, используя возвращаемый дескриптор. Вызовем метод
remove дескриптора, чтобы отключить хук:
import torch
t = torch.tensor([1.0, 2.0], requires_grad=True)
res = t * 3
res = res.sum()
def simple_hook(grad):
print("Hook called!")
return grad
hook_handle = t.register_hook(simple_hook)
hook_handle.remove()
res.backward()
print("Backward complete")
Результат выполнения кода:
Backward complete
Хук не был вызван, так как мы удалили его до выполнения обратного распространения.
Пример
Хуки можно использовать для накопления статистики по градиентам, например, для вычисления среднего или нормы градиента:
import torch
grad_norms = []
def collect_norm(grad):
norm = grad.norm().item()
grad_norms.append(norm)
print(f"Gradient norm: {norm:.4f}")
return grad
t = torch.tensor([1.0, 2.0, 3.0, 4.0], requires_grad=True)
res = t * t
res = res.sum()
t.register_hook(collect_norm)
res.backward()
print(f"All norms: {grad_norms}")
Результат выполнения кода:
Gradient norm: 5.4772
All norms: [5.477225575051661]
Смотрите также
-
метод
backward,
который выполняет обратное распространение ошибки -
атрибут
grad,
содержащий вычисленный градиент тензора -
атрибут
requires_grad,
указывающий, нужны ли градиенты для тензора -
метод
retain_grad,
сохраняющий градиент для промежуточного тензора