From 5d3eeeff54c1371f852e2cedb65c4898a6975fa5 Mon Sep 17 00:00:00 2001 From: rohith500 Date: Thu, 17 Sep 2026 13:13:05 -0400 Subject: [PATCH] fix(pruning): validate multidimensional tensor in channel-structured pruning (#107) --- changelog.d/107.fixed | 1 + src/coreai_opt/pruning/spec/prune.py | 14 +++++++++++--- tests/pruning/test_magnitude_pruner.py | 19 ++++++++++++++++++- 3 files changed, 30 insertions(+), 4 deletions(-) create mode 100644 changelog.d/107.fixed diff --git a/changelog.d/107.fixed b/changelog.d/107.fixed new file mode 100644 index 00000000..169b5a68 --- /dev/null +++ b/changelog.d/107.fixed @@ -0,0 +1 @@ +Validate that channel-structured magnitude pruning is applied to tensors with at least 2 dimensions, raising a descriptive ValueError for 1D tensors diff --git a/src/coreai_opt/pruning/spec/prune.py b/src/coreai_opt/pruning/spec/prune.py index 7333fdad..d04751e0 100644 --- a/src/coreai_opt/pruning/spec/prune.py +++ b/src/coreai_opt/pruning/spec/prune.py @@ -139,14 +139,15 @@ def compute_mask( Returns: torch.Tensor: Binary mask (1 = keep, 0 = prune). """ + # TODO: Replace this with generic abstractions + if isinstance(pruning_scheme, ChannelStructured): + return _MagnitudePruneImpl._compute_channel_mask(weight, sparsity, pruning_scheme.axis) + if sparsity == 0.0: return torch.ones_like(weight) if sparsity >= 1.0: return torch.zeros_like(weight) - # TODO: Replace this with generic abstractions - if isinstance(pruning_scheme, ChannelStructured): - return _MagnitudePruneImpl._compute_channel_mask(weight, sparsity, pruning_scheme.axis) return _MagnitudePruneImpl._compute_unstructured_mask(weight, sparsity) @staticmethod @@ -171,6 +172,13 @@ def _compute_channel_mask( Channel importance is measured by L1 norm. The least-important channels are pruned entirely. """ + if weight.ndim < 2: + raise ValueError( + f"Channel-structured pruning requires a tensor with at least 2 dimensions, " + f"got shape {tuple(weight.shape)} with {weight.ndim} dims. " + f"For 1D tensors, use Unstructured pruning instead." + ) + if not (-weight.ndim <= axis < weight.ndim): raise ValueError( f"Invalid axis. Should be in range [{-weight.ndim}, {weight.ndim}), but got {axis}" diff --git a/tests/pruning/test_magnitude_pruner.py b/tests/pruning/test_magnitude_pruner.py index b3b60e3b..1b54e5ca 100644 --- a/tests/pruning/test_magnitude_pruner.py +++ b/tests/pruning/test_magnitude_pruner.py @@ -23,7 +23,12 @@ OpMagnitudePrunerConfig, PolynomialDecaySchedule, ) -from coreai_opt.pruning.spec import ChannelStructured, PruneImplBase, Unstructured +from coreai_opt.pruning.spec import ( + ChannelStructured, + PruneImplBase, + Unstructured, + _MagnitudePruneImpl, +) @pytest.fixture @@ -462,6 +467,18 @@ def test_channel_structured_axis_out_of_range(self, axis: int) -> None: with pytest.raises(ValueError, match="Invalid axis"): pruner.prepare((torch.randn(1, 4),)) + @pytest.mark.parametrize("target_sparsity", [0.0, 0.5]) + @pytest.mark.parametrize("axis", [0, -1], ids=["axis-0", "axis-neg-1"]) + def test_channel_structured_1d_tensor_raises(self, axis: int, target_sparsity: float) -> None: + """Channel-structured pruning requires >= 2 dimensions; 1D tensors raise ValueError.""" + weight = torch.tensor([1.0, 5.0, 2.0, 8.0]) + scheme = ChannelStructured(axis=axis) + with pytest.raises( + ValueError, + match=r"Channel-structured pruning requires a tensor with at least 2 dimensions", + ): + _MagnitudePruneImpl.compute_mask(weight, target_sparsity, scheme) + def test_linear_unstructured_conv2d_channel_structured(self) -> None: """Apply unstructured to Linear and channel-structured to Conv2d in same model."""