Skip to content

Allow auto-registration to fall back to function outputs when no loss tags are registered. - #423

Merged
copybara-service[bot] merged 1 commit into
mainfrom
test_982154613
Sep 18, 2026
Merged

copybara-service[bot] merged 1 commit into
mainfrom
test_982154613

Conversation

@copybara-service

@copybara-service copybara-service Bot commented Sep 16, 2026 •

Copy link
Copy Markdown

Allow auto-registration to fall back to function outputs when no loss tags are registered.

Currently, KFAC graph matching and curvature estimation rely on registered LossTags as anchor points for backward reachability pruning and for computing loss VJPs. However, certain estimation modes—specifically "fisher_empirical_direct" and "fisher_empirical_direct_synced"—do not require loss VJPs and only need layer tags to compute empirical Fisher blocks directly from gradients.

This change:

  1. Adds fallback_to_outputs_if_no_losses to auto_register_tags and make_jax_graph. When enabled and no LossTags are present, graph reachability pruning anchors to the primary function output (outvars[:1]), ignoring auxiliary outputs (aux).
  2. Updates tracer.py to allow tracing without LossTags when fallback_to_outputs_if_no_losses is specified in auto_registration_kwargs.
  3. Defers _compute_losses_vjp() in BlockDiagonalCurvature.update_curvature_matrix_estimate() to only the specific estimation modes that require it (fisher_gradients, fisher_empirical, curvature propagation, and exact modes), raising a descriptive ValueError if no loss tags are found for those modes.
  4. Adds unit tests in test_graph_matcher.py and test_estimator.py covering fallback auto-registration, curvature updates, unsupported mode guards, and end-to-end Optimizer execution.

@copybara-service
copybara-service Bot force-pushed the test_982154613 branch 2 times, most recently from c3a9381 to bae96ed Compare September 17, 2026 21:45
@copybara-service copybara-service Bot changed the title Allow KFAC auto-registration to fall back to function outputs when no loss tags are registered. Allow auto-registration to fall back to function outputs when no loss tags are registered. Sep 17, 2026
@copybara-service
copybara-service Bot force-pushed the test_982154613 branch 2 times, most recently from 640f033 to 45f96c2 Compare September 18, 2026 16:18
… tags are registered.

Currently, KFAC graph matching and curvature estimation rely on registered LossTags as anchor points for backward reachability pruning and for computing loss VJPs. However, certain estimation modes—specifically "fisher_empirical_direct" and "fisher_empirical_direct_synced"—do not require loss VJPs and only need layer tags to compute empirical Fisher blocks directly from gradients.

This change:
1. Adds `fallback_to_outputs_if_no_losses` to `auto_register_tags` and `make_jax_graph`. When enabled and no LossTags are present, graph reachability pruning anchors to the primary function output (`outvars[:1]`), ignoring auxiliary outputs (`aux`).
2. Updates `tracer.py` to allow tracing without LossTags when `fallback_to_outputs_if_no_losses` is specified in `auto_registration_kwargs`.
3. Defers `_compute_losses_vjp()` in `BlockDiagonalCurvature.update_curvature_matrix_estimate()` to only the specific estimation modes that require it (`fisher_gradients`, `fisher_empirical`, curvature propagation, and exact modes), raising a descriptive ValueError if no loss tags are found for those modes.
4. Adds unit tests in `test_graph_matcher.py` and `test_estimator.py` covering fallback auto-registration, curvature updates, unsupported mode guards, and end-to-end Optimizer execution.

PiperOrigin-RevId: 983909653
@copybara-service
copybara-service Bot merged commit 086619d into main Sep 18, 2026
8 checks passed
@copybara-service
copybara-service Bot deleted the test_982154613 branch September 18, 2026 16:38
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