基于回溯法的数独求解器出现内存超限问题求助
数独回溯算法内存超限问题排查与修复
问题描述
我用以下基于回溯算法的Python代码求解数独,代码能输出正确结果,但测试时出现**内存超限(memory limit exceeded)**错误,怀疑代码陷入循环,相关代码如下:
import numpy as np def sudoku_solver_util(mat, many, f_count): if many[0] >= 9 ** (81 - f_count): return False for i in range(9): for j in range(9): if mat[i][j] == 0: counter = 0 for k in range(1, 10): counter += 1 many[0] += 1 if is_safe(mat, i, j, k, 0, f_count, many): mat[i][j] = k if sudoku_solver_util(mat, many, f_count): return True count = np.count_nonzero(mat) if many[0] >= 9**(81-f_count): return False mat[i][j] = 0 if counter == 9: return False return True def row_is_safe(mat, i, j, k): return not k in mat[i][:] def col_is_safe(mat, i, j, k): return not k in mat[:,j] def grid_is_safe(mat, i, j, k): return not k in small_sudoku(mat, i, j) def is_safe(mat, i, j, k, counter, count, many): return row_is_safe(mat, i, j, k) and col_is_safe(mat, i, j, k) and grid_is_safe(mat, i, j, k) def small_sudoku(mat, i, j): return mat[(i//3)*3:(i//3)*3+3, (j//3)*3:(j//3)*3+3] def sudoku_solver(mat): many = [0] f_count = np.count_nonzero(mat) if sudoku_solver_util(mat, many, f_count): return mat else: return np.zeros((9,9)) # input mat = np.zeros((9,9)) for i in range(0,9): mat[i] = list(map(int, input().split())) # output out_mat = sudoku_solver(mat) for row in out_mat: print(" ".join(map(str, row.astype(int))))
问题根源分析
- 无效的终止条件计算:
9 ** (81 - f_count)会随着空单元格数量剧增变成天文数字(比如空40格就是9^41),many[0]的累加毫无意义,反而会因为存储超大整数占用大量内存,递归中反复计算该值也会浪费资源。 - 低效的遍历逻辑:每次递归都从(0,0)开始扫描整个数独找空单元格,重复遍历会大幅增加递归深度和内存消耗。
- 冗余参数与无效操作:
is_safe函数的counter, count, many参数完全没用到;递归里的np.count_nonzero(mat)计算后也没使用,纯属于多余开销。 - 错误的回溯终止逻辑:在单元格尝试完9个数字后直接
return False,会提前终止整个递归树的其他分支,导致逻辑混乱,可能引发不必要的递归嵌套。
优化后的代码
import numpy as np def sudoku_solver_util(mat): # 定位第一个空单元格 for i in range(9): for j in range(9): if mat[i][j] == 0: # 尝试1-9所有可能数字 for k in range(1, 10): if is_safe(mat, i, j, k): mat[i][j] = k # 递归求解,成功则直接返回 if sudoku_solver_util(mat): return True # 回溯,重置当前单元格 mat[i][j] = 0 # 当前单元格无有效数字,回溯上一层 return False # 所有单元格填满,求解完成 return True def row_is_safe(mat, i, j, k): return k not in mat[i] def col_is_safe(mat, i, j, k): return k not in mat[:, j] def grid_is_safe(mat, i, j, k): start_row = (i // 3) * 3 start_col = (j // 3) * 3 return k not in mat[start_row:start_row+3, start_col:start_col+3] def is_safe(mat, i, j, k): return row_is_safe(mat, i, j, k) and col_is_safe(mat, i, j, k) and grid_is_safe(mat, i, j, k) def sudoku_solver(mat): # 复制原矩阵,避免修改输入数据 mat_copy = mat.copy() if sudoku_solver_util(mat_copy): return mat_copy else: return np.zeros((9,9), dtype=int) # 输入处理 mat = np.zeros((9,9), dtype=int) for i in range(9): mat[i] = list(map(int, input().split())) # 输出结果 out_mat = sudoku_solver(mat) for row in out_mat: print(" ".join(map(str, row)))
优化点说明
- 移除了冗余的计数器和无效终止条件,避免超大整数占用内存。
- 每次递归仅处理第一个空单元格,避免重复扫描整个数独矩阵,降低递归开销。
- 简化
is_safe函数参数,删除无用变量,提升函数调用效率。 - 复制输入矩阵,避免修改原始数据,符合编程规范。
- 修正回溯逻辑,仅在当前单元格无有效数字时返回False,保证递归树的正确遍历。
内容的提问来源于stack exchange,提问作者Manya
相关产品推荐
相关产品推荐

