You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何优化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.

Fast, Loop-Free (or C-Loop) Solutions for Unique Neighbor Counts

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.15 03:35:13