CF1916E Happy Life in University 题解

题意

CF1916E Happy Life in University

\(T\)组测试数据。

你有一个n个节点的树,根为1.每个节点有一个类型a_i, 我们定义\(diff(u,v)\) 为树上\(u\)到v的路径中,节点类型的数量。我们想知道,对于任意的\(u,v\) \(diff(u,lca(u,v))*diff(v,lca(u,v))\)的最大值。

\(\sum_{i=1}^{T}n_i \leq 3*10^5\)

思考。

我们考虑从\(lca\)入手。

我们递归枚举\(lca\)节点,答案不同子树下的\(u,v\),最大的\(diff(u,lca),diff(v,lca)\)

由于是递归枚举。每次出现一个新的\(lca\),\(lca\)节点的类型就会加入到\(lca\)为根的子树的每一个节点开始的路径。涉及子树操作:考虑线段树。

有两种方向:

  • 允许种类计算重复,然后去重

  • 只更新还没有出现这个种类的路径

很明显第一个好做,考虑第一个。

假设没有重复,那么每次处理一个新的\(lca\),就出现一个新的类型。那么子树上所有节点开始的路径所经过的类型+1。

考虑如何去重。由于每次更新会覆盖子树的所有节点,所以之前所有的同类型贡献就重复了。我们要消除之前的权值

我们可以每次加入新的节点,然后线段树删除,记重复计算的点为\(u\),每次让以\(u\)为根的子树权值-1。

对于每个点下面与其相同类型的节点,可以dfs预处理。

由于每个点只会被添加一次,删除一次。线段树修改的复杂度为\(O(log_2n)\)所以总复杂度为\(O(nlog_2n)\)

为啥一开始我没有想出来

一直觉得是启发式合并,想是集合存集合还是bitset,然后一直不行。

这告诉我们做题立场不要太坚定

#include <bits/stdc++.h>
#define ls(x) ((x)<<1)
#define rs(x) ((x)<<1|1)

using std::cin;
using std::pair;
using std::max;
using std::min;
using std::swap;
using std::vector;

typedef long long ll;

const ll inf=1e18;
const ll maxn=3e5+5;

struct Edge{
    ll v,next;
};

struct node{
    ll maxx,tag;
};

node seg[maxn<<2];
vector<ll> vec[maxn];
ll pos[maxn];
ll lst[maxn];
ll a[maxn];
ll dfn[maxn];
Edge e[maxn<<1];
ll head[maxn];
ll size[maxn];
ll T,n,etot,dfncnt,ans;

void init() {

std::fill(a+1,a+1+n,0);
    std::fill(head+1,head+1+n,0);
    std::fill(pos+1,pos+1+n,0);
    std::fill(dfn+1,dfn+1+n,0);
    std::fill(a+1,a+1+n,0);
    std::fill(lst+1,lst+1+n,0);
    std::fill(size+1,size+1+n,0);
    std::fill(e+1,e+1+etot,Edge{0,0});
    std::fill(seg+1,seg+1+(n<<2),node{ll(0),ll(0)});
    etot=0;
    dfncnt=0;
    ans=0;
    for(ll i=1;i<=n;++i) {
        vec[i].clear();
    }
    n=0;
}

void debug() {
    for(ll i=1;i<=n;++i) {
        printf("vec[%lld]:",i);
        for(ll j:vec[i]) {
            printf("%lld ",j);
        }
        puts("");
    }
    return;
}

void add_edge(ll u,ll v) {
    ++etot;
    e[etot]={v,head[u]};
    head[u]=etot;
    return;
}

void push_up(ll x) {
    seg[x].maxx=max(seg[ls(x)].maxx,seg[rs(x)].maxx);
    return;
}

void push_down(ll x) {
    if(!seg[x].tag) {
        return;
    }
    seg[ls(x)].tag+=seg[x].tag;
    seg[ls(x)].maxx+=seg[x].tag;
    seg[rs(x)].tag+=seg[x].tag;
    seg[rs(x)].maxx+=seg[x].tag;
    seg[x].tag=0;
    return;
}

//max[x,y]
ll ask(ll p,ll l,ll r,ll x,ll y) {
    if(x<=l && r<=y) {
        return seg[p].maxx;
    }
    ll mid=(l+r)>>1,ret=-inf;
    push_down(p);
    if(x<=mid) {
        ret=ask(ls(p),l,mid,x,y);
    }
    if(y>mid) {
        ret=max(ask(rs(p),mid+1,r,x,y),ret);
    }
    return ret;
}

void upd(ll p,ll l,ll r,ll x,ll y,ll w) {
    if(x<=l && r<=y) {
        seg[p].tag+=w;
        seg[p].maxx+=w;
        return;
    }
    ll mid=(l+r)>>1,ret=-inf;
    push_down(p);
    if(x<=mid) {
        upd(ls(p),l,mid,x,y,w);
    }
    if(y>mid) {
        upd(rs(p),mid+1,r,x,y,w);
    }
    push_up(p);
    return;
}

void dfs(ll u,ll fa) {
    size[u]=1;
    dfn[u]=++dfncnt;
    ll tmp=pos[a[u]];
    lst[u]=pos[a[u]];
    vec[lst[u]].push_back(u);
    pos[a[u]]=u;
    for(ll i=head[u];i;i=e[i].next) {
        if(e[i].v==fa) {
            continue;
        }
        dfs(e[i].v,u);
        size[u]+=size[e[i].v];
    }
    pos[a[u]]=tmp;
    return;
}

void solve(ll u,ll fa) {
    for(ll i=head[u];i;i=e[i].next) {
        if(e[i].v==fa) {
            continue;
        }
        solve(e[i].v,u);
    }
    upd(1,1,n,dfn[u],dfn[u]+size[u]-1,1);
    for(ll i:vec[u]) {
        upd(1,1,n,dfn[i],dfn[i]+size[i]-1,-1);
    }
    std::multiset<ll,std::greater<ll>> set;
    ll tmp=ask(1,1,n,dfn[u],dfn[u]);
    set.insert(ask(1,1,n,dfn[u],dfn[u]));
    for(ll i=head[u];i;i=e[i].next) {
        if(e[i].v==fa) {
            continue;
        }
        set.insert(ask(1,1,n,dfn[e[i].v],dfn[e[i].v]+size[e[i].v]-1));
    }
    if(set.size()==1) {
        ans=max(ans,ll(1));
        return;
    }
    ll fst,sec;
    fst=*set.begin();
    set.erase(set.begin());
    if(set.empty()) {
        return;
    }
    sec=*set.begin();
    ans=max(ans,fst*sec);
    return;
}

void work() {
    cin>>n;

    for(ll i=2;i<=n;++i) {
        ll v;
        cin>>v;
        add_edge(v,i);
    }

    for(ll i=1;i<=n;++i) {
        cin>>a[i];
    }

    dfs(1,0);

    //debug();

    solve(1,0);

    printf("%lld\n",ans);
    init();
    return;
}

int main() {
    freopen("input","r",stdin);

    cin>>T;

    while(T--) {
        work();
    }

    return 0;
}
posted @ 2026-08-04 08:42  txp2025  阅读(2)  评论(0)    收藏  举报