Skip to content

Fix NNX axis metadata under Linen scan and vmap - #5590

Open
Junyi-Zheng wants to merge 1 commit into
google:mainfrom
Junyi-Zheng:fix/nnx-axis-metadata-5589
Open

Junyi-Zheng wants to merge 1 commit into
google:mainfrom
Junyi-Zheng:fix/nnx-axis-metadata-5589

Conversation

@Junyi-Zheng

Copy link
Copy Markdown

What does this PR do?

Fixes #5589.

NNXMeta.add_axis and remove_axis currently do nothing, so Linen scan and vmap can change variable shapes without updating axis metadata or invoking axis hooks. This can produce incorrect partition specifications and sharding errors.

This change:

  • Updates axis metadata through the NNX SPMD helpers without reinitializing or resharding values.
  • Handles empty metadata tuples and negative mapped axes.
  • Preserves the underlying variable type when converting HiJAX variables.
  • Adds regression tests for metadata, hooks, scan/vmap, sharding, JIT, and gradients.

Validation

  • All 72 bridge tests passed with JAX 0.11.1, including 36 new regression cases.
  • Full core tests ran on Python 3.12–3.14. Each run had one failure in test_tabulate_enum, also reproduced on the unmodified baseline.
  • All pre-commit hooks and library pytype checks passed. Full-library mypy errors were identical to the baseline.
  • Documentation builds, doctests, and built-wheel smoke tests passed.
  • Validation used macOS CPU, including four virtual CPU devices. GPU/TPU execution remains unverified.

Checklist

Update NNXMeta axis metadata through NNX SPMD helpers while preserving values, variable types, and hooks without reinitializing variables. Handle empty metadata tuples and negative mapped axes, and add regression coverage for scan, vmap, sharding, JIT, and gradients.

Fixes google#5589

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.

NNXMeta.add_axis / remove_axis are no-ops, silently dropping axis metadata under Linen scan / vmap

1 participant