分治让我魅力无限

· · 题解

这是一个(自认为)有一定启发性的 O(n\log^2n) 的分治优化dp做法喵

刚看到题还是没什么思路的,先来考虑一些简单的问题

Q:只考虑一个怪物 i,要选择什么宝剑 j 击败它最优?

A:贪心地选 h_i\leq s_jjc_j 最小的,这个很好求,让宝剑对于 s_j 从小到大排序后维护后缀最小 c_j 的值,对于 i 二分求出击败它所需要消耗的最小代价 v_i ,另外不难发现若 h_i<h_{i'}, 则 v_i<v_{i'}

Q:对于一段连续怪物区间 [l,r] ,只用一把宝剑击败这里所有怪物的最小花费是多少?

A:显然是 \min_{i=l}^{r}v_i,注意只用一把宝剑击败这里所有怪物的条件是 r-l+1\le k

我们容易发现上面过程中的关键量是拥有的钱数,消耗钱才能买剑,用剑打完怪物之后我们又能赚到钱

于是可以考虑对钱数划分状态进行dp

以下定义 dp_i 为击败前 i 个怪物后剩下钱数的最大值,c_i 是击败前 i 个怪物获得钱数的前缀和

dp_i=c_i+\max_{j\in[i-k,i-1] \land \max_{p=j+1}^{i}v_p \leq dp_j}dp_j-c_j-\max_{p=j+1}^{i}v_p

初始状态 dp_0=0,判断有没有解,只需要看 dp_n 是大于等于 0 即可

于是我们得到一个 O(n^2) 的做法,可喜可贺

考虑优化,这个式子最麻烦的地方在于取区间最值,按顺序转移不太好做(其实是因为做这题的时候我还不会),于是考虑分治的求,设计分治函数 slove(l,r),表示已经计算准确 [0,l-1]dp 值之后,再求出 [l,r] 的dp值

根据定义,对于 slove(l,r) 先递归求出 slove(l,mid),再将 [l,mid] 的dp值转移给 [mid+1,r],最后递归求出 slove(mid+1,r),就可以求出整个区间的dp值了

这时候再去考虑 [l,mid] 的dp转移至 [mid+1,r],记 mx_i 表示 \max_{j\in[\min\{i,mid\},\max\{i,mid\}]}v_j,对于 L\in[l,mid],R\in[mid+1,r]

dp_R=c_R+\max_{L\in[R-k,R-1]\land \max \{ mx_{L+1},mx_{R}\} \leq dp_L}dp_L-c_L-\max \{ mx_{L+1},mx_{R}\}

于是我们把区间的最大值变成了两点的最大值,看起来好做很多

首先先把 L=mid 的情况特殊考虑,直接贡献给所有 R,因为上面的转移式用到了 mx_{L+1},但是当 L=mid 时就会跨越到右边导致转移错误

我们枚举 L,考虑 LR\in[mid+1,L+k] 的贡献

mx_L<mx_R,则 mx_R\leq dp_L , dp_L-c_L-mx_R\to dp_R

mx_R\leq mx_L,则mx_L\leq dp_L , dp_L-c_L-mx_L\to dp_R

对两种情况分别二分出 R 的区间,区间的左端点和右端点分别打上 dp_L-c_L-mx_Rdp_L-c_L-mx_L 的标记,最后对 R 做扫描线就可以完成转移了

可能我讲的有些复杂了,实现其实不难,细节很少,注意最后对 R 做扫描线的时候需要一个支持插入元素,删除元素,查询最大值的数据结构,容易想到multiset,但是这东西常数很大,可以换成手写的可删堆优化常数(不会写可删堆可以看看我的代码理解一下,挺好懂的)

代码:(自认为码风优良)

