P17141 [NOI 2026] 传送 题解

· · 题解

:::::success[Hint 1] 考虑最终答案形态,在 \Theta(nm\log n) 的复杂度内解决它。 ::::: :::::success[Hint 2] 考虑我们需要求什么什么东西,使用数据结构解决它,并观察性质避免外层二分。 :::::

对应 Hint 1:

我们容易注意到下面几个性质(下面我们将选择随机点到达的点称为随机点,否则称为非随机点):

那么我们枚举一个 E 表示随机点的答案,设 k 表示深度不超过 E 的点的个数,s 表示深度不超过 E 的点的深度和,那么我们有:

E=\frac1n((n-k)E+s)+1

解得 E=\frac{s+n}k,那么我们枚举 E 的整数部分可以做到 \Theta(n^2m)

我们注意到,上面的式子 E 增加右边只会增加 \frac{n-k}{n},那么我们可以直接二分 E 的整数部分,时间复杂度来到 \Theta(nm\log n)

对应 Hint 2:

我们考虑到我们需要维护距离一个点的距离不超过一个值的点的个数及距离和,这个我们可以使用点分治维护,在预处理时维护和每个重心距离不超过某个值的点的个数及距离和以及某个重心的子树中和每个重心距离不超过某个值的点的个数及距离和,这个可以直接 dfs 一遍然后前缀和得到,在每次查询时我们直接从这个点往上在点分树上跳累加答案即可,时间复杂度为 \Theta(n\log n+m\log^3n)

我们考虑将查询距离转化成 \Theta(1),我们使用 \Theta(n\log n)-\Theta(1) 查询 LCA 即可(具体地,我们注意到 \text{LCA}(u,v)=dfn^{-1}_{\min_{i=dfn_u+1}^{dfn_v}dfn_{fa_{dfn^{-1}_i}}}),这时我们的时间复杂度来到了 \Theta(n\log n+m\log^2n)

我们考虑省去外层二分,直观上来讲我们猜测相邻点为根的 \lfloor E\rfloor 不会相差很大。我们首先将上面的等式改写为:

E=1+E-\frac1n\left(\sum_{u=1}^n\max(0,E-\text{dist}(u,y_i))\right)

f_x(E)=\sum_{u=1}^n\max(0,E-\text{dist}(u,x)),设相邻的两个 y_ip,q,期望分别设为 E_p,E_q,由于 \text{dist}(u,p)\le\text{dist}(u,q)+1,那么 \max(0,E_p-\text{dist}(u,p))\ge\max(0,E_p-1-\text{dist}(u,q)),即 f_p(E_p)\ge f_q(E_p-1),由于 f_p(E_p)=f_q(E_q)=nf_q(x) 单调不降,我们有 E_q\ge E_p-1。同理,我们有 E_p\ge E_q-1,即 |E_p-E_q|\le 1

那么我们只在某个点进行二分,其它的点可以直接省去二分。

这样我们就做完了这道题。
时间复杂度为 \Theta(n\log n+m),空间复杂度为 \Theta(n+m)

:::::success[code]

