树套树神力!

· · 题解

upd on 2026/07/19:修复了数学公式没有使用 \LaTeX 的问题。

前言:个人认为本场 AT 题目难度 F>E>G,但是 E、F 题耗太久,G 题考场上 10min 极速写完树套树没时间调了 qwq。

主包主包,题解的做法太吃操作了,有没有简单做法?
有的兄弟,有的。树套树秒了它!

前置知识:树套树、主席树(了解思想即可)。

不知道大家是怎么理解树套树的,我的理解就是主席树的动态升级版,因为主席树只支持静态区间可差分问题,但是树套树支持修改,这道题用到这里就够了。

回到这道题,如果给你一个静态区间,让你选择最少的元素,使得选择元素的和大于 k,这不就是一个橙题贪心吗?直接从大到小选即可。

如果有多次询问,阁下应该如何应对呢?我会主席树!二分需要选几个元素,主席树维护值域 lr 的元素个数个元素总和,检查选出来的元素和有没有大于 k。这样就算是静态的一次询问也是 \log V\log n 的,因为最外层有一个二分。

当然二分可以优化掉:由于线段树是二叉树,我们在线段树上二分就可以减少一只 \log。具体的,就是在当前节点判断右子树的元素和是否大于 k,如果小于等于 k,说明整个右子树选完都不够,答案加上右子树元素个数,往左子树搜。否则说明右子树的元素不用全选,所以往右子树搜即可。

现在,你一定已经学会了在静态区间上 q\log V 查询了,我们把询问放在动态区间上。

主席树就不能胜任这个任务了,我们可以搬上树状数组套值域线段树了。思想仍然类似主席树,我们在树状数组把对应位置的根节点取出来,通过差分将这个区间的真实元素情况还原出来,树套树上二分即可实现查询。

修改也类似,就是修改树状数组的对应位置的线段树即可。时间复杂度 \mathcal O((n+q)\log V\log n),我的实现比较粗糙,最慢点也只跑了 1500ms+,完全不用担心时间。但是空间如果不回收是 \mathcal O((n+q)\log n\log V),在全开 long long 的情况下会 MLE。可以选择回收节点,空间减少到 \mathcal O(n\log V\log n)long long 就可以放心开了。

这道题的树套树的代码不长哦

:::success[code]

#define int long long
#define lc(x) tr[x].son[0]
#define rc(x) tr[x].son[1]
#define mid ((l+r)>>1)
int n,q,a[N];
int addnode[N],addcnt = 0,delnode[N],delcnt = 0,b[N];
struct Segment_Tree{
    int son[2];
    int cnt,sum;
}tr[N<<6];
int stcnt = 0,pool[N],top = 0;
int Newnode(){
    return top?pool[top--]:++stcnt;
}
void del(int x){
    lc(x) = rc(x) = tr[x].cnt = tr[x].sum = 0;
    pool[++top] = x;
    return ;
}
void add(int &x,int l,int r,int v){
    if(!x) x = Newnode();
    tr[x].sum += v;
    tr[x].cnt += 1; 
    if(l==r) return ;
    if(v<=mid) add(lc(x),l,mid,v);
    else add(rc(x),mid+1,r,v);
    return ;
}
void del(int &x,int l,int r,int v){
    tr[x].sum -= v;
    tr[x].cnt -= 1; 
    if(l==r){
        if(!tr[x].cnt) del(x),x = 0;
        return ;
    }
    if(v<=mid) del(lc(x),l,mid,v);
    else del(rc(x),mid+1,r,v);
    if(!tr[x].cnt) del(x),x = 0;
    return ;
}
void add(int x,int v){
    for(;x<=n;x+=x&-x) add(b[x],1,1e9,v);
    return ;
}
void del(int x,int v){
    for(;x<=n;x+=x&-x) del(b[x],1,1e9,v);
}
int query_t(int l,int r,int k){
    if(l==r) return (k+l-1)/l;
    int cnt = 0,sum = 0;
    for(int i=1;i<=addcnt;++i) if(addnode[i]) sum += tr[rc(addnode[i])].sum,cnt += tr[rc(addnode[i])].cnt;
    for(int i=1;i<=delcnt;++i) if(delnode[i]) sum -= tr[rc(delnode[i])].sum,cnt -= tr[rc(delnode[i])].cnt;
    if(sum>=k){
        for(int i=1;i<=addcnt;++i) if(addnode[i]) addnode[i] = rc(addnode[i]);
        for(int i=1;i<=delcnt;++i) if(delnode[i]) delnode[i] = rc(delnode[i]);
        return query_t(mid+1,r,k);
    }
    for(int i=1;i<=addcnt;++i) if(addnode[i]) addnode[i] = lc(addnode[i]);
    for(int i=1;i<=delcnt;++i) if(delnode[i]) delnode[i] = lc(delnode[i]);
//  if(cnt)printf("choose %d~%d sum:%lld,cnt:%lld,remain:%lld\n",mid+1,r,sum,cnt,k-sum);
    return cnt + query_t(l,mid,k-sum);
}
int query(int l,int r,int k){
    addcnt = 0,delcnt = 0;
    int tot = 0;
    for(int i=r;i;i-=i&-i) if(b[i]) addnode[++addcnt] = b[i],tot += tr[b[i]].sum;
    for(int i=l-1;i;i-=i&-i) if(b[i]) delnode[++delcnt] = b[i],tot -= tr[b[i]].sum;
    if(tot<k) return -1;
    return query_t(1,1e9,k);
}
signed main(){
    read(n,q);
    for(int i=1;i<=n;++i){
        read(a[i]);
        add(i,a[i]);
    }
    while(q--){
        int c,x,l,r,k;
        read(c,x,l,r,k);
        del(c,a[c]);
        a[c] = x;
        add(c,a[c]);
        write(query(l,r,k),'\n');
    }
    return 0;
}

::: 如果有我写的有问题的地方,欢迎大犇在评论区指正。