DP记忆化递归:存储调用结果与直接嵌套max的计算异常排查
记忆化DFS实现findMaxForm时递归调用结果不一致问题
我在实现基于记忆化DFS的DP解决方案(findMaxForm问题)时,遇到了一个异常情况:将递归调用结果存入变量后取max,与直接在max函数中嵌套递归调用,两种写法得到的结果不一致。针对给定测试用例,正确输出应为17,但错误写法输出14。关闭记忆化时结果正确,因此怀疑记忆化逻辑存在问题。
测试用例代码
int main() { Solution obj = Solution(); vector<string> arr {"0","11","1000","01","0","101","1","1","1","0","0","0","0","1","0", "0110101","0","11","01","00","01111","0011","1","1000","0","11101", "1","0","10","0111"}; // output: 17 cout << obj.findMaxForm(arr, 9, 80) << endl; return 0; }
原始实现代码
#include <iostream> #include <vector> using namespace std; class Solution { public: int dfsHelper(vector<string>& strs, int idx, pair<int, int>& target, pair<int, int> curr, int result, vector<vector<vector<int>>>& memo) { if (curr.first > target.first || curr.second > target.second) return 0; if (idx >= strs.size()) return result; if (memo[idx][curr.first][curr.second] != -1) return memo[idx][curr.first][curr.second]; pair<int, int> addition {0, 0}; for (char& ch: strs[idx]){ if (ch == '0') addition.first += 1; else addition.second += 1; } // int leave = dfsHelper(strs, idx+1, target, curr, result, memo); // int take = dfsHelper(strs, idx+1, target, {curr.first+addition.first, curr.second+addition.second}, result+1, memo); // memo[idx][curr.first][curr.second] = max(leave, take); memo[idx][curr.first][curr.second] = max( dfsHelper(strs, idx+1, target, curr, result, memo), dfsHelper(strs, idx+1, target, {curr.first+addition.first, curr.second+addition.second}, result+1, memo) ); return memo[idx][curr.first][curr.second]; } int findMaxForm(vector<string>& strs, int m, int n) { pair<int, int> target {m, n}; vector<vector<vector<int>>> memo(strs.size(), vector<vector<int>>(m+1, vector<int>(n+1, -1))); return dfsHelper(strs, 0, target, {0, 0}, 0, memo); } };
问题根源分析
max参数求值顺序的不确定性:C++标准未定义max函数参数的求值顺序,编译器可能先计算右侧的take分支。当take分支执行时,会修改memo中的状态,后续执行左侧leave分支时,会读取到已经被修改的缓存值,导致错误的缓存命中,最终得到错误结果。而将两个递归调用存入变量的写法,是严格按顺序执行leave再执行take,此时memo未被污染,能正确获取缓存。- 冗余的
result参数导致状态不一致:递归函数中传递的result参数是累计的子集大小,但memo的状态定义是[idx][curr.first][curr.second],未包含result。当缓存被错误覆盖时,result的累计逻辑会混乱,进一步加剧结果错误。
修正方案
- 移除冗余的
result参数,让递归函数直接返回当前状态下能得到的最大子集大小,保证memo状态定义的自洽性。 - 预先计算每个字符串的0和1的数量,避免递归中重复计算,提升效率。
- 始终用变量存储两个分支的递归结果,再取
max,确保执行顺序可控,避免memo被提前污染。
修正后的代码
#include <iostream> #include <vector> #include <algorithm> using namespace std; class Solution { public: int dfsHelper(vector<string>& strs, int idx, int remain0, int remain1, vector<vector<vector<int>>>& memo, vector<pair<int, int>>& cnts) { if (idx >= strs.size()) return 0; if (memo[idx][remain0][remain1] != -1) return memo[idx][remain0][remain1]; // 不选当前字符串的分支 int skip = dfsHelper(strs, idx + 1, remain0, remain1, memo, cnts); // 选当前字符串的分支(剩余容量足够时) int take = 0; int need0 = cnts[idx].first; int need1 = cnts[idx].second; if (remain0 >= need0 && remain1 >= need1) { take = 1 + dfsHelper(strs, idx + 1, remain0 - need0, remain1 - need1, memo, cnts); } memo[idx][remain0][remain1] = max(skip, take); return memo[idx][remain0][remain1]; } int findMaxForm(vector<string>& strs, int m, int n) { // 预计算每个字符串的0、1数量 vector<pair<int, int>> cnts; for (string& s : strs) { int c0 = 0, c1 = 0; for (char ch : s) { ch == '0' ? c0++ : c1++; } cnts.emplace_back(c0, c1); } // memo维度:当前索引、剩余0的数量、剩余1的数量 vector<vector<vector<int>>> memo(strs.size(), vector<vector<int>>(m + 1, vector<int>(n + 1, -1))); return dfsHelper(strs, 0, m, n, memo, cnts); } }; // 测试用例 int main() { Solution obj = Solution(); vector<string> arr {"0","11","1000","01","0","101","1","1","1","0","0","0","0","1","0", "0110101","0","11","01","00","01111","0011","1","1000","0","11101", "1","0","10","0111"}; cout << obj.findMaxForm(arr, 9, 80) << endl; // 输出17 return 0; }
内容的提问来源于stack exchange,提问作者Mohamed Samir
相关产品推荐
相关产品推荐

