f-dangel / f-dangel/backpack

Missing implementation of supported layers for DiagHessian and BatchDiagHessian

Open
#316 0 comments 0 reactions 0 assignees View on GitHub
Dominant language
Python
Stars
617
Forks
57
PR merge metrics
No merged PRs in 30d

Description

There are multiple layers which are specified as [being supported](https://docs.backpack.pt/en/master/supported-layers.html) for second order derivatives that actually do not work when trying to calculate the Hessian diagonal using `backpack-for-pytorch<=1.6.0`.

So far, I've run into this problem with the following layers:

- [ ] backpack.custom_module.branching.ScaleModule, torch.nn.Identity
- [ ] torch.nn.BatchNorm1d, torch.nn.BatchNorm2d, torch.nn.BatchNorm3d
- [ ] backpack.custom_module.branching.SumModule

This can be tested with a script such as the following:
```python
import torch
from backpack import backpack, extend
from backpack.extensions import DiagHessian, BatchDiagHessian
from backpack.custom_module.branching import Parallel, SumModule

model = extend(
torch.nn.Sequential(
*[
torch.nn.Conv2d(3, 16, kernel_size=(3, 3)),
Parallel(
torch.nn.Identity(), torch.nn.BatchNorm2d(16), merge_module=SumModule()
),
torch.nn.AdaptiveAvgPool2d(output_size=1),
torch.nn.Flatten(),
torch.nn.Linear(16, 2),
]
).cuda()
)
criterion = extend(torch.nn.CrossEntropyLoss())

batch = torch.randn((2, 3, 8, 8)).cuda()
target = torch.tensor([[1.0, 0.0], [0.0, 1.0]]).cuda()

model.eval()
model.zero_grad()
loss = criterion(model(batch), target)

with backpack(DiagHessian(), BatchDiagHessian()):
loss.backward()

hessian_diag = torch.cat(
[p.diag_h.view(-1) for p in model.parameters()], dim=1
)
hessian_diag_batch = torch.cat(
[p.diag_h_batch.view(batch.shape[0], -1) for p in model.parameters()], dim=1
)
```
I'm guessing that these require independent fixes, but think it is a good idea to collect all layers with missing support summarised here.

Contributor guide

No contributing guide indexed for this repository

Assessment

This issue has not been assessed yet.

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.