如何优化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)
优化效果说明
- 内存开销降低:无需预先生成滑动窗口和布尔数组,仅用少量变量记录连续计数。
- 执行速度提升:避免numpy高阶函数的调用栈开销,对于7x7矩阵的频繁调用场景,性能提升显著。
- 额外信息支持:同步返回连珠的类型(行/列/对角线)、起始坐标和出现次数,满足游戏AI的扩展需求。
- 提前终止能力:若仅需判断是否存在连珠,可在检测到第一个匹配项时立即返回,无需遍历整个矩阵。
进一步优化方向
- 局部检测:游戏AI场景中每次仅落一子,只需检查落子位置附近的行、列、对角线,无需全量扫描。
- JIT编译:用numba对扫描函数进行即时编译,进一步降低Python循环的执行开销。
内容的提问来源于stack exchange,提问作者goldenotaste
相关产品推荐
相关产品推荐

