[省选联考 2024]迷宫守卫 题解

主要思路就是贪心 + dp,其中细节很多,所以不简单。

首先显然,我们一定是不惜代价地让前面的值大于后面的值,于是我们考虑,让第一个数最大该如何做。

我们发现要让第一个数是某个数的代价,是有单调性的,也就是说,我们要让第一个数越大,那么付出的代价一定大于等于让他比较小,所以显然使用二分,二分一个值 \(mid\) 表示让第一个数是 \(\ge mid\) 的,二分出来最大的 \(mid\) 使得付出的代价 \(\le k\),至于计算代价,可以使用简单的 dp,设 \(f_{u}\) 表示如果 Bob 走到 \(u\) 这个子树,想要让 Bob 无论如何走都必须走到一个值大于等于当前 \(mid\) 的最小代价。我们考虑转移,首先左子树是一定要走的,所以必须加上 \(f_{lson_u}\),而右子树呢,我们可以让 Bob 走,那么代价就是 \(f_{rson_u}\),但是我们也可以唤醒石像,这时候 Bob 必须先走左面,而我们现在只考虑第一个数最大,所以右面就被排除了,所以我们的代价仅仅是 \(w_u\)。这时候的转移就是

\[f_u = f_{lson_u} + \min(f_{rson_u}, w_u) \]

边界是叶子节点,走到了叶子节点,那么第一个数就是叶子这个值了,所以如果这个值 \(\ge mid\),那么代价就是 \(0\),否则就是 \(inf\),记住这个 \(inf\) 是要 \(> 10^{12}\) 的,因为他必须大于输入的 \(K\) 的取值,意味着如果走到了这个叶子并且这个叶子的值 \(< mid\),那么第一个数必然不是 \(\ge mid\) 的。

现在我们解决了重要一步,然后考虑,如何确定别的数的取值。

image

如图,假设我们现在能走到的最大的值在白色箭头所指(\(9\) 号),Bob 一定会沿着红色箭头所指前行,那么走到 \(9\) 之后,会回到他的父亲,然后去兄弟节点。去兄弟节点这一必然事件是我们解决问题的极大关键。因为他去了兄弟节点,那么下面的一些值一定是在他兄弟节点的子树上产生的,产生完之后,就会回到父亲,然后到父亲的兄弟,往复下去,最终到达根节点结束程序。那么在第一个数已经确定了的情况下,他兄弟节点的子树上产生的答案如何计算?显然先算第一个值最大是什么,然后再考虑别的,我们发现这是一个明显的子问题。

设 solve(u, sum) 表示 \(x\) 的子树中,我们提供 \(sum\) 的魔力值,所能获得的最优解,返回两个值,一个是消耗代价,一个是最优解序列,最终答案就是 solve(1, K) 返回的序列。

关于 \(sum\) 的计算是一个关键点,我们就拿上图举例,假设现在我们正在解决 \(4\) 的兄弟节点即 \(5\),那么他的 \(sum\) 就是原本的 K,减去比他深度低的 solve 消耗的代价,再减去比他深度高的 solve 为了达到效果,所要预留的代价。具体来说,我们已经运行完成了 solve(8, sum),那么他的代价我们就要减去,然后呢,我们在运行 solve(5, sum) 的同时其实还在运行 solve(1, sum) 中,那么 solve(1, sum) 已经确定了他的最优解序列中的第一个数,那么想要实现这个第一个数,我们必然要给人家预留一些魔力值,所以减去。

为了计算预留多少魔力值,我们模拟 dp 过程,我们发现 \(f_1\) 是从 \(f_2\)\(f_3\) 转移来的,最终我们走了 \(2\) 节点,这个结果是 Bob 选择的,也就说我们消耗了一定代价(可能是 \(0\))使得 Bob 如果走了 \(3\),最终的值一定还大于走 \(2\) 呢。那么消耗的代价正好是上述转移式的右侧。如果走了右子节点,那么就是转移式的左侧。

