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

如何优化NumPy中二维数组的五子连珠检测性能?

五子连珠检测优化需求

需要实现一个函数,检测尺寸≥5的二维方阵(实际常用7x7)中是否存在某一数字的横向、纵向、斜向五子连珠模式(至少检测到一次即可)。例如如下矩阵中存在3处1的五子连珠:

import numpy as np
A = np.array(
[
    [0, 1, 0, 0, 0, 0],
    [1, 0, 1, 0, 1, 0],
    [1, 0, 0, 1, 0, 0],
    [1, 0, 1 ,0, 1, 0],
    [1, 1, 0, 0, 0, 1],
    [1, 0, 0, 0, 0, 0]
]
)

已基于numpy的sliding_window_view实现检测逻辑,但该函数需用于游戏AI树搜索,会被频繁调用,当前性能仍有优化空间。cProfile显示大量reduce、dictcomp等调用耗时较长,推测np.any(np.all(windows))的方式会创建大量布尔数组导致低效。希望找到更优的检测方案,若能同时获取模式的位置和出现次数则更佳。

现有实现代码:

import numpy as np
import cProfile
import pstats
import time

class Board:
    def __init__(self, size):
        self.data = np.zeros((size, size), dtype=np.byte)
        self.size = size
        # 生成每行每列的5格滑动窗口
        self.rowWindows = np.lib.stride_tricks.sliding_window_view(self.data, window_shape=(1,5))
        self.colWindows = np.lib.stride_tricks.sliding_window_view(np.transpose(self.data), window_shape=(1,5))

        # 生成两个对角线方向的滑动窗口
        # 由于不同对角线的窗口尺寸不同,用object数组存储
        self.antiDiagonalWindow = np.array(
            [
                np.lib.stride_tricks.sliding_window_view(np.fliplr(self.data).diagonal(offset=i), window_shape=(5,))
                for i in range(-self.size + 5, self.size - 5 + 1, 1)
            ],
            dtype=object,
        )
        self.diagonalWindow = np.array(
            [
                np.lib.stride_tricks.sliding_window_view(self.data.diagonal(offset=i), window_shape=(5,))
                for i in range(-self.size + 5, self.size - 5 + 1, 1)
            ],
            dtype=object,
        )

    def hasFiveInRow(self, value):
        return (
           np.any(np.all(self.rowWindows == value, -1),)
            or np.any(np.all(self.colWindows == value,-1), )
            # 对角线窗口需要拼接成二维数组才能统一处理
            or np.any(np.all(np.concatenate(self.antiDiagonalWindow) == value, -1), )
            or  np.any(np.all(np.concatenate(self.diagonalWindow) == value, -1), )
        )

def benchMark():
    b = Board(size=7)
    b.data[:]=np.random.randint(low=0, high=3, size=(7,7))

    for i in range(100_000):
        val = b.hasFiveInRow(1)

# t0 = time.time()
# benchMark()
# print(time.time() - t0)
with cProfile.Profile() as p:
    benchMark()
    res = pstats.Stats(p)
    res.sort_stats(pstats.SortKey.TIME)
    res.print_stats()

优化方案

核心思路:减少中间数组创建,支持提前终止检测

原方案依赖滑动窗口生成大量临时布尔数组,且必须遍历所有窗口才能返回结果。优化方向为:

  • 直接遍历矩阵,实时统计连续目标值,达到5个时可立即返回存在性(若仅需判断存在)
  • 扫描过程中同步记录连珠的位置和出现次数
  • 避免numpy高阶函数的额外开销,用原生循环降低调用栈复杂度

优化后代码实现

import numpy as np
import cProfile
import pstats
import time

