Skip to content

Fix batch scalar broadcasting for block linear operators - #136

Open
AHMETHAKANBEZIR1 wants to merge 1 commit into
cornellius-gp:mainfrom
AHMETHAKANBEZIR1:fix/block-batch-scaling
Open

AHMETHAKANBEZIR1 wants to merge 1 commit into
cornellius-gp:mainfrom
AHMETHAKANBEZIR1:fix/block-batch-scaling

Conversation

@AHMETHAKANBEZIR1

Copy link
Copy Markdown

Fixes #135.

Problem

BlockLinearOperator._mul_constant forwards the visible batch constants to a base operator that has an extra internal block dimension. The constants therefore align with the blocks rather than the visible batch. Unequal batch/block counts can raise; equal counts silently return incorrect matrices.

Change

Insert a singleton dimension for the block axis when the constant is a Tensor. Each visible batch constant then applies to every block in that batch. The existing block representation and scalar path are retained.

The new unittest covers SumBatch, BlockDiag and BlockInterleaved operators with 30 subcases: float32/64, equal/unequal batch and block counts, one block, multiple batch dimensions and broadcast batch scalars. The reference assembles matrices directly with Tensor sum/einsum, independently of the block operators. It compares dense outputs, matrix multiplication and gradients with respect to blocks, scales and the RHS.

Validation

On fresh main 1c5e24bc774352461be818b9d347e93ddee7d7f7:

  • Before the fix: all 30 new subcases fail (14 assertion failures, 16 shape errors).
  • Python 3.12.14 / PyTorch 2.10.0+cpu: entire unittest suite: 5,047 tests passed; the five related operator test modules: 917 tests passed.
  • Python 3.14.6 / PyTorch 2.13.0+cpu: all 30 new forward/gradient subcases passed.
  • All 15 configured pre-commit hooks completed successfully or reported no applicable files, including full-repository flake8/ufmt/ASCII checks; git diff --check passed.
  • Full Sphinx documentation build with -W: passed.

All tests were run from this source checkout on Windows CPU. GPU/Inductor and other Torch/Python combinations were not run. No public API/docstring change. Existing sparse/indexing warnings in the full suite remain unchanged.

For local hook execution, an old installed require-ascii.py launcher used backslashes in its Windows shebang, which pre-commit misparsed. Only the cached launcher's shebang separators were normalized; its validation logic and the repository hook configuration were unchanged. No hook was bypassed.

#130 changes block slicing; this PR only changes batch scalar multiplication in a separate method.

AI disclosure

OpenAI Codex autonomously reproduced, implemented and tested this contribution at the account owner's request. No independent human code review is claimed. The commit includes Codex co-authorship.

Co-authored-by: Codex <noreply@openai.com>
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.

Block operator batch scaling broadcasts constants over the internal block axis

1 participant