题解:P14328 [JOI2022 预选赛 R2] 糖 2 / Candies 2

· · 题解

题目传送门

题意

n 个糖果,第 i 颗糖美味度为 a_i,任意连续 K 颗糖果中最多只能吃两颗,求美味度之和的最大值。

思路

首先想到定义 dp_{i,j} 为在前 i 颗糖果中吃的最后一块糖是第 j 颗时的最大美味度之和。初始化 dp_{i,0}=a_i,容易想到 O(n^3) 的转移。枚举 i,j,w,若 w=0K \le i-w,就说明第 i 颗糖与第 w 颗糖不在连续 K 颗糖中,此时可以吃掉第 i 颗糖,dp_{i,j} \leftarrow \max(dp_{i,j},dp_{j,w}+a_i) ,否则此时无法吃掉第 i 颗糖,dp_{i,j} \leftarrow \max(dp_{i,j},dp_{j,w})

代码:

#include<bits/stdc++.h>
using namespace std;
#define int long long
const int N=3e3+5;
int n,k,a[N],dp[N][N],ans;
signed main(){
    cin>>n>>k;
    for(int i=1;i<=n;i++) cin>>a[i];
    for(int i=1;i<=n;i++) dp[i][0]=a[i];
    for(int i=1;i<=n;i++){
        for(int j=0;j<i;j++){
            for(int w=0;w<j;w++){
                if(i-w>=k||w==0) dp[i][j]=max(dp[i][j],dp[j][w]+a[i]);
                dp[i][j]=max(dp[i][j],dp[j][w]);
            }
        }
    }
    for(int i=0;i<n;i++) ans=max(ans,dp[n][i]);
    cout<<ans;
    return 0;
}

由于 2 \le K \le n \le 3000O(n^3) 的做法显然无法通过本题,考虑优化(蒟蒻只能想到优先队列了)。

我们可以用优先队列优化第三层枚举,用大根堆记录所有 dp_{i,j}i-w < K 的情况就被我们解决了,此时仅需枚举下标 \{ w \mid w \le i-K \},复杂度 O(n^2(n-K))

代码:

#include<bits/stdc++.h>
using namespace std;
#define int long long
const int N=3e3+5;
int n,k,a[N],dp[N][N],ans;
priority_queue<int> q[N];
signed main(){
    cin>>n>>k;
    for(int i=1;i<=n;i++) cin>>a[i];
    for(int i=1;i<=n;i++) dp[i][0]=a[i];
    for(int i=1;i<=n;i++){
        for(int j=0;j<i;j++){
            if(!q[j].empty())
                dp[i][j]=max(dp[i][j],q[j].top());
            for(int w=0;w<=i-k||w==0;w++){
                dp[i][j]=max(dp[i][j],dp[j][w]+a[i]);
            }
            q[i].push(dp[i][j]);
        }
    }
    for(int i=0;i<n;i++) ans=max(ans,dp[n][i]);
    cout<<ans;
    return 0;
}

还是超时,继续优化。考虑继续用优先队列,每次枚举结束后存入 \{ dp_{j,i+1-K} \mid i+1-K < j \le i \},这样在枚举到下一个位置时,堆中存的是所有满足 w \le i-Kdp_{j,w}。这样我们就彻底解决了第三层循环,复杂度为 O(n^2 \log n),跑起来也并不慢。记录。

代码

#include<bits/stdc++.h>
using namespace std;
#define int long long
const int N=3e3+5;
int n,k,a[N],dp[N][N],ans;
priority_queue<int> q[N],q1[N];
signed main(){
    cin>>n>>k;
    for(int i=1;i<=n;i++) cin>>a[i];
    for(int i=1;i<=n;i++) dp[i][0]=a[i];
    for(int i=1;i<=n;i++){
        q1[i].push(dp[i][0]);
        for(int j=0;j<i;j++){
            if(!q[j].empty())
                dp[i][j]=max(dp[i][j],q[j].top());
            if(!q1[j].empty())
                dp[i][j]=max(dp[i][j],q1[j].top()+a[i]);
            q[i].push(dp[i][j]);
        }
        if(i+1-k>=0) for(int j=i;j>i+1-k;j--)
            q1[j].push(dp[j][i+1-k]);
    }
    for(int i=0;i<n;i++) ans=max(ans,dp[n][i]);
    cout<<ans;
    return 0;
}

完结撒花!