Skip to content

Commit e600825

Browse files
committed
fix: review
1 parent bb74339 commit e600825

1 file changed

Lines changed: 7 additions & 17 deletions

File tree

‎mkl_fft/_fft_utils.py‎

Lines changed: 7 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -80,21 +80,13 @@ def _compute_fwd_scale(norm, n, shape):
8080

8181
def _compute_nd_scale_shape(x, s, axes, norm=None, invreal=False):
8282
"""
83-
Resolve the lengths that a norm-scaled N-D transform normalizes over.
83+
Lengths a norm-scaled N-D transform normalizes over.
8484
8585
``_compute_fwd_scale`` falls back to the full array shape when ``s`` is
86-
None. That over-normalizes when only a subset of axes is transformed, and
87-
for c2r transforms the basis is the *output* length along the last
88-
transformed axis rather than the input one.
89-
90-
This mirrors what the ``numpy_fft`` and ``scipy_fft`` interfaces already
91-
do by calling ``_cook_nd_args`` before delegating. Only the scale basis is
92-
resolved here; ``s`` itself is left alone so that dispatch in
93-
``_c2c_fftnd_impl`` is unchanged.
94-
95-
``norm`` is accepted only to skip the work for the unscaled norms, whose
96-
scale is 1.0 regardless of shape. Invalid values fall through to
97-
``_compute_fwd_scale``, which validates them.
86+
None, which over-normalizes a subset-of-axes transform; for c2r the basis
87+
is the output length ``2 * (n - 1)``. Mirrors what the interfaces already
88+
do via ``_cook_nd_args``, but leaves ``s`` alone so dispatch is unchanged.
89+
``norm`` is only used to skip the work when the scale is 1.0 anyway.
9890
"""
9991

10092
if s is not None or norm in (None, "backward"):
@@ -104,17 +96,15 @@ def _compute_nd_scale_shape(x, s, axes, norm=None, invreal=False):
10496
ss = list(x.shape)
10597
last = len(ss) - 1
10698
elif len(axes) == 0:
107-
# transforming no axes is an identity, so the scale is 1.0;
108-
# np.prod(()) is 1, which gives that for every norm
99+
# identity transform; np.prod(()) == 1 gives scale 1.0
109100
return ()
110101
else:
111102
ss = [x.shape[ai] for ai in axes]
112103
last = axes[-1]
113104
if invreal:
114105
ss[-1] = 2 * (x.shape[last] - 1)
115106
except (IndexError, TypeError):
116-
# invalid axes; leave the scale alone and let the transform itself
117-
# raise, so the error matches what NumPy reports
107+
# invalid axes; let the transform raise, matching NumPy's error
118108
return s
119109
return tuple(ss)
120110

0 commit comments

Comments
 (0)