数独生成器移除单元格性能瓶颈求助(高级难度耗时超30秒)
数独高级难度单元格移除性能优化
问题根源
你的代码在高级难度下耗时过长,主要因为三个低效点:
- 每次移除单元格都依赖错误的浅拷贝(原代码
board.copy()仅拷贝外层列表,内部子列表仍为引用,会意外修改原棋盘),且拷贝操作本身带来额外开销; - 朴素递归回溯的
solve函数按固定顺序尝试填数,空单元格较多时会产生巨量无效回溯; is_valid函数频繁创建临时列表用于列和宫的检查,重复计算浪费资源。
优化方案
1. 修复is_valid,消除临时列表
将列和宫的检查改为直接遍历,避免生成不必要的列表:
def is_valid(array, r, c, e): size = 3 # 检查行 if e in array[r]: return False # 检查列 for i in range(9): if array[i][c] == e: return False # 检查宫 start_r = (r // size) * size start_c = (c // size) * size for i in range(start_r, start_r + size): for j in range(start_c, start_c + size): if array[i][j] == e: return False return True
2. 重写solve,用最少约束启发式减少回溯
优先填充候选数最少的单元格,能大幅降低回溯次数,这是提升解算速度的核心优化:
def solve(board): # 收集所有空单元格及其候选数 empty = [] for r in range(9): for c in range(9): if board[r][c] == 0: candidates = [] for num in range(1, 10): if is_valid(board, r, c, num): candidates.append(num) if not candidates: return False # 无候选数,直接无解 empty.append((len(candidates), r, c)) if not empty: return True # 无空单元格,已解 # 按候选数数量排序,优先处理最容易确定的单元格 empty.sort() _, r, c = empty[0] for num in range(1, 10): if is_valid(board, r, c, num): board[r][c] = num if solve(board): return True board[r][c] = 0 # 回溯 return False
3. 优化remove_cells,修复拷贝bug+随机化尝试顺序
改用copy.deepcopy保证棋盘独立性,同时随机打乱单元格顺序,更快达到目标移除数量:
import copy import random def remove_cells(board, difficulty = Difficulty.EASY): # 定义各难度保留的单元格数量 keep_config = { Difficulty.ADVANCED: random.randint(20, 29), Difficulty.INTERMEDIATE: random.randint(30, 39), Difficulty.EASY: random.randint(40, 49) } target_keep = keep_config[difficulty] current_keep = 81 # 生成随机顺序的单元格列表 cells = [(i, j) for i in range(9) for j in range(9)] random.shuffle(cells) for i, j in cells: if current_keep <= target_keep: break if board[i][j] == 0: continue temp_val = board[i][j] board[i][j] = 0 # 深拷贝棋盘,避免影响原棋盘状态 test_board = copy.deepcopy(board) if solve(test_board): current_keep -= 1 else: board[i][j] = temp_val return board
进阶优化:缓存已用数字(可选)
如果想进一步提速,可以预缓存行、列、宫的已用数字,让is_valid判断变为O(1)操作:
# 初始化缓存(生成完整棋盘时同步更新) row_used = [[False]*10 for _ in range(9)] col_used = [[False]*10 for _ in range(9)] box_used = [[False]*10 for _ in range(9)] # 优化后的is_valid def is_valid(r, c, e): box_idx = (r // 3) * 3 + (c // 3) return not row_used[r][e] and not col_used[c][e] and not box_used[box_idx][e]
内容的提问来源于stack exchange,提问作者charles
相关产品推荐
相关产品推荐