现在考虑一种情况,仍以上图举例,如果我们正在运算 solve(1, sum),并且刚好算完 solve(5, sum),结果我们发现算出来的最优解的字典序小于 solve(1, sum) 算出来的字典序(虽然 solve(1, sum) 还没算完,但是他的第一项绝对算出来了,所以我们只需要把两个序列的第一项比较就行了),那么 Bob 为什么会走 \(4\) 不走 \(5\)?当然是因为 \(5\) 的父亲也就是 \(2\) 的石像被唤醒了,注意刚才我们算预留魔力值时,只算了路径节点的兄弟节点的子树需要消耗的代价,并没有算路径节点,什么时候算?这时候算呀,我们发现显然是 \(2\) 的石像被唤醒了,所以我们就要把魔力值减去 \(w_2\),于是就解决了。

Show Code
#include
#define int long long
#define fi first
#define se second
#define mp make_pair
using namespace std;
const int N = 1.5e5 + 5, inf = 2e12;
const int BUFSIZE = 1 << 24;
char ibuf[BUFSIZE], *is = ibuf, *it = ibuf;
inline char getch(){
    if(is == it)
        it = (is = ibuf) + fread(ibuf, 1, BUFSIZE, stdin);
    return is == it ? EOF : *is++;
}
inline int mread(){
    int res = 0, ch = getch();
    while(!(isdigit(ch)) and ch != EOF) ch = getch();
    while(isdigit(ch)) res = res * 10 + (ch ^ 48), ch = getch();
    return res;
}
int t = mread(), w[N], a[N], n, k, nl[N], nr[N], f[N], belong[N], sum;
vector ve[N];
void build(int x, int ll, int rr){
    ve[x].clear();
    nl[x] = ll, nr[x] = rr;
    if(ll == rr){
        ve[x].push_back(a[nl[x]]);
        return;
    }
    int mid = (ll + rr) >> 1;
    build(x * 2, ll, mid);
    build(x * 2 + 1, mid + 1, rr);
    for(auto t : ve[x * 2])
    ve[x].push_back(t);
    for(auto t : ve[x * 2 + 1])
    ve[x].push_back(t);
    sort(ve[x].begin(), ve[x].end());
    return;
}
int dp(int x, int mid){
    if(nl[x] == nr[x]){
        if(a[nl[x]] >= mid)
        return f[x] = 0;
        else
        return f[x] = inf;
    }
    return f[x] = dp(x * 2, mid) + min(dp(x * 2 + 1, mid), w[x]);
}
pair > solve(int x, int sum){
    int la = sum;
    vector ans;
    if(nl[x] == nr[x]){
        ans.push_back(a[nl[x]]);
        return mp(0, ans);
    }
    int ls = 0, rs = nr[x] - nl[x];
    while(ls < rs){
        int mid = (ls + rs + 1) >> 1;
        if(dp(x, ve[x][mid]) <= sum)
        ls = mid;
        else
        rs = mid - 1;
    }
    ls = ve[x][ls];
    dp(x, ls);
    int y = belong[ls];
    while(x != y){
        if(y & 1){
            sum -= f[y ^ 1];
        }
        else{
            sum -= min(w[y >> 1], f[y ^ 1]);
        }
        y >>= 1;
    }
    y = belong[ls];
    ans.push_back(ls);
    while(x != y){
        if(y & 1){
            sum += f[y ^ 1];
        }
        else{
            sum += min(w[y >> 1], f[y ^ 1]);
        }
        auto t = solve(y ^ 1, sum);
        if(t.se[0] < ls){
            sum -= w[y >> 1];
            t = solve(y ^ 1, sum);
        }
        sum -= t.fi;
        for(int tmp : t.se)
        ans.push_back(tmp);
        y >>= 1;
    }
    return mp(la - sum, ans);
}
signed main(){
    while(t --){
        n = mread(), k = mread();
        n = (1 << n);
        for(int i = 1; i < n; i ++)
        w[i] = mread();
        for(int i = 1; i <= n; i ++){
            a[i] = mread();
            belong[a[i]] = i + n - 1;
        }
        build(1, 1, n);
        auto t = solve(1, k).se;
        for(auto tmp : t)
        printf("%lld ", tmp);
        printf("\n");
    }
    return 0;
}
posted @ 2024-03-10 21:04  cndark_moon  阅读(53)  评论(0)    收藏  举报