Skip to content

Support Conv1d and Conv3d weights in Adahessian trace - #547

Open
shaneraphel wants to merge 1 commit into
jettify:masterfrom
shaneraphel:adahessian-conv1d-conv3d
Open

shaneraphel wants to merge 1 commit into
jettify:masterfrom
shaneraphel:adahessian-conv1d-conv3d

Conversation

@shaneraphel

@shaneraphel shaneraphel commented Sep 25, 2026 •

Copy link
Copy Markdown

Summary

Adahessian.get_trace averaged the Hessian diagonal over the spatial dims only when the weight was 4D. A Conv1d (3D) or Conv3d (5D) weight left tmp_output unbound and step() raised UnboundLocalError.

The trace now averages hv.abs() over dims 2..ndim-1 with keepdim=True. For 4D that is exactly dims [2, 3], so Conv2d is unchanged.

Fixes #500.

Test plan

  • pytest tests/test_adahessian_conv.py: Conv1d, Conv2d, and Conv3d each take a step; the 4D generic mean equals the old explicit dims
  • pytest tests/test_basic.py -k Adahessian: 3 passed

Prepared with an AI assistant. I reviewed the diff and ran the tests on CPU.

get_trace averaged the Hessian diagonal over the spatial dims only for
4D kernels. Conv1d (3D) and Conv3d (5D) weights left tmp_output
unbound. Average over dims 2..ndim-1 instead; 4D is unchanged.

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

adahessian bug: it does not support the 3D conv

1 participant