P17141 [NOI 2026] 传送 题解
题目链接:P17141 [NOI 2026] 传送,#18984. 传送
首先随便推推可以看出来对于从 $ x $ 到 $ y $ 的答案等于选取一个包含 $ y $ 的联通块 $ S $,答案为:
\[\min\bigg(dis(x, y), \frac{n+\sum_{u\in S} dis(u, y)}{|S|}\bigg)
\]
然后这个 $ S $ 显然是加入的点距离 $ y $ 远的不如 $ y $ 近的,而且如果选了某个距离的某个点,则距离 $ y $ 为这个距离的所有点都会选(等价于和一个更小的数取加权平均数),因此问题转化为了对于给定点 $ u $ 求距离 $ u $ 小于等于 $ x $ 的个数和距离和,要求做到 $ O(logn) $。考虑点分树处理即可,点分树的大致思想是重构成每次取子树重心往父亲节点连边,然后高度是 $ O(logn) $ 的,然后保存当前重心的联通块的信息和贡献到其父亲的信息,这样对于一个点 $ v $,在前面的点分树点上面不会统计到他,因为当前联通块不包含他,在当前联通块包含他但是子节点不包含他的时候恰好会贡献,在之后的节点因为当前联通块包含他且子节点的联通块也包含他,相减会消去,因此只会贡献一次。
此时时间复杂度是 $ O(nlog^2n) $ 的,可能卡卡能过去,但是我是大常数 vector 选手,所以继续观察,发现相邻节点选的联通块 $ S $ 包含的节点的最大深度的差不超过 1,因此只需要对根节点进行一次二分,其他可以 $ O(1) $ check,总复杂度 $ O(nlogn) $,居然还跑了 3s,常数多大可想而知。
代码:
#include<bits/stdc++.h>
#define ll long long
#define i128 __int128
#define pb push_back
using namespace std;
int n,tot,dep[500005],dfn[500005],st[20][500005],sz[500005],cfa[500005],rad[500005];
bool vis[500005];
vector<int> e[500005];
struct node{
vector<int> cnt;
vector<ll> sum;
}all[500005],sub[500005];
pair<ll,int> val[500005];
ll gcd(ll x,ll y)
{
while(y)
{
x%=y;
swap(x,y);
}
return x;
}
int low(int x,int y)
{
return dep[x]<dep[y]?x:y;
}
void dfslca(int u,int fa)
{
dfn[u]=++tot;
st[0][tot]=fa;
for(auto v:e[u])
{
if(v==fa)continue;
dep[v]=dep[u]+1;
dfslca(v,u);
}
}
int getlca(int u,int v)
{
if(u==v)return u;
int l=dfn[u],r=dfn[v];
if(l>r)swap(l,r);
++l;
int k=__lg(r-l+1);
return low(st[k][l],st[k][r-(1<<k)+1]);
}
int get2dis(int u,int v)
{
return dep[u]+dep[v]-2*dep[getlca(u,v)];
}
void getsz(int u,int fa)
{
sz[u]=1;
for(auto v:e[u])
{
if(v==fa||vis[v])continue;
getsz(v,u);
sz[u]+=sz[v];
}
}
int getmid(int u,int fa,int sum)
{
for(auto v:e[u])
{
if(v!=fa&&!vis[v]&&sz[v]*2>sum)return getmid(v,u,sum);
}
return u;
}
void getdis(int u,int fa,int d,vector<int> &vc)
{
vc.pb(d);
for(auto v:e[u])
{
if(v==fa||vis[v])continue;
getdis(v,u,d+1,vc);
}
}
void getinfo(node &ret,vector<int> &vc)
{
int mx=0;
for(auto x:vc)
{
mx=max(mx,x);
}
ret.cnt.assign(mx+1,0);
ret.sum.assign(mx+1,0);
for(auto x:vc)
{
ret.cnt[x]++;
}
ll cnt=0,sum=0;
for(int i=0;i<=mx;i++)
{
cnt+=ret.cnt[i];
sum+=1ll*i*ret.cnt[i];
ret.cnt[i]=cnt;
ret.sum[i]=sum;
}
}
int build(int rt,int fa)
{
getsz(rt,-1);
int mid=getmid(rt,-1,sz[rt]);
cfa[mid]=fa;
vector<int> dis;
getdis(mid,-1,0,dis);
getinfo(all[mid],dis);
vis[mid]=1;
for(auto v:e[mid])
{
if(vis[v])continue;
vector<int> cur;
getdis(v,mid,1,cur);
node tmp;
getinfo(tmp,cur);
int son=build(v,mid);
sub[son]=tmp;
}
return mid;
}
void merge(node &a,int lim,int m,int op,ll &cnt,ll &sum)
{
if(lim<0||!a.cnt.size())return ;
lim=min(lim,(int)a.cnt.size()-1);
cnt+=1ll*op*a.cnt[lim];
sum+=1ll*op*(a.sum[lim]+1ll*a.cnt[lim]*m);
}
pair<ll,int> query(int u,int r)
{
ll cnt=0,sum=n,son=-1;
for(int mid=u;mid!=-1;son=mid,mid=cfa[mid])
{
int d=get2dis(u,mid);
int lim=r-d;
if(lim<0)continue;
merge(all[mid],lim,d,1,cnt,sum);
if(son!=-1)merge(sub[son],lim,d,-1,cnt,sum);
}
return {sum,(int)cnt};
}
bool cmp(pair<ll,int> a,pair<ll,int> b)
{
return (i128)a.first*b.second<(i128)b.first*a.second;
}
bool cmp1(pair<ll,int> a,pair<ll,int> b)
{
return (i128)a.first*b.second<=(i128)b.first*a.second;
}
int getr(int u)
{
int l=0,r=n-1;
while(l<r)
{
int mid=l+r>>1;
if(cmp(query(u,mid+1),query(u,mid)))l=mid+1;
else r=mid;
}
return l;
}
void solve(int u,int fa)
{
for(auto v:e[u])
{
if(v==fa)continue;
int d=rad[u];
pair<ll,int> now=query(v,d);
if(d>0)
{
pair<ll,int> pre=query(v,d-1);
if(cmp1(pre,now))
{
d--;
now=pre;
}
else if(d+1<n)
{
pair<ll,int> nxt=query(v,d+1);
if(cmp(nxt,now))
{
d++;
now=nxt;
}
}
}
else if(d+1<n)
{
pair<ll,int> nxt=query(v,d+1);
if(cmp(nxt,now))
{
d++;
now=nxt;
}
}
rad[v]=d;
val[v]=now;
solve(v,u);
}
}
vector<pair<ll,int> > teleport(int c,int _n,int m,vector<int> u,vector<int> v,vector<int> x,vector<int> y)
{
n=_n;
for(int i=0;i<n-1;i++)
{
e[u[i]].pb(v[i]);
e[v[i]].pb(u[i]);
}
dfslca(0,0);
for(int i=1;i<20;i++)
{
for(int j=1;j+(1<<i)-1<=n;j++)
{
st[i][j]=low(st[i-1][j],st[i-1][j+(1<<(i-1))]);
}
}
build(0,-1);
rad[0]=getr(0);
val[0]=query(0,rad[0]);
solve(0,-1);
vector<pair<ll,int> > ans(m);
for(int i=0;i<m;i++)
{
int d=get2dis(x[i],y[i]);
pair<ll,int> p=val[y[i]];
if((i128)d*p.second<p.first)ans[i]={d,1};
else
{
int gd=gcd(p.first,p.second);
ans[i]={p.first/gd,p.second/gd};
}
}
return ans;
}

浙公网安备 33010602011771号