LeetCode 1416:回溯解法无法识别有效数组分区问题排查
问题排查:LeetCode 1416 恢复数组回溯代码错误
题目描述
一个程序本应输出整数数组,但遗漏了空格,将数组输出为数字字符串s。已知数组中所有整数都在[1, k]范围内,且无前置零。给定字符串s和整数k,返回可被打印为s的可能数组的数量。
我的代码
class Solution: def numberOfArrays(self, s: str, k: int) -> int: ans = 0 def backtrack(s, index, buffer): nonlocal ans, k # If we have reached the end of the string, # we have found a valid array of integers. if index == len(s): ans += 1 return # If the current digit is '0', we cannot form # a valid number using any of the following digits. if s[index] == "0": return # Try forming a number using the current digit and # the following digits. If the number is valid, # continue the backtracking process with the remaining # digits in the string. for i in range(index, len(s)): buffer *= 10 buffer += int(s[i]) if buffer <= k: backtrack(s, i + 1, buffer) else: # If the number is not valid, stop the backtracking # process and undo any changes made to the buffer. buffer //= 10 backtrack(s, 0, 0) return ans
测试情况
通过的测试用例
输入: s = "1000", k = 10000 输出: 1 输入: s = "1000", k = 10 输出: 0 输入: s = "1317", k = 2000 输出: 8
失败的测试用例
输入: s = "2020", k = 30 输出: 0 预期: 1
问题原因分析
核心错误在于回溯函数的buffer参数设计逻辑错误:
- 你传递的
buffer是上一个分割出的数字的值,但在当前递归层级中,错误地将这个buffer作为当前数字的初始值继续累加,导致构建出的是「上一个数字+当前位」的拼接值,而非从当前index开始的新数字。 - 以测试用例
s="2020", k=30为例:- 第一次递归分割出
20后,进入backtrack(s, 2, 20)(此时index=2对应字符'2')。 - 在这个递归层级中,用传入的
buffer=20开始累加:20*10 + 2 = 202,远大于30,回退buffer到20;接着循环到i=3,计算20*10 +0=200,仍大于30,循环结束,未触发有效递归。 - 但实际上,在
index=2的位置,需要构建的是从'2'开始的新数字:2和20,其中20符合条件,应递归到index=4(字符串末尾)使ans加1。
- 第一次递归分割出
修复方案
去掉传递上一个数字的buffer参数,在每个递归层级中重新构建当前数字:
class Solution: def numberOfArrays(self, s: str, k: int) -> int: ans = 0 n = len(s) def backtrack(index): nonlocal ans if index == n: ans += 1 return if s[index] == "0": return current_num = 0 for i in range(index, n): current_num = current_num * 10 + int(s[i]) if current_num > k: break # 后续数字只会更大,直接终止循环 backtrack(i + 1) backtrack(0) return ans
额外优化
上述回溯会存在重复计算,长字符串场景下会超时,建议加入记忆化搜索(动态规划):
class Solution: def numberOfArrays(self, s: str, k: int) -> int: MOD = 10**9 +7 n = len(s) memo = [-1]*n def dp(index): if index == n: return 1 if s[index] == "0": return 0 if memo[index] != -1: return memo[index] res =0 current_num =0 for i in range(index, n): current_num = current_num *10 + int(s[i]) if current_num >k: break res += dp(i+1) res %= MOD memo[index] = res return res return dp(0)
内容的提问来源于stack exchange,提问作者Ruan
相关产品推荐
相关产品推荐

