You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

使用map<pair<int,ll>,ll>实现记忆化反而触发TLE?求解析

CSES Apple Division问题:记忆化递归反而超时的原因分析

问题背景

在解决CSES Apple Division问题(将数组分成两个子集,使两子集和的差值最小)时,由于数组长度n≤21,且元素值可达1e9,无法使用常规的二维DP数组(sum维度过大)。原暴力递归代码可以正常运行,但改用map<pair<int, ll>, ll>存储状态实现记忆化后,反而触发了TLE(时间限制1秒)。

原递归代码:

ll helper(vector<int> & arr,int size,ll currt_sum,ll total_sum){
    if(size==0){
        return abs(2*currt_sum- total_sum);
    }
    // PICK
    ll v1 = helper(arr,size-1,currt_sum+arr[size-1],total_sum);
    // Not Pick
    ll v2 = helper(arr,size-1,currt_sum,total_sum);
    return min(v1,v2);
}

记忆化代码:

ll helper(vector<int> & arr,int size,ll x,ll t,map<pair<int,ll>,ll > &m1){
    if(size==0){
        return abs(2*x-t);
    }
    if(m1.find({size,x})!=m1.end()) return m1[{size,x}];
    ll v1 = helper(arr,size-1,x+arr[size-1],t,m1);
    ll v2 = helper(arr,size-1,x,t,m1);
    return m1[{size,x}] = min(v1,v2);
}
void solve()
{
    int n;
    cin>>n;
    vector<int> v1(n);
    for (int i = 0; i < n; ++i)
    {
        cin>>v1[i];
    }
    map<pair<int,ll>,ll> m1;
    ll t = accumulate(v1.begin(), v1.end(),0LL);
    cout<<helper(v1,n,0,t,m1)<<"\n";
    return;
}

(注:ll代表long long int)


超时原因分析

  • map的高常数开销
    map底层基于红黑树实现,每次状态的查询(find)和插入操作的时间复杂度是O(log k)(k为已存储状态数)。对于n=21的场景,总状态数最多是2^21≈2e6,虽然理论上记忆化能避免重复计算,但map的常数开销远大于单次递归调用的开销——红黑树的节点旋转、键值比较等操作的耗时,比直接递归两次的耗时还要高,累积下来反而让总运行时间超过了1秒。

  • 原暴力递归的实际开销可控
    n=21时,暴力递归的总调用次数是2^21 + 1≈2097153次,这个数量级在C++中属于可接受范围:单次递归调用仅做加法和min比较,操作非常简单,1秒内完全可以完成所有调用。记忆化引入的额外开销反而抵消了避免重复计算的收益,甚至拉低了整体效率。

  • 状态重复率有限
    对于随机生成的数组,不同递归路径产生的(size, curr_sum)重复状态数量并不多,记忆化能减少的重复计算有限,但付出的map操作代价却很高,最终得不偿失。


优化方案

如果要使用记忆化,建议替换map为哈希表(unordered_map),但需要为pair<int, ll>自定义哈希函数;或者更简单的方式,直接用二进制枚举所有子集:遍历从0到2^n-1的所有整数,每个整数的二进制位表示是否选取对应元素,计算子集和后记录最小差值。这种方法时间复杂度同样是O(2^n),但没有递归调用的栈开销,实际运行效率更高。

示例二进制枚举代码:

void solve() {
    int n;
    cin >> n;
    vector<int> arr(n);
    ll total = 0;
    for (int i = 0; i < n; ++i) {
        cin >> arr[i];
        total += arr[i];
    }
    ll min_diff = LLONG_MAX;
    for (int mask = 0; mask < (1 << n); ++mask) {
        ll sum = 0;
        for (int i = 0; i < n; ++i) {
            if (mask & (1 << i)) {
                sum += arr[i];
            }
        }
        min_diff = min(min_diff, abs(total - 2 * sum));
    }
    cout << min_diff << endl;
}

内容的提问来源于stack exchange,提问作者Red_RanGer

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.09 18:22:47