Skip to content

Fix sparse Jacobians with empty dependencies and singleton batch axes - #48

Open
Giuseppe Carleo (gcarleo) wants to merge 1 commit into
microsoft:mainfrom
gcarleo:fix/sparse-empty-and-batch-axes
Open

Giuseppe Carleo (gcarleo) wants to merge 1 commit into
microsoft:mainfrom
gcarleo:fix/sparse-empty-and-batch-axes

Conversation

@gcarleo

@gcarleo Giuseppe Carleo (gcarleo) commented Sep 21, 2026 •

Copy link
Copy Markdown
Contributor

Sparse propagation can fail for batched attention and complex Slater wavefunctions:

  • find_out_idx removes singleton position axes, but remove_zero_entries indexes the original operands. Restore the broadcast position frame before selecting nonzero entries; a size-one batch ahead of a multi-head contraction otherwise produces out-of-bounds mask indices.
  • A sparse Jacobian may have no rows. Treat its largest dependency index as -1, allowing the existing zero-Jacobian densification path instead of reducing an empty array without an identity.

Regression tests compare multi-head attention values, gradients, and Laplacians against JAX autodiff for unbatched, singleton, nested-singleton, and two-item batches. The empty-mask test now also covers zero rows.

Validation with JAX 0.10.2 and two CPU devices:

JAX_NUM_CPU_DEVICES=2 python -m pytest -o addopts='' -q test --ignore=test/experimental
# 148 passed, 328 subtests passed
pre-commit run --all-files

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