如何用Numpy高效计算含最小值的DTW递推公式?
Great question! Your nested Python loops are the bottleneck here—Python-level loops are notoriously slow for large matrices like (2000,2000). Let's convert that recurrence into vectorized NumPy operations to get massive speedups, just like your simpler slicing example.
First, Recap the Recurrence
We need to compute:
g[i,j] = d[i,j] + min(g[i-1,j], g[i-1,j-1], g[i-1,j-2])
The key challenge is handling boundary cases (when j=0 or j=1, some of the previous indices go out of bounds). Your original code uses nan to ignore invalid entries, which we can replicate efficiently with padding and inf (since min will ignore inf just like nanmin ignores nan).
Vectorized Implementation
Instead of looping over each j, we can compute entire rows at once using NumPy slicing and padding. Here's how:
import numpy as np # Assume g is pre-initialized with your starting row g[0,:], and d is your distance matrix g_rows, g_cols = g.shape for i in range(1, g_rows): # Get the three previous row slices, padded to handle boundaries # 1. g[i-1,j] (same column, no padding needed) prev_same = g[i-1, :] # 2. g[i-1,j-1] (shift right by 1, pad left with inf for j=0) prev_left1 = np.pad(g[i-1, :-1], (1, 0), mode='constant', constant_values=np.inf) # 3. g[i-1,j-2] (shift right by 2, pad left with inf for j=0 and j=1) prev_left2 = np.pad(g[i-1, :-2], (2, 0), mode='constant', constant_values=np.inf) # Compute the minimum of the three candidates for every j min_prev = np.minimum(np.minimum(prev_same, prev_left1), prev_left2) # Update the entire row in one go g[i, :] = d[i, :] + min_prev
Why This Is Faster
- Vectorized Operations: NumPy executes these slicing and min operations in optimized C code, not slow Python loops. For a 2000x2000 matrix, this should reduce runtime from minutes to seconds.
- Boundary Handling: Padding with
infensures invalid entries (likeg[i-1,-1]) don't affect the minimum, just like your originalnanminlogic.
Extra Optimization Tip
If you want even more speed, you could use Numba to JIT-compile the loop (it plays nicely with NumPy operations). But the vectorized approach above should already give you a huge improvement without adding extra dependencies.
内容的提问来源于stack exchange,提问作者jkjk