#include<bits/stdc++.h>
//#include "grader.cpp"
#include "teleport.h" 
//#define int long long
#define lll __int128
#define fst ios::sync_with_stdio(0);cin.tie(0);cout.tie(0);
#define rep(i,x,y) for(int i=x;i<=(y);++i) 
#define per(i,x,y) for(int i=x;i>=(y);--i)
#define rpr(i,x,y,z) for(int i=x;i<=(y);i+=z)
#define epe(i,x,y,z) for(int i=x;i>=(y);i-=z)
#define repe(i,x,y) for(i=x;i<=(y);++i) 
#define endl '\n'
#define INF 1e9 
#define pb push_back
#define pob pop_back
#define pf push_front
#define pof pop_front 
#define fi first
#define se second
#define lcm(x,y) x/__gcd(x,y)*y
#define ull unsigned long long
#define prr make_pair
#define pii pair<int,int> 
#define gt(s) getline(cin,s)
#define at(x,y) for(const auto &x:y)
#define ff fflush(stdout)
#define mt(x,y) memset(x,y,sizeof(x))
#define idg isdigit
#define fp(s) string ssss=s;freopen((ssss+".in").c_str(),"r",stdin);freopen((ssss+\
".out").c_str(),"w",stdout);
#define sstr stringstream 
#define all(x) x.begin(),x.end()
#define mcy(a,b) memcpy(a,b,sizeof(b))
#define ui unsigned
#define si signed
#define eb emplace_back
#define pff(x) ((x)*(x))
#define eush emplace
#define double long double
#define pdi pair<double,int>
#ifdef __unix__
#define gc getchar_unlocked
#else
#define gc _getchar_nolock
#endif
#ifdef __unix__
#define pc putchar_unlocked
#else
#define pc _putchar_nolock
#endif
using namespace std;
const int N=5e5+5,M=20;
struct queries{
    int x,y,id;
    queries(){
        x=y=id=0;
    }
    queries(int x,int y,int id):x(x),y(y),id(id){}
};
vector<int>g[N],sumg[N],sum2g[N];
vector<pair<long long,int> >ans;
int fa[N],siz[N],mx[N],dep[N],h[N],cnt,rt,ansk[N],mn[M][N],dfn[N],tim,idfn[N],n;
bool vis[N];
vector<long long>sum[N],sum2[N];
vector<queries>qu[N];
pair<long long,int>now;
void dfssiz(int u,int ft){
    siz[u]=1;
    at(v,g[u])
        if(v!=ft&&!vis[v]){
            dfssiz(v,u);
            siz[u]+=siz[v];
        }
}
void dfsrt(int u,int ft){
    mx[u]=cnt-siz[u];
    at(v,g[u])
        if(v!=ft&&!vis[v]){
            dfsrt(v,u);
            mx[u]=max(mx[u],siz[v]);
        }
    if(mx[u]<mx[rt]) rt=u;
}
void dfssiz2(int u,int ft){
    siz[u]=1;
    h[u]=dep[u]=dep[ft]+1;
    at(v,g[u])
        if(v!=ft&&!vis[v]){
            dfssiz2(v,u);
            siz[u]+=siz[v];
            h[u]=max(h[u],h[v]);
        }
}
void update2(int u,int ft,int p){
    sum2[p][dep[u]]+=dep[u];
    ++sum2g[p][dep[u]];
    at(v,g[u])
        if(v!=ft&&!vis[v]) update2(v,u,p);
}
void update(int u,int ft,int p){
    sum[p][dep[u]]+=dep[u];
    ++sumg[p][dep[u]];
    at(v,g[u])
        if(v!=ft&&!vis[v]) update(v,u,p);
}
void init(int u,int ft){
    dfssiz(u,ft);
    cnt=siz[u];
    rt=0;
    mx[0]=INF;
    dfsrt(u,ft); 
    fa[rt]=ft;
    dep[0]=-1;
    dfssiz2(u,0);
    sum2[rt].resize(h[u]+1);
    sum2g[rt].resize(h[u]+1); 
    update2(u,0,rt); 
//  cout<<u<<' '<<ft<<' '<<rt<<endl;
    rep(i,1,h[u]) sum2[rt][i]+=sum2[rt][i-1];
    rep(i,1,h[u]) sum2g[rt][i]+=sum2g[rt][i-1];
//  rep(i,0,h[u]) cout<<sum2[rt][i]<<' ';
//  cout<<endl;
//  rep(i,0,h[u]) cout<<sum2g[rt][i]<<' ';
//  cout<<endl;
    dfssiz2(rt,0);
    sum[rt].resize(h[rt]+1);
    sumg[rt].resize(h[rt]+1);
    update(rt,0,rt);
    rep(i,1,h[rt]) sum[rt][i]+=sum[rt][i-1];
    rep(i,1,h[rt]) sumg[rt][i]+=sumg[rt][i-1]; 
//  rep(i,0,h[rt]) cout<<sum[rt][i]<<' ';
//  cout<<endl;
//  rep(i,0,h[rt]) cout<<sumg[rt][i]<<' ';
//  cout<<endl<<endl;
    vis[rt]=1;
    int p=rt; 
    at(v,g[rt])
        if(!vis[v]) init(v,p);
}
long long query(int rt,int k){
    if(k<0) return 0;
    if(k>=(int)sum[rt].size()) return sum[rt].back();
    return sum[rt][k];
}
long long query2(int rt,int k){
    if(k<0) return 0;
    if(k>=(int)sum2[rt].size()) return sum2[rt].back();
    return sum2[rt][k];
}
int queryg(int rt,int k){
    if(k<0) return 0;
    if(k>=(int)sumg[rt].size()) return sumg[rt].back();
    return sumg[rt][k];
}
int query2g(int rt,int k){
    if(k<0) return 0;
    if(k>=(int)sum2g[rt].size()) return sum2g[rt].back();
    return sum2g[rt][k];
}
int getlca(int u,int v){
    if(u==v) return u;
    u=dfn[u];
    v=dfn[v];
    if(u>v) swap(u,v);
    ++u;
    int len=__lg(v-u+1);
//  cout<<u<<' '<<v<<' '<<len<<' '<<mn[len][u]<<' '<<mn[len][v-(1ll<<len)+1]<<endl; 
    return idfn[min(mn[len][u],mn[len][v-(1ll<<len)+1])];
}
int getdis(int u,int v){
    return dep[u]+dep[v]-(dep[getlca(u,v)]<<1);
}
pair<long long,int>calc(int k,int rt){
    long long sum=query(rt,k);
    int res=queryg(rt,k),u=rt,ft=fa[rt];
//  cout<<sum<<' '<<res<<' '<<u<<' '<<ft<<' '<<k<<endl;
    while(ft){
        sum-=query2(u,k-getdis(rt,ft)-1)+1ll*query2g(u,k-getdis(rt,ft)-1)*(getdis(rt,ft)+1);
        sum+=query(ft,k-getdis(rt,ft))+1ll*queryg(ft,k-getdis(rt,ft))*getdis(rt,ft);
        res-=query2g(u,k-getdis(rt,ft)-1);
        res+=queryg(ft,k-getdis(rt,ft));
        u=fa[u];
        ft=fa[ft];
//      cout<<sum<<' '<<res<<' '<<u<<' '<<ft<<' '<<k<<endl;
    }
    long long a=sum+n;
    int b=res;
    long long g=__gcd(a,1ll*b);
    a/=g;
    b/=g;
    return prr(a,b);
}
int lower(int l,int r,int rt){
    int ans=-1;
    while(l<=r){
        int mid=l+r>>1;
        pair<long long,int>k=calc(mid,rt);
        if(k.fi>=1ll*mid*k.se){
            ans=mid;
            l=mid+1;
        }else r=mid-1;
    }
    return ans;
}
void dfs(int u,int ft){
    dep[u]=dep[ft]+1;
    dfn[u]=++tim;
    idfn[tim]=u;
    mn[0][dfn[u]]=dfn[ft];
    at(v,g[u])
        if(v!=ft) dfs(v,u);
}
void solveans(int u,int ft){
    at(p,qu[u]){
        int x=p.x,y=p.y,id=p.id;
//      cout<<"  "<<x<<' '<<y<<' '<<id<<' '<<getdis(x,y)<<' '<<getlca(x,y)<<endl;
        if(1ll*getdis(x,y)*now.se<=now.fi) ans[id]=prr(getdis(x,y),1);
        else ans[id]=now;
    }
//  cout<<" "<<u<<' '<<ansk[u]<<endl; 
    at(v,g[u])
        if(v!=ft){
            now=calc(ansk[u],v);
            if(now.fi>=1ll*ansk[u]*now.se){
                pair<long long,int>now2=calc(ansk[u]+1,v);
                if(now2.first>=1ll*(ansk[u]+1)*now2.se){
                    ansk[v]=ansk[u]+1;
                    now=now2;
                }else ansk[v]=ansk[u];
            }else{
                ansk[v]=ansk[u]-1;
                now=calc(ansk[v],v);
            }
//          cout<<" "<<u<<' '<<v<<' '<<ansk[v]<<' '<<now.fi<<' '<<now.se<<endl;
            solveans(v,u);
        }
}
vector<pair<long long,int> >teleport(int c,int nn,int m,vector<int>u,vector<int>v,vector<int>x,vector<int>y){
    n=nn;
    ans.resize(m);
    rep(i,0,n-2) ++u[i];
    rep(i,0,n-2) ++v[i];
    rep(i,0,m-1) ++x[i];
    rep(i,0,m-1) ++y[i];
    rep(i,0,n-2) g[u[i]].pb(v[i]);
    rep(i,0,n-2) g[v[i]].pb(u[i]);
    rep(i,0,m-1) qu[y[i]].eb(x[i],y[i],i);
    init(1,0);
//  rep(i,1,n) cout<<fa[i]<<' ';
//  cout<<endl;
    dep[0]=-1;
    dfs(1,0);
    rep(i,1,__lg(n))
        rep(j,1,n-(1ll<<i)+1) mn[i][j]=min(mn[i-1][j],mn[i-1][j+(1ll<<i-1)]);
    ansk[1]=lower(0,n,1);
    now=calc(ansk[1],1);
    solveans(1,0);
    return ans;
}
/*
0 5 1
0 1
1 2
2 3
1 4
4 3 
*/

:::::