Skip to content

Estimate the Adahessian trace on both complex components - #556

Open
shaneraphel wants to merge 1 commit into
jettify:masterfrom
shaneraphel:adahessian-complex-trace
Open

shaneraphel wants to merge 1 commit into
jettify:masterfrom
shaneraphel:adahessian-complex-trace

Conversation

@shaneraphel

@shaneraphel shaneraphel commented Sep 25, 2026 •

Copy link
Copy Markdown

Summary

The Adahessian half of #458. torch.randint_like rejects a complex parameter, so get_trace raised RuntimeError: check_random_bounds handles only integral, floating-point and boolean types and the step never ran. AdaBound's complex failure in the same issue is a different call and is handled separately.

The Hutchinson vector is now a Rademacher draw on view_as_real of the parameter, packed back with view_as_complex so it still matches the complex gradient that autograd produced. The absolute value that removes the ±1 sign is taken per component. A complex conv kernel still averages its spatial dimensions, as the real 4D path does; the component axis is not one of those dimensions. Moments are stored on the two components and the write-back goes through view_as_real, so neither part is dropped.

A real parameter does not take this branch. Two steps on [0.5, -0.2, 0.3] still finish at [0.43060994148254395, -0.17224396765232086, 0.2583659589290619].

Complex rank 3 and 5 are still rejected. The real implementation has no branch for those ranks either, and a complex conv1d or conv3d would need the same spatial reduction the real kernels are still missing.

Test plan

  • pytest tests/test_adahessian_complex.py
  • Real path matches the previous two-step result.
  • One step on [0.5+0.1j, -0.2+0.3j] matches Adahessian on view_as_real of the same values, bitwise, including the random signs.
  • A 2×2 complex matrix takes a step and stays complex64.

This change was produced with assistance from an AI coding tool. I reproduced the randint error, checked the complex step against the real algorithm on both components, and ran the tests locally.

randint_like rejects a complex parameter, so Adahessian raised
RuntimeError before the Hutchinson vector existed (issue 458).
Signs are drawn on the real and imaginary parts and packed back
into a complex vector. The trace and the moments stay on those
two components, and a real parameter still follows the old update.

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.

1 participant