如何优化Numpy二维数组中唯一邻域元素计数的实现?
Hey there! Great question—ditching those slow Python loops for vectorized NumPy operations is always a smart move, especially when you're working with larger arrays. Let's walk through how to solve this without explicit loops, and why your initial np.unique attempt didn't quite work.
First, let's clarify the np.unique confusion: the axis parameter works across entire rows/columns/slices of a high-dimensional array, but we need to compute unique values per individual position's neighborhood—so it's not a straightforward fit. Instead, we need to collect all neighbors for each position first, then count unique values (excluding the element itself) for each spot.
Option 1: C-Backed Loop (Way Faster Than Python Loops)
This approach uses np.apply_along_axis, which runs loops in C under the hood (far faster than pure Python loops) while keeping code clean and readable. It works for any label type (not just integers) and doesn't require knowing your label set upfront.
import numpy as np def count_unique_neighbors(labels, neighborhood='4'): # Pad the array to handle edge cases (use 'constant' instead of 'edge' if you want 0s for out-of-bounds) padded = np.pad(labels, pad_width=1, mode='edge') rows, cols = labels.shape # Define neighborhood shifts if neighborhood == '4': shifts = [(-1, 0), (1, 0), (0, -1), (0, 1)] # Up, Down, Left, Right elif neighborhood == '8': shifts = [(-1, -1), (-1, 0), (-1, 1), (0, -1), (0, 1), (1, -1), (1, 0), (1, 1)] # All 8 directions else: raise ValueError("Neighborhood must be '4' or '8'") # Collect all neighbor arrays (each matches the original array's shape) neighbors = [] for dr, dc in shifts: neighbor = padded[1+dr : rows+1+dr, 1+dc : cols+1+dc] neighbors.append(neighbor) # Stack into a 3D array: (num_neighbors, rows, cols) neighbors_stack = np.stack(neighbors, axis=0) # Helper function to count unique neighbors (excluding current value) def _count_unique(arr): current_val = arr[-1] # Filter out current value, then count unique entries filtered = arr[:-1][arr[:-1] != current_val] return len(np.unique(filtered)) # Combine neighbors with the original value (for the helper function) combined = np.concatenate([neighbors_stack, labels[np.newaxis, :, :]], axis=0) # Apply the helper to every position (C-backed loop, fast!) return np.apply_along_axis(_count_unique, axis=0, arr=combined)
Option 2: Fully Vectorized (No Python/C Loops)
If you want to avoid loops entirely, this method uses boolean broadcasting and label mapping. It's blazing fast when your label set is small, but uses more memory if you have hundreds/thousands of unique labels.
import numpy as np def count_unique_neighbors_vectorized(labels, neighborhood='4'): padded = np.pad(labels, pad_width=1, mode='edge') rows, cols = labels.shape # Define neighborhood shifts if neighborhood == '4': shifts = [(-1, 0), (1, 0), (0, -1), (0, 1)] elif neighborhood == '8': shifts = [(-1, -1), (-1, 0), (-1, 1), (0, -1), (0, 1), (1, -1), (1, 0), (1, 1)] else: raise ValueError("Neighborhood must be '4' or '8'") # Collect all neighbor arrays neighbors = [padded[1+dr:rows+1+dr, 1+dc:cols+1+dc] for dr, dc in shifts] # Map labels to integer indices (for broadcasting) all_labels = np.unique(labels) label_to_idx = {lbl: idx for idx, lbl in enumerate(all_labels)} labels_idx = np.vectorize(label_to_idx.get)(labels) # Track which labels are present in each position's neighborhood num_labels = len(all_labels) present = np.zeros((num_labels, rows, cols), dtype=bool) for neighbor in neighbors: neighbor_idx = np.vectorize(label_to_idx.get)(neighbor) # Mark positions where each label appears in the neighbor for idx in range(num_labels): present[idx] |= (neighbor_idx == idx) # Count valid unique neighbors (exclude the current label) current_idx_expanded = labels_idx[np.newaxis, :, :] valid = present & (np.arange(num_labels)[:, np.newaxis, np.newaxis] != current_idx_expanded) return np.sum(valid, axis=0)
Test It Out!
Let's use your example input to verify:
# Example 6x7 input array labels = np.array([ [1, 1, 2, 2, 2, 3, 3], [1, 1, 2, 2, 3, 3, 3], [1, 2, 2, 3, 3, 3, 4], [2, 2, 3, 3, 4, 4, 4], [2, 3, 3, 4, 4, 5, 5], [3, 3, 4, 4, 5, 5, 5] ]) # 4-neighborhood counts counts_4 = count_unique_neighbors(labels, '4') print("4-Neighborhood Unique Counts:\n", counts_4) # 8-neighborhood counts counts_8 = count_unique_neighbors_vectorized(labels, '8') print("\n8-Neighborhood Unique Counts:\n", counts_8)
Which One Should You Use?
- Option 1 is great for most cases: it's flexible, easy to tweak, and fast enough for large arrays. It works with any label type (strings, floats, etc.).
- Option 2 is ideal if you have a small number of unique labels and need maximum speed—just be aware of the memory overhead if your label set grows large.
内容的提问来源于stack exchange,提问作者ikkjo

