3 Numba Tricks for Python Runtime Optimization

Numba compiles a numeric Python loop to machine code without ever having to leave your Python environment, rewrite anything in C, or vectorize some chunk of code that does not want to be vectorized. When Numba code disappoints, it’s nearly never the compiler. The offender usually ends up being the boundary around the compiled code: not crossing it, not making it wide enough, or crossing it during every run. Here are three tricks, all of them the same question asked three ways.
Note that everything below was checked against Numba 0.67.0.
pip install numba
Trick 1: Compiling the Loop Instead of Interpreting It
The baseline is a reduction over a NumPy array, and this is slow for the ordinary reason: the interpreter dispatches on types once per element, ten million times. The decorated version of the code differs by exactly one line. Numba reads the types on the first call to total_jit(), compiles a specialization for them, and every call after that it just runs native code:
import time
import numpy as np
from numba import njit
def total_plain(x):
total = 0.0
for i in range(x.shape[0]):
total += np.sqrt(x[i]) * np.sin(x[i])
return total
@njit
def total_jit(x):
total = 0.0
for i in range(x.shape[0]):
total += np.sqrt(x[i]) * np.sin(x[i])
return total
def total_numpy(x):
return np.sum(np.sqrt(x) * np.sin(x))
def benchmark(func, x, repeats=3):
"""Run func repeats times and return (best_time, mean_time, result)."""
times = []
result = None
for _ in range(repeats):
start = time.perf_counter()
result = func(x)
times.append(time.perf_counter() - start)
return min(times), sum(times) / len(times), result
x = np.random.default_rng(0).random(10_000_000)
# Measure JIT compilation separately (first call compiles)
start = time.perf_counter()
total_jit(x)
compile_time = time.perf_counter() - start
print(f"Numba first call (includes compilation): {compile_time:.4f} sn")
# Plain Python loop is slow on 10M elements, so run it only once
results = {
"Plain Python loop": benchmark(total_plain, x, repeats=1),
"Numba @njit": benchmark(total_jit, x, repeats=5),
"NumPy vectorized": benchmark(total_numpy, x, repeats=5),
}
baseline = results["Plain Python loop"][0]
print(f"{'Method':<20} {'Best (s)':>10} {'Mean (s)':>10} {'Speedup':>10} Result")
print("-" * 72)
for name, (best, mean, value) in results.items():
print(f"{name:<20} {best:>10.4f} {mean:>10.4f} {baseline / best:>9.1f}x {value:.6f}")
# Sanity check that all methods agree
values = [r[2] for r in results.values()]
print("nResults match:", np.allclose(values, values[0]))
Output:
Numba first call (includes compilation): 0.2799 s
Method Best (s) Mean (s) Speedup Result
------------------------------------------------------------------------
Plain Python loop 3.0789 3.0789 1.0x 3641603.675817
Numba @njit 0.0384 0.0385 80.2x 3641603.675817
NumPy vectorized 0.0552 0.0606 55.7x 3641603.675816
Results match: True
What hasn’t changed is the constraint underneath it. Nopython mode (@njit) “produces much faster code, but has limitations,” and those limitations are important. Using Numba well is mostly a matter of keeping the hot function inside the subset of Python and NumPy it can assign types to.
Trick 2: Spreading the Loop Across Every Core
The compiled function above still runs on one core. Adding parallel=True means that Numba now tries to parallelize the function, and swapping range() for prange() tells it which loop you mean. The body does not change at all. If we work this function into our script:
from numba import njit, prange
@njit(parallel=True)
def total_parallel(x):
total = 0.0
for i in prange(x.shape[0]):
total += np.sqrt(x[i]) * np.sin(x[i])
return total
And modify our results in order to run the new experiment as follows:
results = {
"Plain Python loop": benchmark(total_plain, x, repeats=1),
"Numba @njit": benchmark(total_jit, x, repeats=5),
"NumPy vectorized": benchmark(total_numpy, x, repeats=5),
"Parallel @njit": benchmark(total_parallel, x, repeats=5),
}
And here is our output:
Numba first call (includes compilation): 0.2431 s
Method Best (s) Mean (s) Speedup Result
------------------------------------------------------------------------
Plain Python loop 3.0229 3.0229 1.0x 3641603.675817
Numba @njit 0.0374 0.0379 80.7x 3641603.675817
NumPy vectorized 0.0547 0.0608 55.3x 3641603.675816
Parallel @njit 0.0087 0.1222 347.3x 3641603.675816
Results match: True
This is quite a dramatic increase in speedup.
The reason this is safe is that total += ... is a pattern Numba recognizes as a reduction. As a result, it splits the range across threads, gives each one a private accumulator, and finally combines them at the end. The same is true for -=, *=, /=, max and min.
Trick 3: Paying the Compile Cost Only Once
Compilation happens on the first call, so a fresh process pays for that computation again every time. For a script you run once a day that is negligible. However, for a tool you are running twenty times an hour it may consume the majority of the runtime. cache=True writes the compiled result for total_cached() to disk beside the source, and a later run loads it instead of recompiling. Let’s add this function to our script:
@njit(parallel=True, cache=True)
def total_cached(x):
total = 0.0
for i in prange(x.shape[0]):
total += np.sqrt(x[i]) * np.sin(x[i])
return total
Once again modify results to run the new experiment and report back to us once it has:
results = {
"Plain Python loop": benchmark(total_plain, x, repeats=1),
"Numba @njit": benchmark(total_jit, x, repeats=5),
"NumPy vectorized": benchmark(total_numpy, x, repeats=5),
"Parallel @njit": benchmark(total_parallel, x, repeats=5),
"Cached @njit": benchmark(total_cached, x, repeats=5),
}
And the output:
Numba first call (includes compilation): 0.1248 s
Method Best (s) Mean (s) Speedup Result
------------------------------------------------------------------------
Plain Python loop 3.0025 3.0025 1.0x 3641603.675817
Numba @njit 0.0384 0.0384 78.2x 3641603.675817
NumPy vectorized 0.0550 0.0563 54.5x 3641603.675816
Parallel @njit 0.0086 0.0383 349.1x 3641603.675816
Cached @njit 0.0086 0.0094 349.9x 3641603.675816
Results match: True
In this case, we see a modest (nearly imperceptible) speedup over the parallel @njit implementation.
A global variable the function reads is frozen at its compile-time value and won’t rebind on cache load. Cache invalidation also “fails to recognize changes in symbols defined in a different file,” meaning that editing a helper function elsewhere can leave you running the previously compiled code. And caching a parallel=True function has a rougher history than caching a plain one, so confirm the cache is really being hit before you rely on it.
Wrapping Up
Three decorators, one function body, one question underneath all of them. Is the work inside the compiled boundary, is the whole machine inside it, and are you “paying” to cross it more than once. Compile the loop, then widen it, then stop recompiling it.
Matthew Mayo (@mattmayo13) holds a master’s degree in computer science and a graduate diploma in data mining. As managing editor of KDnuggets & Statology, and contributing editor at Machine Learning Mastery, Matthew aims to make complex data science concepts accessible. His professional interests include natural language processing, language models, machine learning algorithms, and exploring emerging AI. He is driven by a mission to democratize knowledge in the data science community. Matthew has been coding since he was 6 years old.



