递归求和组合函数的记忆化实现异常问题排查
记忆化优化导致how_sum函数输出异常的原因排查
问题背景
需要实现how_sum函数:输入目标值target(整数)与整数列表nums,返回任意一组和为target的元素组合;若无符合条件组合则返回None。初始递归版本功能正常,添加记忆化(memo)优化后输出异常,注释掉memo[target] = rv + [num]这行代码后结果恢复正确,需排查该赋值语句导致错误的原因。
错误输出示例
def main(): print(how_sum(7, [2, 3])) # [3, 2, 2] print(how_sum(7, [5, 3, 4, 7])) # [3, 2, 2] (错误:应返回[4,3]或其他合法组合) print(how_sum(7, [2, 4])) # [3, 2, 2] (错误:应返回None) print(how_sum(8, [2, 3, 5])) # [2, 2, 2, 2] (错误:应返回None) print(how_sum(500, [7, 14])) # [3, 7, 7, ..., 7] (错误:应返回None)
正确输出示例
def main(): print(how_sum(7, [2, 3])) # [3, 2, 2] print(how_sum(7, [5, 3, 4, 7])) # [4, 3] print(how_sum(7, [2, 4])) # None print(how_sum(8, [2, 3, 5])) # None print(how_sum(500, [7, 14])) # None
问题代码
def how_sum(target: int, nums: list[int], memo: dict[int, list[int]] = {}) -> list[int] | None: if target in memo: return memo[target] if target == 0: return [] if target < 0: return None for num in nums: remainder = target - num rv = how_sum(remainder, nums, memo) if rv is not None: memo[target] = rv + [num] # 注释此行后结果正常 return rv + [num] memo[target] = None return None
错误原因分析
1. 默认参数的可变对象陷阱
Python中函数的默认参数在函数定义时就完成初始化,而非每次调用时。这里memo={}会创建一个全局共享的字典,所有未手动传入memo的函数调用都会复用这个字典。
比如第一个测试用例how_sum(7, [2,3])会把target=7对应的结果[3,2,2]存入memo;后续调用how_sum(7, [5,3,4,7])时,函数发现target=7已在memo中,直接返回之前缓存的结果,完全忽略当前传入的nums参数已经变化,导致输出错误。
2. 记忆化逻辑未关联nums参数
当前的记忆化仅以target作为缓存键,但同一个target在不同nums列表下的结果可能完全不同(比如target=7在[2,3]下有解,在[2,4]下无解)。即使修复了默认参数的问题,现有逻辑也无法区分不同nums对应的缓存,依然会出现结果错误。
修复方案
方案1:修复默认参数问题(解决当前测试用例错误)
将默认参数改为None,在函数内部初始化新字典,确保每次调用(未传memo时)都使用独立的缓存:
def how_sum(target: int, nums: list[int], memo: dict[int, list[int]] = None) -> list[int] | None: if memo is None: memo = {} if target in memo: return memo[target] if target == 0: return [] if target < 0: return None for num in nums: remainder = target - num rv = how_sum(remainder, nums, memo) if rv is not None: memo[target] = rv + [num] return rv + [num] memo[target] = None return None
方案2:支持不同nums的缓存(更严谨)
如果需要在多次调用不同nums时也能正确缓存,可将(target, tuple(nums))作为缓存键(列表不可哈希,需转成元组):
def how_sum(target: int, nums: list[int], memo: dict[tuple[int, tuple[int]], list[int]] = None) -> list[int] | None: if memo is None: memo = {} key = (target, tuple(nums)) if key in memo: return memo[key] if target == 0: return [] if target < 0: return None for num in nums: remainder = target - num rv = how_sum(remainder, nums, memo) if rv is not None: memo[key] = rv + [num] return rv + [num] memo[key] = None return None
内容的提问来源于stack exchange,提问作者CStudent
相关产品推荐
相关产品推荐

