@@ -80,21 +80,13 @@ def _compute_fwd_scale(norm, n, shape):
8080
8181def _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