Repository navigation
Expand file tree
/
Copy pathfft_bench.py
More file actions
241 lines (212 loc) · 10.4 KB
/
Copy pathfft_bench.py
File metadata and controls
241 lines (212 loc) · 10.4 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
# Copyright (c) 2017-2025 Intel Corporation.
#
# SPDX-License-Identifier: MIT
import argparse
import contextlib
import importlib
import inspect
import numpy as np
import scipy.fft
import os
import perf
import re
# Mark which FFT submodules are available...
fft_modules = {'numpy.fft': np.fft, 'scipy.fft': scipy.fft}
def valid_shape(shape_str):
shape = re.sub(r'[^\d]+', 'x', shape_str).strip('x').split('x')
shape = tuple(int(i) for i in shape)
if len(shape) < 1 or any(i < 1 for i in shape):
raise argparse.ArgumentTypeError(f'parsed shape {shape} has '
'non-positive entries or less than '
'one dimension.')
return shape
def valid_dtype(dtype_str):
dtype = np.dtype(dtype_str)
if dtype.kind not in ('f', 'c'):
raise argparse.ArgumentTypeError('only complex or real floating-point '
'data-types are allowed')
return dtype
# Parse args
parser = argparse.ArgumentParser(description='Benchmark FFT using NumPy and '
'SciPy.')
fft_group = parser.add_argument_group(title='FFT problem arguments')
fft_group.add_argument('-t', '--threads', '--num-threads', '--core-number',
type=int, default=perf.set_threads(no_guessing=True)[0],
help='Number of threads to use for FFT computation. '
'%(prog)s will attempt to use mkl-service to get/set '
'number of threads globally, and will also try to '
'set number of workers in scipy.fft. (default in this '
'environment: %(default)d)')
fft_group.add_argument('-m', '--modules', '--submodules', nargs='*',
default=tuple(fft_modules.keys()),
choices=tuple(fft_modules.keys()),
help='Use FFT functions from MODULES. (default: '
'%(default)s)')
fft_group.add_argument('-d', '--dtype', default=np.dtype('complex128'),
type=valid_dtype,
help='use DTYPE as the FFT domain. DTYPE must be '
'specified such that it is parsable by numpy.dtype. '
'(default: %(default)s)')
fft_group.add_argument('-r', '--rfft', default=False, action='store_true',
help='do not copy superfluous harmonics when FFT '
'output is conjugate-even, i.e. for real inputs.')
fft_group.add_argument('-P', '--overwrite-x', '--in-place', default=False,
action='store_true', help='Allow overwriting the input '
'buffer with the FFT outputs')
fft_group.add_argument('-s', '--seed', default=7777, type=int,
help='Seed for random number generator')
fft_group.add_argument('--scipy-backend', default='stock',
choices=('stock', 'mkl'),
help='Which uarray backend to use for scipy.fft. '
'"stock" = default pocketfft. "mkl" = register '
'mkl_fft.interfaces.scipy_fft via scipy.fft.set_backend '
'around the timed region. Ignored for numpy.fft. '
'(default: %(default)s)')
fft_group.add_argument('--numpy-backend', default='stock',
choices=('stock', 'mkl'),
help='Which implementation to use for numpy.fft. '
'"stock" = numpy.fft as installed. "mkl" = route '
'numpy.fft through mkl_fft via '
'mkl_fft.patch_numpy_fft(). Ignored for scipy.fft. '
'(default: %(default)s)')
timing_group = parser.add_argument_group(title='Timing arguments')
timing_group.add_argument('-i', '--inner-loops', '--batch-size',
type=int, default=16, metavar='IL',
help='time the benchmark IL times for each printed '
'measurement. Copying is not timed. (default: '
'%(default)s)')
timing_group.add_argument('-o', '--outer-loops', '--samples', '--repetitions',
type=int, default=24, metavar='OL',
help='print OL measurements. (default: %(default)s)')
output_group = parser.add_argument_group(title='Output arguments')
output_group.add_argument('-p', '--prefix', default='python',
help='Output PREFIX as the first value in outputs '
'(default: %(default)s)')
output_group.add_argument('-H', '--no-header', default=True,
action='store_false', dest='header',
help='do not output CSV header. This can be useful '
'if running multiple benchmarks back to back.')
output_group.add_argument('-v', '--verbose', default=False,
action='store_true',
help='Print excessive debug messages')
parser.add_argument('shape', type=valid_shape,
help='FFT shape to run, specified as a tuple of positive '
'decimal integers, delimited by any non-digit characters. '
'For example, both (101, 203, 305) and 101x203x305 denote '
'the same 3D FFT.')
args = parser.parse_args()
# Resolve optional mkl_fft scipy uarray backend. Doing this once up front
# lets us fail fast with a clear message if the user asked for "mkl" but
# the package is missing, rather than silently falling through to pocketfft
# the way a plain scipy.fft.set_backend call would.
_mkl_scipy_backend = None
if args.scipy_backend == 'mkl':
try:
import mkl_fft.interfaces.scipy_fft as _mkl_scipy_backend
except ImportError as e:
parser.error(f'--scipy-backend mkl requested but mkl_fft is not '
f'importable in this environment: {e}')
# Route numpy.fft through mkl_fft unless the installed NumPy already does,
# failing fast if routing is not possible.
if args.numpy_backend == 'mkl':
try:
import mkl_fft
except ImportError as e:
parser.error(f'--numpy-backend mkl requested but mkl_fft is not '
f'importable in this environment: {e}')
if not np.fft.fft.__module__.startswith('mkl_fft'):
if not hasattr(mkl_fft, 'patch_numpy_fft'):
parser.error(f'--numpy-backend mkl requested but mkl_fft '
f'{mkl_fft.__version__} has no patch_numpy_fft()')
mkl_fft.patch_numpy_fft()
if not np.fft.fft.__module__.startswith('mkl_fft'):
parser.error(f'--numpy-backend mkl requested but numpy.fft.fft is '
f'still {np.fft.fft.__module__} after patching')
# Print environment info (conda env, MKL version)
perf.print_environment_info()
# Get timer
timer = perf.get_timer()
if args.verbose:
print(f'TAG: timer = {timer.name}')
# Set threads
threads, threading_info_source = perf.set_threads(num_threads=args.threads,
verbose=args.verbose)
if args.verbose:
print(f'TAG: threading_info_source = {threading_info_source}')
# Get function from shape
assert len(args.shape) >= 1
func_name = {1: 'fft', 2: 'fft2'}.get(len(args.shape), 'fftn')
if args.rfft:
func_name = 'r' + func_name
if args.rfft and args.dtype.kind == 'c':
parser.error('--rfft makes no sense for an FFT of complex inputs. The '
'FFT output will not be conjugate even, so the whole output '
'matrix must be computed!')
# Generate input data
rs, rs_name = perf.get_generator_and_name(seed=args.seed)
if args.verbose:
print(f'TAG: random = {rs_name}')
arr = rs.standard_normal(args.shape)
if args.dtype.kind == 'c':
arr = arr + rs.standard_normal(args.shape) * 1j
arr = np.asarray(arr, dtype=args.dtype)
if args.verbose:
print(f'TAG:{perf.arg_signature(arr)}')
# Print header
print("", flush=True)
if args.header:
print('prefix,module,function,threads,dtype,size,place,time', flush=True)
# Run benchmarks. One for each selected module
for mod_name in args.modules:
# Determine arguments to benchmark and get function
mod = fft_modules[mod_name]
func = getattr(mod, func_name)
kwargs = {}
time_kwargs = dict(timer=timer, batch_size=args.inner_loops,
repetitions=args.outer_loops,
refresh_buffer=False, verbose=args.verbose)
in_place = False
actual_threads = threads
# Inspect function to see if it allows overwrite_x.
# For example, numpy.fft functions do not accept overwrite_x.
sig = inspect.signature(func)
if 'overwrite_x' in sig.parameters:
in_place = kwargs['overwrite_x'] = args.overwrite_x
time_kwargs['refresh_buffer'] = in_place
else:
# Skip this if we needed overwrite_x but didn't get it
if args.overwrite_x:
continue
if 'workers' in sig.parameters:
actual_threads = kwargs['workers'] = args.threads
# Route scipy.fft through mkl_fft's uarray backend when requested.
# numpy.fft has no uarray dispatch, so the context is a no-op there.
if mod_name == 'scipy.fft' and _mkl_scipy_backend is not None:
backend_ctx = scipy.fft.set_backend(_mkl_scipy_backend, only=True)
effective_backend = 'mkl'
else:
backend_ctx = contextlib.nullcontext()
effective_backend = 'stock' if mod_name == 'scipy.fft' else 'n/a'
if args.verbose:
print(f'TAG: scipy_backend = {effective_backend}')
with backend_ctx:
# threads warm-up, inside the backend context so the warmup path
# matches the timed path exactly (same dispatcher, same planner).
buf = np.empty_like(arr)
np.copyto(buf, arr)
x1 = func(buf)
del x1
del buf
perf_times = perf.time_func(func, arr, kwargs, **time_kwargs)
# Tag the prefix with the effective backend so CSV rows from the
# two scipy passes (stock vs mkl in the same env) stay distinguishable.
if mod_name == 'scipy.fft':
row_prefix = f'{args.prefix}-scipy-{effective_backend}'
elif args.numpy_backend == 'mkl':
row_prefix = f'{args.prefix}-numpy-mkl'
else:
row_prefix = args.prefix
for t in perf_times:
print(f'{row_prefix},{mod_name},{func_name},{actual_threads},'
f'{arr.dtype.name},{"x".join(str(i) for i in args.shape)},'
f'{"in-place" if in_place else "out-of-place"},{t:.5g}')