feat: scale a batched tensor by a vector of per-batch scalars - #20
Merged
Merged
Conversation
A plain numeric array multiplying (or dividing) a tensor is now read as batch-wise scalar weights: its shape broadcasts against the batch axes with numpy's right-alignment, and the coefficient axis is never matched against it. This makes a spectral filter — exp.(-vals) * vecs on the batched eigenbasis from spectrum — work as it does upstream. Closes #12
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Closes #12.
A plain numeric array multiplying or dividing a tensor is now read as batch-wise scalar weights:
weights * ω,ω * weightsandω / weightsall work. The weights' shape broadcasts against the tensor's batch axes with numpy's right-alignment (reusingcompatible_batches/align_batch_pair), and the coefficient axis is never matched against them. Weights may also expand the batch shape, so an unbatched form times a length-Bvector yields batch(B,)— the same rule as upstream's_broadcast_batch_scalars+_broadcast_coeffs_to_batch.This makes a spectral filter work as it does in Python:
Tests
The
@test_brokenintest/test_pyclasses.jl(product_operations→per-batch scalar vector scaling) is flipped to real assertions and extended to cover the reverse order, division, batch expansion, multi-dimensional batch broadcasting, and the error path when no batch axis can absorb the weights. The upstream test that motivated the issue,test_operators.py::test_spectrum_eigenvector_batch_scaling_with_numpy, is ported into the operators section.Full suite passes: 2381 tests, none broken.
Two deliberate divergences
The dispatch signature is
AbstractArray{<:Number}, not theAbstractVectorthe issue suggested. Multi-dimensional batch shapes need the wider signature and upstream'snp.broadcast_shapesaccepts any rank — but it does mean a matrix times a tensor is now a batch-broadcast scale rather than aMethodError, so Julia's usual linear-algebra reading ofA * xdoesn't apply here.Incompatible weights raise
AssertionError, where upstream raisesTypeError. That matches how+and the pointwise products already guard in this file, so it is consistent with the repo rather than with Python.Not covered
weights .* ω(broadcast syntax) still fails —AbstractTensorhas noBase.broadcastable. Upstream reaches it via__array_ufunc__, which this test file already declares out of scope. Filed separately as #19.