如何加速自适应中值滤波(Adaptive Median Filter)的Python代码?
Your original code is slow for large images primarily because it uses nested per-pixel loops, and for each pixel, it repeatedly extracts windows and computes min/median/max from scratch—operations that scale poorly with image size and window size. Here are two optimized approaches, focusing on vectorization and precomputation to drastically speed things up:
Key Issues in Original Code
- Per-pixel window extraction: For every pixel, you slice the image to get the window and compute stats (min/median/max), which is O(k²) per pixel where k is window size.
- Redundant computations: The same window stats are recalculated multiple times as you increment window size for pixels that don't meet the initial condition.
- Off-by-one loop error: Your outer loops run
range(a, H+a+1)which causes window slices to go out of bounds when using the maximum window sizesMax.
Optimization Strategy
The core idea is to precompute window statistics (min, median, max) for all required window sizes upfront using NumPy's vectorized sliding window operations. Then, use boolean masks to apply the adaptive median logic without per-pixel loops.
Optimized Vectorized Implementation
This version eliminates per-pixel loops entirely by leveraging precomputed stats and vectorized boolean operations:
import numpy as np from numpy.lib.stride_tricks import sliding_window_view def AdaptiveMedianFilterVectorized(img, initial_s=3, sMax=7): if len(img.shape) == 3: raise Exception("Single channel image only") H, W = img.shape pad_size = sMax // 2 # Pad the image with zeros (use np.pad for efficiency) padded_img = np.pad(img, pad_width=pad_size, mode='constant', constant_values=0) # Precompute min, median, max for all odd window sizes from initial_s to sMax window_sizes = list(range(initial_s, sMax + 1, 2)) window_stats = {} for sz in window_sizes: # Create sliding windows of size sz x sz windows = sliding_window_view(padded_img, window_shape=(sz, sz)) # Compute stats across window axes zmin = np.min(windows, axis=(-2, -1)) zmed = np.median(windows, axis=(-2, -1)) zmax = np.max(windows, axis=(-2, -1)) window_stats[sz] = (zmin, zmed, zmax) # Initialize output and unprocessed mask output = np.zeros_like(img) unprocessed = np.ones((H, W), dtype=bool) # Original pixel values (aligned with output) zxy = padded_img[pad_size:H+pad_size, pad_size:W+pad_size] for sz in window_sizes: if not np.any(unprocessed): break # Exit early if all pixels are processed half_sz = sz // 2 # Calculate offset to align precomputed stats with output pixels offset = pad_size - half_sz # Extract stats for current window size, aligned to output zmin = window_stats[sz][0][offset:offset+H, offset:offset+W] zmed = window_stats[sz][1][offset:offset+H, offset:offset+W] zmax = window_stats[sz][2][offset:offset+H, offset:offset+W] # Apply only to unprocessed pixels mask = unprocessed # Level A condition: Zmed is not min or max of the window cond_level_A = (zmed - zmin > 0) & (zmed - zmax < 0) cond_level_A = cond_level_A & mask if np.any(cond_level_A): # Level B condition: Current pixel is not min or max of the window cond_level_B = (zxy - zmin > 0) & (zxy - zmax < 0) cond_level_B = cond_level_B & cond_level_A # Assign values based on Level B result output[cond_level_B] = zxy[cond_level_B] output[cond_level_A & ~cond_level_B] = zmed[cond_level_A & ~cond_level_B] # Mark these pixels as processed unprocessed[cond_level_A] = False # For the largest window size, assign median to remaining unprocessed pixels if sz == sMax and np.any(unprocessed): output[unprocessed] = zmed[unprocessed] return output
Alternative: Precomputed Stats with Minimal Loops
If you prefer a more incremental approach (closer to your original code but faster), this version precomputes stats and uses per-pixel loops with O(1) lookups instead of window computations:
import numpy as np from numpy.lib.stride_tricks import sliding_window_view def AdaptiveMedianFilterPrecomputed(img, initial_s=3, sMax=7): if len(img.shape) == 3: raise Exception("Single channel image only") H, W = img.shape pad_size = sMax // 2 padded_img = np.pad(img, pad_width=pad_size, mode='constant', constant_values=0) # Precompute stats for all window sizes window_sizes = list(range(initial_s, sMax + 1, 2)) window_stats = {} for sz in window_sizes: windows = sliding_window_view(padded_img, window_shape=(sz, sz)) window_stats[sz] = ( np.min(windows, axis=(-2, -1)), np.median(windows, axis=(-2, -1)), np.max(windows, axis=(-2, -1)) ) output = np.zeros_like(img) # Iterate over each pixel in the original image (aligned to padded image) for i in range(H): for j in range(W): padded_x = i + pad_size padded_y = j + pad_size current_s = initial_s final_val = None while current_s <= sMax: half_sz = current_s // 2 # Get precomputed stats for this window size at the current pixel stat_idx_x = padded_x - half_sz stat_idx_y = padded_y - half_sz zmin, zmed, zmax = window_stats[current_s] zmin = zmin[stat_idx_x, stat_idx_y] zmed = zmed[stat_idx_x, stat_idx_y] zmax = zmax[stat_idx_x, stat_idx_y] # Level A check if (zmed - zmin > 0) and (zmed - zmax < 0): # Level B check zxy = padded_img[padded_x, padded_y] if (zxy - zmin > 0) and (zxy - zmax < 0): final_val = zxy else: final_val = zmed break else: current_s += 2 # If we reached max window size, use its median if final_val is None: half_sz = sMax // 2 stat_idx_x = padded_x - half_sz stat_idx_y = padded_y - half_sz final_val = window_stats[sMax][1][stat_idx_x, stat_idx_y] output[i, j] = final_val return output
Why This Is Faster
- Precomputation: We compute window stats once per window size instead of per pixel, reducing redundant calculations.
- Vectorization: The vectorized version uses NumPy's optimized array operations instead of Python loops, which are orders of magnitude faster for large datasets.
- Eliminated Window Slicing: No repeated slicing of the image for each pixel—stats are looked up from precomputed arrays.
Performance Notes
- The median computation is the most expensive part of precomputation. For very large images or larger
sMax, consider usingscipy.ndimage.median_filter(from SciPy) which is optimized for speed, though it requires installing SciPy. - Ensure your input image is of a numeric dtype (e.g.,
uint8,float32) to avoid unnecessary type conversions during computations.
内容的提问来源于stack exchange,提问作者Mario

