FazBrowse GitHub Viewer | Trending |
URL:
| Home
Tools: [Download Repo ZIP]   [Original HTTPS Page]

Merge pull request #10559 from charris/backport-10352 · numpy/numpy@3f82455 · GitHub

/ numpy Public

Commit 3f82455

Browse files
authored
Merge pull request #10559 from charris/backport-10352
BUG: Fix einsum optimize logic for singleton dimensions
2 parents a0f117d + ab3e91c commit 3f82455

2 files changed

Lines changed: 38 additions & 7 deletions

File tree

‎numpy/core/einsumfunc.py‎

Lines changed: 19 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -706,10 +706,13 @@ def einsum_path(*operands, **kwargs):
706706
for cnum, char in enumerate(term):
707707
dim = sh[cnum]
708708
if char in dimension_dict.keys():
709-
if dimension_dict[char] != dim:
710-
raise ValueError("Size of label '%s' for operand %d does "
711-
"not match previous terms."
712-
% (char, tnum))
709+
# For broadcasting cases we always want the largest dim size
710+
if dimension_dict[char] == 1:
711+
dimension_dict[char] = dim
712+
elif dim not in (1, dimension_dict[char]):
713+
raise ValueError("Size of label '%s' for operand %d (%d) "
714+
"does not match previous terms (%d)."
715+
% (char, tnum, dimension_dict[char], dim))
713716
else:
714717
dimension_dict[char] = dim
715718

@@ -1101,13 +1104,22 @@ def einsum(*operands, **kwargs):
11011104
if specified_out and ((num + 1) == len(contraction_list)):
11021105
handle_out = True
11031106

1104-
# Call tensordot
1107+
# Handle broadcasting vs BLAS cases
11051108
if blas:
1106-
11071109
# Checks have already been handled
11081110
input_str, results_index = einsum_str.split('->')
11091111
input_left, input_right = input_str.split(',')
1110-
1112+
if 1 in tmp_operands[0] or 1 in tmp_operands[1]:
1113+
left_dims = {dim: size for dim, size in
1114+
zip(input_left, tmp_operands[0].shape)}
1115+
right_dims = {dim: size for dim, size in
1116+
zip(input_right, tmp_operands[1].shape)}
1117+
# If dims do not match we are broadcasting, BLAS off
1118+
if any(left_dims[ind] != right_dims[ind] for ind in idx_rm):
1119+
blas = False
1120+
1121+
# Call tensordot if still possible
1122+
if blas:
11111123
tensor_result = input_left + input_right
11121124
for s in idx_rm:
11131125
tensor_result = tensor_result.replace(s, "")

‎numpy/core/tests/test_einsum.py‎

Lines changed: 19 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -481,6 +481,25 @@ def check_einsum_sums(self, dtype, do_opt=False):
481481
r = np.arange(4).reshape(2, 2) + 7
482482
assert_equal(np.einsum('z,mz,zm->', p, q, r), 253)
483483

484+
# singleton dimensions broadcast (gh-10343)
485+
p = np.ones((10,2))
486+
q = np.ones((1,2))
487+
assert_array_equal(np.einsum('ij,ij->j', p, q, optimize=True),
488+
np.einsum('ij,ij->j', p, q, optimize=False))
489+
assert_array_equal(np.einsum('ij,ij->j', p, q, optimize=True),
490+
[10.] * 2)
491+
492+
p = np.ones((1, 5))
493+
q = np.ones((5, 5))
494+
for optimize in (True, False):
495+
assert_array_equal(np.einsum("...ij,...jk->...ik", p, p,
496+
optimize=optimize),
497+
np.einsum("...ij,...jk->...ik", p, q,
498+
optimize=optimize))
499+
assert_array_equal(np.einsum("...ij,...jk->...ik", p, q,
500+
optimize=optimize),
501+
np.full((1, 5), 5))
502+
484503
def test_einsum_sums_int8(self):
485504
self.check_einsum_sums('i1')
486505

0 commit comments

Comments
 (0)

Back | FazBrowse Home | New Git URL