3 条题解

  • 28
    @ 2025-3-25 16:35:05

    主播主播,你的树剖确实强悍,但是太难操作了,有没有算法简单,而且时间复杂度还优的做法呢?

    有的兄弟,有的,这里是一篇单 log\log 的题解!

    考虑并查集啊,只需要维护树上每个节点向上走,离它最近的没被走过的节点就好了。

    你发现你每次向上跳只会跳没被走过的点,于是每个点至多被跳到一次,所以时间复杂度是 O(nlogn)O(n\log n) 的。

    #include <iostream>
    #include <vector>
    #define ll int
    using namespace std;
    const ll M=4e5+10;
    const ll N=5e5+10;
    inline void read(ll &x){
        x=0;bool flag(0);char ch=getchar();
        while(!isdigit(ch)) flag=ch=='-',ch=getchar();
        while(isdigit(ch)) x=(x<<1)+(x<<3)+(ch^48),ch=getchar();
        flag?x=-x:0;
    }
    ll dep[N],zg[N],stt;
    ll findf(ll x){
        if(zg[x]==x) return x;
        return zg[x]=findf(zg[x]);
    }
    void _merge(ll x,ll y){
        ll fx=findf(x),fy=findf(y);
        if(fx!=fy) zg[fy]=fx; 
    }
    ll fa[N][20],n,m,lin[M];
    vector<ll> A[N];
    void init(){for(ll i=1;i<=n;i++) zg[i]=i;}
    void dfs(ll u,ll f){
        fa[u][0]=f;dep[u]=dep[f]+1;
        for(ll i=1;i<20;i++){
            fa[u][i]=fa[fa[u][i-1]][i-1];
            if(fa[u][i]==0) break;
        }
        ll rt=A[u].size();
        for(ll i=0;i<rt;i++){
            if(A[u][i]==f) continue;
            dfs(A[u][i],u);
        }
    }
    ll lca(ll a,ll b){
        if(dep[a]<dep[b]) swap(a,b);
        for(ll i=19;i>=0;i--) if(dep[fa[a][i]]>=dep[b]) a=fa[a][i];
        if(dep[a]>dep[b]) a=fa[a][0];
        for(ll i=19;i>=0;i--) if(fa[a][i]!=fa[b][i]) a=fa[a][i],b=fa[b][i];
        if(a!=b) a=fa[a][0],b=fa[b][0];
        return a;
    }
    bool vis[N];
    void biaoji(ll from,ll to){
        ll now=from;
        while(dep[now]>=dep[to]){
            vis[now]=1; 
            ll nxt=fa[findf(now)][0];
            _merge(to,now);
            now=nxt;
        }
    }
    int main(){
        read(n),read(m),read(stt);
        init();
        for(ll i=1;i<n;i++){
            ll x,y;
            read(x),read(y);
            A[x].push_back(y);
            A[y].push_back(x);
        }
        for(ll i=1;i<=m;i++) read(lin[i]);
        dfs(1,0);
        long long ans=0;
        vis[stt]=1;
        for(ll i=1;i<=m;i++){
            if(vis[lin[i]]) continue;
            ll lc=lca(stt,lin[i]);
            vis[stt]=vis[lc]=vis[lin[i]]=1;
            ans+=(dep[stt]-dep[lc])+(dep[lin[i]]-dep[lc]);
            biaoji(stt,lc);
            biaoji(lin[i],lc);
            stt=lin[i];
        }
        cout<<ans;
        return 0; 
    }
    

    信息

    ID
    106
    时间
    1000ms
    内存
    256MiB
    难度
    8
    标签
    (无)
    递交数
    75
    已通过
    10
    上传者