Skip to content

Commit badf839

Browse files
committed
chore: improve
1 parent 27f7037 commit badf839

1 file changed

Lines changed: 10 additions & 5 deletions

File tree

‎mkl_fft/_fft_utils.py‎

Lines changed: 10 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -23,6 +23,8 @@
2323
# OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
2424
# OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
2525

26+
import math
27+
2628
import numpy as np
2729

2830
# pylint: disable=no-name-in-module
@@ -70,13 +72,16 @@ def _compute_fwd_scale(norm, n, shape):
7072
return 1.0
7173

7274
ss = n if n is not None else shape
73-
# The 1-D callers in `_mkl_fft.py` always pass a plain scalar here
74-
# (either `n` or `x.shape[axis]`), so avoid the overhead of wrapping
75-
# it in a 0-d numpy array via `np.prod` on that hot path. Sequences
76-
# (used by the N-D callers via `_nd_fwd_scale`) still go through
77-
# `np.prod`, unchanged.
75+
# `np.prod` dominates the Python-side cost of a small normalized transform,
76+
# so take cheaper routes for the two shapes the callers actually pass: a
77+
# scalar `n` (1-D) and a sequence (`_nd_fwd_scale`). It stays the fallback
78+
# because `numpy.fft` also accepts array-like `n` and `s` (e.g. a 0-d
79+
# `n=np.array(8)`, a 1-D `s=np.array([4, 4])`), which `math.prod` cannot
80+
# handle uniformly.
7881
if isinstance(ss, (int, np.integer)):
7982
nn = ss
83+
elif isinstance(ss, (list, tuple)):
84+
nn = math.prod(ss)
8085
else:
8186
nn = np.prod(ss)
8287
fsc = 1 / nn if nn != 0 else 1

0 commit comments

Comments
 (0)