题解:P15761 [JAG 2025 Summer Camp #1] Colored Tree and Path

· · 题解

题解怎么全是根号或者 \log^2 之类的东西。这篇题解是 O(n \log n) 的。

考虑主席树,每个结点从父亲结点的版本更新,额外加上当前结点的值,也就是说每个结点的版本维护的是到根路径上的值的权值线段树。

对于路径,设结点 u 的版本为 f_u,则路径 (u,v) 可以表示为 f_u+f_v-f_{\textup{lca}(u,v)}-f_{\textup{fa}(\textup{lca}(u,v))}

现在差不多做完了,线段树上二分即可:若左子树中两条路径全部值相等,则递归右子树,否则递归左子树,这样就能找到第一个不相等的,减 1 即可。但是线段树没办法判断整棵子树相等,怎么办呢?

考虑哈希,每个不同的 c_i 维护一个不同的哈希值,线段树上每个结点的哈希值就是子树中哈希值之和。判断哈希是否相等即可。

代码不建议看,比较史,特别是因为写了个双哈希。 :::success[code]

#include<bits/stdc++.h>
using namespace std;
#define int long long
#define ui unsigned int
#define fi first
#define se second
#define pii pair<int,int>
#define lowbit(x) ((x)&(-(x)))
#define popc(x) __builtin_popcountll(x)
#define ctz(x) __builtin_ctzll(x)
#define clz(x) __builtin_clzll(x)
#define double long double
#define sqrt(x) sqrtl(x)
#define pow(x,y) powl(x,y)
#define cbrt(x) cbrtl(x)
#define sin(x) sinl(x)
#define cos(x) cosl(x)
#define tan(x) tanl(x)
#define push emplace
#define pb emplace_back
#define pf emplace_front
const int N=2e5+10,mod=999911659;
const int base=1657,base2=167;
ui pw[N];
int pw2[N];
struct persgt
{
    struct node
    {
        int ls,rs;
        int sum;
        ui hsh;
        int hsh2;
    }tr[N<<5];
    #define ls(u) tr[u].ls
    #define rs(u) tr[u].rs
    int rt[N],utot;
    int& operator[](int x){return rt[x];}
    void pushup(int u)
    {
        tr[u].sum=tr[ls(u)].sum+tr[rs(u)].sum;
        tr[u].hsh=tr[ls(u)].hsh+tr[rs(u)].hsh;
        tr[u].hsh2=(tr[ls(u)].hsh2+tr[rs(u)].hsh2)%mod;
    }
    int copy(int u)
    {
        tr[++utot]=tr[u];
        return utot;
    }
    void modify(int& u,int v,int l,int r,int p)
    {
        u=copy(v);
        if(l==r)
        {
            tr[u].sum++;
            tr[u].hsh+=pw[l];
            tr[u].hsh2=(tr[u].hsh2+pw2[l])%mod;
            return;
        }
        int mid=l+r>>1;
        if(p<=mid) modify(ls(u),ls(v),l,mid,p);
        else modify(rs(u),rs(v),mid+1,r,p);
        pushup(u);
    }
    int query(int u1,int v1,int l1,int f1,int u2,int v2,int l2,int f2,int l,int r)
    {
        int ss=tr[u1].sum+tr[v1].sum-tr[l1].sum-tr[f1].sum-tr[u2].sum-tr[v2].sum+tr[l2].sum+tr[f2].sum;
        int sh=tr[u1].hsh+tr[v1].hsh-tr[l1].hsh-tr[f1].hsh-tr[u2].hsh-tr[v2].hsh+tr[l2].hsh+tr[f2].hsh;
        int sh2=(tr[u1].hsh2+tr[v1].hsh2-tr[l1].hsh2-tr[f1].hsh2-tr[u2].hsh2-tr[v2].hsh2+tr[l2].hsh2+tr[f2].hsh2)%mod;
        if(!ss&&!sh&&!sh2) return r+1;
        if(l==r) return l;
        ss=tr[ls(u1)].sum+tr[ls(v1)].sum-tr[ls(l1)].sum-tr[ls(f1)].sum-tr[ls(u2)].sum-tr[ls(v2)].sum+tr[ls(l2)].sum+tr[ls(f2)].sum;
        sh=tr[ls(u1)].hsh+tr[ls(v1)].hsh-tr[ls(l1)].hsh-tr[ls(f1)].hsh-tr[ls(u2)].hsh-tr[ls(v2)].hsh+tr[ls(l2)].hsh+tr[ls(f2)].hsh;
        sh2=(tr[ls(u1)].hsh2+tr[ls(v1)].hsh2-tr[ls(l1)].hsh2-tr[ls(f1)].hsh2-tr[ls(u2)].hsh2-tr[ls(v2)].hsh2+tr[ls(l2)].hsh2+tr[ls(f2)].hsh2)%mod;
        int mid=l+r>>1;
        if(!ss&&!sh&&!sh2) return query(rs(u1),rs(v1),rs(l1),rs(f1),rs(u2),rs(v2),rs(l2),rs(f2),mid+1,r);
        else return query(ls(u1),ls(v1),ls(l1),ls(f1),ls(u2),ls(v2),ls(l2),ls(f2),l,mid);
    }
    #undef ls
    #undef rs
}tr;
int fa[N],siz[N],dep[N],hs[N],top[N];
int a[N];
vector<int>adj[N];
int n;
void dfs1(int u,int f)
{
    fa[u]=f;
    siz[u]=1;
    dep[u]=dep[f]+1;
    tr.modify(tr[u],tr[f],1,n,a[u]);
    for(auto v:adj[u])
    {
        if(v==f) continue;
        dfs1(v,u);
        siz[u]+=siz[v];
        if(siz[v]>siz[hs[u]]) hs[u]=v;
    }
}
void dfs2(int u,int t)
{
    top[u]=t;
    if(hs[u]) dfs2(hs[u],t);
    for(auto v:adj[u])
    {
        if(v==fa[u]||v==hs[u]) continue;
        dfs2(v,v);
    }
}
int lca(int x,int y)
{
    while(top[x]!=top[y])
    {
        if(dep[top[x]]<dep[top[y]]) swap(x,y);
        x=fa[top[x]];
    }
    return dep[x]<dep[y]?x:y;
}
signed main()
{
    freopen("path.in","r",stdin);
    freopen("path.out","w",stdout);
    ios::sync_with_stdio(0);
    cin.tie(0),cout.tie(0);
    cin>>n;
    pw[0]=pw2[0]=1;
    for(int i=1;i<=n;i++) pw[i]=pw[i-1]*base,pw2[i]=pw2[i-1]*base2%mod;
    for(int i=1;i<n;i++)
    {
        int u,v;
        cin>>u>>v;
        adj[u].pb(v),adj[v].pb(u);
    }
    for(int i=1;i<=n;i++) cin>>a[i];
    dfs1(1,0);
    dfs2(1,1);
    int q;
    cin>>q;
    while(q--)
    {
        int u1,v1,u2,v2;
        cin>>u1>>v1>>u2>>v2;
        int l1=lca(u1,v1),f1=fa[l1],l2=lca(u2,v2),f2=fa[l2];
        cout<<tr.query(tr[u1],tr[v1],tr[l1],tr[f1],tr[u2],tr[v2],tr[l2],tr[f2],1,n)-1<<'\n';
    }
    return 0; 
}

:::