class OptimizedBoard:
    def __init__(self, size):
        self.data = np.zeros((size, size), dtype=np.byte)
        self.size = size

    def check_five_in_row(self, value):
        found = False
        count = 0
        positions = []

        # 检测行方向
        for row_idx in range(self.size):
            current_streak = 0
            for col_idx in range(self.size):
                if self.data[row_idx, col_idx] == value:
                    current_streak += 1
                    if current_streak == 5:
                        found = True
                        count += 1
                        positions.append(('row', row_idx, col_idx - 4))
                        # 处理连续超过5个的情况,每出现一组连续5就计数一次
                        current_streak -= 1
                else:
                    current_streak = 0
            # 若仅需判断存在,可在此处提前返回:if found: return True, count, positions

        # 检测列方向
        for col_idx in range(self.size):
            current_streak = 0
            for row_idx in range(self.size):
                if self.data[row_idx, col_idx] == value:
                    current_streak += 1
                    if current_streak == 5:
                        found = True
                        count += 1
                        positions.append(('col', row_idx - 4, col_idx))
                        current_streak -= 1
                else:
                    current_streak = 0
            # 若仅需判断存在,可在此处提前返回

        # 检测正对角线(左上→右下)
        # 起点在第一行
        for j in range(self.size - 4):
            current_streak = 0
            for k in range(self.size - j):
                row, col = k, j + k
                if col >= self.size:
                    break
                if self.data[row, col] == value:
                    current_streak += 1
                    if current_streak == 5:
                        found = True
                        count += 1
                        positions.append(('diag_down', row - 4, col - 4))
                        current_streak -= 1
                else:
                    current_streak = 0
        # 起点在第一列
        for i in range(1, self.size - 4):
            current_streak = 0
            for k in range(self.size - i):
                row, col = i + k, k
                if row >= self.size:
                    break
                if self.data[row, col] == value:
                    current_streak += 1
                    if current_streak == 5:
                        found = True
                        count += 1
                        positions.append(('diag_down', row - 4, col - 4))
                        current_streak -= 1
                else:
                    current_streak = 0

        # 检测反对角线(右上→左下)
        # 起点在第一行
        for j in range(4, self.size):
            current_streak = 0
            for k in range(j + 1):
                row, col = k, j - k
                if row >= self.size:
                    break
                if self.data[row, col] == value:
                    current_streak += 1
                    if current_streak == 5:
                        found = True
                        count += 1
                        positions.append(('diag_up', row - 4, col + 4))
                        current_streak -= 1
                else:
                    current_streak = 0
        # 起点在最后一列
        for i in range(1, self.size - 4):
            current_streak = 0
            for k in range(self.size - i):
                row, col = i + k, self.size - 1 - k
                if row >= self.size:
                    break
                if self.data[row, col] == value:
                    current_streak += 1
                    if current_streak == 5:
                        found = True
                        count += 1
                        positions.append(('diag_up', row - 4, col + 4))
                        current_streak -= 1
                else:
                    current_streak = 0

        return found, count, positions

def optimized_benchMark():
    b = OptimizedBoard(size=7)
    b.data[:]=np.random.randint(low=0, high=3, size=(7,7))

    for i in range(100_000):
        found, count, pos = b.check_five_in_row(1)

# 对比性能测试
print("原方案性能测试:")
with cProfile.Profile() as p:
    benchMark()
    res = pstats.Stats(p)
    res.sort_stats(pstats.SortKey.TIME)
    res.print_stats(10)

print("\n优化方案性能测试:")
with cProfile.Profile() as p:
    optimized_benchMark()
    res = pstats.Stats(p)
    res.sort_stats(pstats.SortKey.TIME)
    res.print_stats(10)

优化效果说明

  1. 内存开销降低:无需预先生成滑动窗口和布尔数组,仅用少量变量记录连续计数。
  2. 执行速度提升:避免numpy高阶函数的调用栈开销,对于7x7矩阵的频繁调用场景,性能提升显著。
  3. 额外信息支持:同步返回连珠的类型(行/列/对角线)、起始坐标和出现次数,满足游戏AI的扩展需求。
  4. 提前终止能力:若仅需判断是否存在连珠,可在检测到第一个匹配项时立即返回,无需遍历整个矩阵。

进一步优化方向

  • 局部检测:游戏AI场景中每次仅落一子,只需检查落子位置附近的行、列、对角线,无需全量扫描。
  • JIT编译:用numba对扫描函数进行即时编译,进一步降低Python循环的执行开销。

内容的提问来源于stack exchange,提问作者goldenotaste

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.06 07:09:53