| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -706,10 +706,13 @@ def einsum_path(*operands, **kwargs): | |||
| 706 | 706 | for cnum, char in enumerate(term): | |
| 707 | 707 | dim = sh[cnum] | |
| 708 | 708 | 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)) | ||
| 713 | 716 | else: | |
| 714 | 717 | dimension_dict[char] = dim | |
| 715 | 718 | ||
@@ -1101,13 +1104,22 @@ def einsum(*operands, **kwargs): | |||
| 1101 | 1104 | if specified_out and ((num + 1) == len(contraction_list)): | |
| 1102 | 1105 | handle_out = True | |
| 1103 | 1106 | ||
| 1104 | - # Call tensordot | ||
| 1107 | + # Handle broadcasting vs BLAS cases | ||
| 1105 | 1108 | if blas: | |
| 1106 | - | ||
| 1107 | 1109 | # Checks have already been handled | |
| 1108 | 1110 | input_str, results_index = einsum_str.split('->') | |
| 1109 | 1111 | 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: | ||
| 1111 | 1123 | tensor_result = input_left + input_right | |
| 1112 | 1124 | for s in idx_rm: | |
| 1113 | 1125 | tensor_result = tensor_result.replace(s, "") | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -481,6 +481,25 @@ def check_einsum_sums(self, dtype, do_opt=False): | |||
| 481 | 481 | r = np.arange(4).reshape(2, 2) + 7 | |
| 482 | 482 | assert_equal(np.einsum('z,mz,zm->', p, q, r), 253) | |
| 483 | 483 | ||
| 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 | + | ||
| 484 | 503 | def test_einsum_sums_int8(self): | |
| 485 | 504 | self.check_einsum_sums('i1') | |
| 486 | 505 | ||
| Back | FazBrowse Home | New Git URL |
0 commit comments