From 48a0e7753d0b53560ebf0388f9e7b26cf497e111 Mon Sep 17 00:00:00 2001 From: AHMETHAKANBEZIR1 Date: Wed, 30 Sep 2026 20:20:05 +0300 Subject: [PATCH] Preserve visible batch dimensions when scaling block operators Co-authored-by: Codex --- .../operators/block_linear_operator.py | 3 + test/operators/test_block_linear_operator.py | 70 +++++++++++++++++++ 2 files changed, 73 insertions(+) create mode 100644 test/operators/test_block_linear_operator.py diff --git a/linear_operator/operators/block_linear_operator.py b/linear_operator/operators/block_linear_operator.py index 4a8e9e51..d5d32f12 100644 --- a/linear_operator/operators/block_linear_operator.py +++ b/linear_operator/operators/block_linear_operator.py @@ -156,6 +156,9 @@ def _mul_constant( # This preserves the block structure from linear_operator.operators.constant_mul_linear_operator import ConstantMulLinearOperator + if torch.is_tensor(other): + # Apply each batch constant to all blocks in that batch. + other = other.unsqueeze(-1) return self.__class__(ConstantMulLinearOperator(self.base_linear_op, other)) def _transpose_nonbatch( diff --git a/test/operators/test_block_linear_operator.py b/test/operators/test_block_linear_operator.py new file mode 100644 index 00000000..c360fd19 --- /dev/null +++ b/test/operators/test_block_linear_operator.py @@ -0,0 +1,70 @@ +#!/usr/bin/env python3 + +import unittest + +import torch + +from linear_operator.operators import ( + BlockDiagLinearOperator, + BlockInterleavedLinearOperator, + DenseLinearOperator, + SumBatchLinearOperator, +) + + +class TestBlockLinearOperatorBatchScaling(unittest.TestCase): + def test_batch_scaling_forward_and_gradients(self) -> None: + cases = [ + ((2,), 3, (2,)), + ((3,), 3, (3,)), + ((2,), 1, (2,)), + ((2, 3), 4, (2, 1)), + ((2, 3), 3, (1, 3)), + ] + for operator in (SumBatchLinearOperator, BlockDiagLinearOperator, BlockInterleavedLinearOperator): + for dtype in (torch.float32, torch.float64): + for batch_shape, n_blocks, scale_shape in cases: + with self.subTest( + operator=operator.__name__, dtype=dtype, case=(batch_shape, n_blocks, scale_shape) + ): + shape = (*batch_shape, n_blocks, 2, 2) + blocks = torch.arange(1, torch.Size(shape).numel() + 1, dtype=dtype).reshape(shape) + blocks.requires_grad_(True) + scales = torch.linspace(-2, 3, torch.Size(scale_shape).numel(), dtype=dtype).reshape( + scale_shape + ) + scales.requires_grad_(True) + linear_op = operator(DenseLinearOperator(blocks)) + + # Assemble the reference without using a block LinearOperator. + if operator is SumBatchLinearOperator: + dense = blocks.sum(-3) + else: + identity = torch.eye(n_blocks, dtype=dtype) + if operator is BlockDiagLinearOperator: + dense = torch.einsum("...bij,bc->...bicj", blocks, identity) + else: + dense = torch.einsum("...bij,bc->...ibjc", blocks, identity) + dense = dense.reshape(*batch_shape, 2 * n_blocks, 2 * n_blocks) + dense = dense * scales[..., None, None] + + scaled_op = linear_op * scales[..., None, None] + self.assertIsInstance(scaled_op, operator) + torch.testing.assert_close(scaled_op.to_dense(), dense) + + rhs = torch.linspace(-1, 1, dense.numel(), dtype=dtype).reshape(dense.shape) + rhs.requires_grad_(True) + actual = scaled_op @ rhs + expected = dense @ rhs + torch.testing.assert_close(actual, expected) + + weights = torch.linspace(1, 2, actual.numel(), dtype=dtype).reshape(actual.shape) + inputs = (blocks, scales, rhs) + actual_grads = torch.autograd.grad((actual * weights).sum(), inputs) + expected_grads = torch.autograd.grad((expected * weights).sum(), inputs) + for actual_grad, expected_grad in zip(actual_grads, expected_grads): + torch.testing.assert_close(actual_grad, expected_grad) + + +if __name__ == "__main__": + unittest.main()