We want to compute d = x + y (DAXPY simplified form) and compare:
- Serial version
- Parallel version using
joblib.Parallel
import numpy as np
from joblib import Parallel, delayed
import time
# DAXPY: d = x + y
def daxpy_serial(x, y):
return x + y
def daxpy_parallel(x, y, n_jobs=-1, chunk_size=1000000):
n = len(x)
chunks = [(x[i:i+chunk_size], y[i:i+chunk_size]) for i in range(0, n, chunk_size)]
results = Parallel(n_jobs=n_jobs, backend="threading")(
delayed(np.add)(xc, yc) for xc, yc in chunks
)
return np.concatenate(results)
N = 10**7
x = np.random.rand(N)
y = np.random.rand(N)
# Serial
t0 = time.time()
d_ser = daxpy_serial(x, y)
t1 = time.time()
# Parallel
t2 = time.time()
d_par = daxpy_parallel(x, y)
t3 = time.time()
print("Equal results?", np.allclose(d_ser, d_par))
print(f"Serial : {t1 - t0:.6f} s")
print(f"Parallel : {t3 - t2:.6f} s")Results :
Equal results? True
Serial : 0.0078 s
Parallel : 0.0193 s
The serial approach, which is already vectorized and uses optimized C/BLAS routines, is faster for basic operations than the parallel version, which adds task overhead.
Lets make the task harder to see difrance:
import numpy as np
import time
from joblib import Parallel, delayed, cpu_count
# Serial version with heavy computation
def daxpy_serial(x, y):
# artificial heavy task per element
return np.sin(x) + np.exp(y) + np.sqrt(x * y + 1)
# Parallel version with joblib
def daxpy_parallel(x, y, n_jobs=None, chunk_size=10**5):
if n_jobs is None:
n_jobs = cpu_count()
n = len(x)
chunks = [(x[i:i+chunk_size], y[i:i+chunk_size]) for i in range(0, n, chunk_size)]
results = Parallel(n_jobs=n_jobs, backend="threading")(
delayed(daxpy_serial)(xc, yc) for xc, yc in chunks
)
return np.concatenate(results)
if __name__ == "__main__":
N = 10**8 # larger to see speed difference
x = np.random.rand(N)
y = np.random.rand(N)
# Serial timing
t0 = time.time()
d_ser = daxpy_serial(x, y)
t1 = time.time()
# Parallel timing
t2 = time.time()
d_par = daxpy_parallel(x, y)
t3 = time.time()
print("Equal results?", np.allclose(d_ser, d_par))
print(f"Serial : {t1 - t0:.3f} s")
print(f"Parallel : {t3 - t2:.3f} s")results
Equal results? True Serial : 1.555 s Parallel : 0.378 s
if we make the task more complex we will see the difference even more than this. So for basic oparations the Python use the C/BLAS which is super optimized but when the task is more complex the parallel computation is faster.