Skip to content

Commit 56d2a70

Browse files
vkverma9534charris
authored andcommitted
BUG: validate contraction axes in tensordot (numpy#30521)
1 parent 2e501ca commit 56d2a70

2 files changed

Lines changed: 20 additions & 1 deletion

File tree

‎numpy/_core/numeric.py‎

Lines changed: 14 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1022,7 +1022,8 @@ def tensordot(a, b, axes=2):
10221022
* (2,) array_like
10231023
Or, a list of axes to be summed over, first sequence applying to `a`,
10241024
second to `b`. Both elements array_like must be of the same length.
1025-
1025+
Each axis may appear at most once; repeated axes are not allowed.
1026+
For example, ``axes=([1, 1], [0, 0])`` is invalid.
10261027
Returns
10271028
-------
10281029
output : ndarray
@@ -1053,6 +1054,13 @@ def tensordot(a, b, axes=2):
10531054
first in both sequences, the second axis second, and so forth.
10541055
The calculation can be referred to ``numpy.einsum``.
10551056
1057+
For example, if ``a.shape == (2, 3, 4)`` and ``b.shape == (3, 4, 5)``,
1058+
then ``axes=([1, 2], [0, 1])`` sums over the ``(3, 4)`` dimensions of
1059+
both arrays and produces an output of shape ``(2, 5)``.
1060+
1061+
Each summation axis corresponds to a distinct contraction index; repeating
1062+
an axis (for example ``axes=([1, 1], [0, 0])``) is invalid.
1063+
10561064
The shape of the result consists of the non-contracted axes of the
10571065
first tensor, followed by the non-contracted axes of the second.
10581066
@@ -1170,6 +1178,11 @@ def tensordot(a, b, axes=2):
11701178
axes_b = [axes_b]
11711179
nb = 1
11721180

1181+
if len(set(axes_a)) != len(axes_a):
1182+
raise ValueError("duplicate axes are not allowed in tensordot")
1183+
if len(set(axes_b)) != len(axes_b):
1184+
raise ValueError("duplicate axes are not allowed in tensordot")
1185+
11731186
a, b = asarray(a), asarray(b)
11741187
as_ = a.shape
11751188
nda = a.ndim

‎numpy/_core/tests/test_numeric.py‎

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -4254,6 +4254,12 @@ def test_raise(self):
42544254

42554255
class TestTensordot:
42564256

4257+
def test_rejects_duplicate_axes(self):
4258+
a = np.ones((2, 3, 3))
4259+
b = np.ones((3, 3, 4))
4260+
with pytest.raises(ValueError):
4261+
np.tensordot(a, b, axes=([1, 1], [0, 0]))
4262+
42574263
def test_zero_dimension(self):
42584264
# Test resolution to issue #5663
42594265
a = np.ndarray((3, 0))

0 commit comments

Comments
 (0)