#ifdef __unix__
#define gc getchar_unlocked
#else
#define gc _getchar_nolock
#endif
#include<bits/stdc++.h>
using namespace std;
#define int long long
int dp[500005],sum[500005];
int a[500005];
int mx[500005],mn[500005];
int vl[500005];
int n,m,k;
struct node{
    int s,c;
}q[500005];
struct DelHeap
{
    priority_queue<int> pq,del;
    void insert(int x)
    {
        pq.push(x);
    }
    void erase(int x)
    {
        del.push(x);
    }
    void clear()
    {
        while(!del.empty() && pq.top() == del.top())
        {
            pq.pop();
            del.pop();
        }
    }
    bool empty()
    {
        clear();
        return pq.empty();
    }
    int rbegin()
    {
        clear();
        return pq.top();
    }
}st1,st2;
vector<int> e[500005][2];
vector<int> g[500005][2];
bool cmp(node a,node b)
{
    if(a.s != b.s) return a.s < b.s;
    return a.c < b.c;
}
void solve(int l,int r)
{
    if(l == r) return;
    int mid = l + r >> 1;
    solve(l,mid);
    mx[mid] = vl[mid],mx[mid + 1] = vl[mid + 1];
    for(int i = mid + 2; i <= r; i++) mx[i] = max(vl[i],mx[i - 1]);
    for(int R = mid + 1; R <= min(mid + k,r); R++) if(dp[mid] >= mx[R]) dp[R] = max(dp[R],dp[mid] - mx[R] - sum[mid] + sum[R]);
    for(int L = mid - 1; L >= max(l,mid - k + 1); L--)
    {
        mx[L] = max(vl[L],mx[L + 1]);
        if(dp[L] < mx[L + 1]) continue;
        int start = mid + 1;
        int ll = start,rr = min(r,L + k) + 1;
        while(rr - ll > 1 && mx[start] < mx[L + 1])
        {
            int mmid = ll + rr >> 1;
            if(mx[mmid] < mx[L + 1]) ll = mmid;
            else rr = mmid;
        }
        if(mx[start] < mx[L + 1])
        {
            e[start][0].push_back(dp[L] - mx[L + 1] - sum[L]);
            e[ll][1].push_back(dp[L] - mx[L + 1] - sum[L]);
            ll++;
        }
        rr = min(r,L + k) + 1;
        start = ll;
        if(start >= rr || mx[start] < mx[L + 1]) continue; 
        while(rr - ll > 1)
        {
            int mmid = ll + rr >> 1;
            if(mx[mmid] <= dp[L]) ll = mmid;
            else rr = mmid;
        }
        if(mx[start] >= mx[L + 1] && mx[start] <= dp[L])
        {
            g[start][0].push_back(dp[L] - sum[L]);
            g[ll][1].push_back(dp[L] - sum[L]);
        }
    }
    for(int R = mid + 1; R <= r; R++)
    {
        for(int val : e[R][0]) st1.insert(val);
        for(int val : g[R][0]) st2.insert(val);
        if(!st1.empty()) dp[R] = max(dp[R],(st1.rbegin()) + sum[R]);
        if(!st2.empty()) dp[R] = max(dp[R],(st2.rbegin()) - mx[R] + sum[R]);
        for(int val : e[R][1]) st1.erase(val);
        for(int val : g[R][1]) st2.erase(val);
        e[R][0].clear();
        e[R][1].clear();
        g[R][0].clear();
        g[R][1].clear();
    }
    solve(mid + 1,r);
}
void read(int &x) 
{
    int f = 1; x = 0; char ch = gc();
    while (!(ch >= '0' && ch <= '9')) 
    {
        if (ch == '-') f = -1;
        ch = gc();
    }
    while (ch >= '0' && ch <= '9') 
    {
        x = x * 10 + (ch - '0');
        ch = gc();
    } 
    x *= f;
}
main()
{
    cin >> n >> m >> k >> dp[0];
    for(int i = 1; i <= n; i++)
    {
        read(a[i]),read(sum[i]);
        sum[i] += sum[i - 1];
    }
    for(int i = 1; i <= m; i++)
        read(q[i].s),read(q[i].c);
    sort(q + 1,q + m + 1,cmp);
    mn[m + 1] = 1e18;
    for(int i = m; i >= 1; i--)
        mn[i] = min(mn[i + 1],q[i].c);
    for(int i = 1; i <= n; i++)
    {
        if(q[m].s < a[i]) 
        {
            cout << "No";
            return 0;
        }
        int l = 0,r = m;
        while(r - l > 1)
        {
            int mid = l + r >> 1;
            if(q[mid].s >= a[i]) r = mid;
            else l = mid;
        }
        vl[i] = mn[r];
        dp[i] = -1e18;
    }
    solve(0,n);
    if(dp[n] >= 0)cout << "Yes\n";
    else cout << "No\n";
}

时间复杂度 O(n\log^2n),瓶颈在于分治函数里面的二分和可删堆

这题想到分治还是因为出现了区间最值这种很容易合并的操作,[l,r] 的最值可以由 [l,mid][mid+1,r] 合并得来,便于用分治处理