P17141 [NOI 2026] 传送 题解
题目描述
给定一棵 \(n\) 个点的树,每花费 \(1\) 单位时间,可以沿一条边移动,或等概率传送到 \(n\) 个点。
\(m\) 次询问,给定起点和终点 \((x,y)\) ,求 \(x\to y\) 的期望最短用时。
数据范围
- \(1\le n\le 5\cdot 10^5,1\le m\le 10^6\) 。
时间限制 \(\texttt{3.5s}\) ,空间限制 \(\texttt{1GB}\) 。
分析
固定终点 \(y\) ,假设走传送门期望用时为 \(T\) ,显然到 \(y\) 距离不超过 \(T\) 的点会直接走到 \(y\) ,否则走传送门。
记 \(\lfloor T\rfloor=i\) , \(dis(j,y)\le T\) 的点有 \(a_i\) 个,距离和为 \(b_i\) ,根据上式可算出 \(T=\frac{b_i+n}{a_i}\) 。
二分 \(i\) ,如果 \(T\ge i+1\) 则表明还有优化空间(让距离 \(i+1\) 的点从走传送门改为直达可以得到更优策略)。
建立点分树,每个点用 vector 统计点数 & 前缀和,即可在 \(\mathcal O(\log n)\) 的时间内对给定的 \(i,y\) 计算出 \(a_i,b_i\) (将 lca 计算代价视为 \(\mathcal O(1)\) )。
记终点 \(y\) 的答案为 \(T_y\) ,我们希望预处理所有 \(T_y\) ,从而询问直接输出 \(\min(dis(x,y),T_y)\) 。
注意到如果 \(u,v\) 相邻且 \(|T_u-T_v|\gt 1\) ,我们可以先用 \(\min(T_u,T_v)\) 的时间走到其中一个点,再用 \(1\) 单位时间走过去,矛盾。因此 \(|T_u-T_v|\le 1\) 。
预处理一个 \(T_u\) ,只需要对 \(i-1,i,i+1\) 分别做一次查询即可确定 \(T_v\) ,时间复杂度 \(\mathcal O(n\log n+m)\) 。
#include<bits/stdc++.h>
#include"teleport.h""
#define ll long long
#define fi first
#define se second
#define mp make_pair
#define pii pair<ll,int>
using namespace std;
const int maxn=5e5+5;
const ll inf=1e18;
int n,rt,all;
int d[maxn],fa[maxn],son[maxn],top[maxn];
int p[maxn],mx[maxn],sz[maxn],r1[maxn],r2[maxn];
int a[maxn];ll b[maxn];
bool vis[maxn];
pii t[maxn];
vector<int> g[maxn];
vector<pii> s1[maxn],s2[maxn];
// s1[u][i] = \sum_{v in subtree(u), d[v]\le i} pair(dis(u,v),1)
// s2[u][i] = \sum_{v in subtree(u), d[v]\le i} pair(dis(p[u],v),1)
pii operator+(const pii &a,const pii &b) {return mp(a.fi+b.fi,a.se+b.se);}
pii operator-(const pii &a,const pii &b) {return mp(a.fi-b.fi,a.se-b.se);}
void operator+=(pii &a,const pii &b) {a.fi+=b.fi,a.se+=b.se;}
void operator-=(pii &a,const pii &b) {a.fi-=b.fi,a.se-=b.se;}
bool cmp(const pii &a,const pii &b) {return a.fi*b.se<b.fi*a.se;}
void dfs1(int u,int f)
{
sz[u]=1,a[d[u]-1]++,b[d[u]-1]+=d[u]-1;
for(auto v:g[u])
{
if(v==f) continue;
d[v]=d[u]+1,fa[v]=u;
dfs1(v,u),sz[u]+=sz[v];
if(sz[v]>=sz[son[u]]) son[u]=v;
}
}
void dfs2(int u,int topf)
{
top[u]=topf;
if(son[u]) dfs2(son[u],topf);
for(auto v:g[u])
{
if(v==fa[u]||v==son[u]) continue;
dfs2(v,v);
}
}
int lca(int u,int v)
{
while(top[u]!=top[v])
{
if(d[top[u]]<d[top[v]]) swap(u,v);
u=fa[top[u]];
}
return d[u]<d[v]?u:v;
}
int getdis(int u,int v)
{
return d[u]+d[v]-2*d[lca(u,v)];
}
void getroot(int u,int fa)
{
sz[u]=1,mx[u]=0;
for(auto v:g[u])
{
if(vis[v]||v==fa) continue;
getroot(v,u);
sz[u]+=sz[v],mx[u]=max(mx[u],sz[v]);
}
mx[u]=max(mx[u],all-sz[u]);
if(!rt||mx[u]<mx[rt]) rt=u;
}
int dfs3(int u,int fa,int rt,int dep)
{
sz[u]=1;
r1[rt]=max(r1[rt],dep);
if(p[rt]) r2[rt]=max(r2[rt],getdis(p[rt],u));
for(auto v:g[u]) if(!vis[v]&&v!=fa) sz[u]+=dfs3(v,u,rt,dep+1);
return sz[u];
}
void solve(int u)
{
vis[u]=true,dfs3(u,0,u,0);
for(auto v:g[u])
{
if(vis[v]) continue;
all=sz[v],getroot(v,rt=0),p[rt]=u,solve(rt);
}
}
pii query(int i,int y)
{
if(i<=0) return mp(inf,1);
pii res=mp(0,0);
for(int u=y;u;u=p[u])
{
int d=getdis(u,y);
auto tmp=i-d>=0?s1[u][min(i-d,r1[u])]:mp(0ll,0);
res+=mp(tmp.fi+1ll*d*tmp.se,tmp.se);
if(!p[u]) continue;
d=getdis(p[u],y);
tmp=i-d>=0?s2[u][min(i-d,r2[u])]:mp(0ll,0);
res-=mp(tmp.fi+1ll*d*tmp.se,tmp.se);
}
return mp(res.fi+n,res.se);
}
void dfs4(int u)
{
for(auto v:g[u])
{
if(v==fa[u]) continue;
int i=t[u].fi/t[u].se;
t[v]=min(min(query(i-1,v),query(i,v),cmp),query(i+1,v),cmp);
dfs4(v);
}
}
vector<pii> 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++) u[i]++,v[i]++,g[u[i]].push_back(v[i]),g[v[i]].push_back(u[i]);
d[1]=1,dfs1(1,0),dfs2(1,1);
all=n,getroot(1,0),solve(rt);
for(int i=1;i<=n;i++) t[i]=mp(inf,1),s1[i].resize(r1[i]+1),s2[i].resize(r2[i]+1);
for(int i=1;i<=n;i++) b[i]+=b[i-1],a[i]+=a[i-1],t[1]=min(t[1],mp(b[i]+n,a[i]),cmp);
for(int i=1;i<=n;i++)
for(int d,u=i;u;u=p[u])
{
d=getdis(u,i),s1[u][d]+=mp(d,1);
if(p[u]) d=getdis(p[u],i),s2[u][d]+=mp(d,1);
}
for(int i=1;i<=n;i++)
{
for(int j=1;j<=r1[i];j++) s1[i][j]+=s1[i][j-1];
for(int j=1;j<=r2[i];j++) s2[i][j]+=s2[i][j-1];
}
dfs4(1);
for(int i=1;i<=n;i++)
{
ll g=__gcd(t[i].fi,(ll)t[i].se);
t[i].fi/=g,t[i].se/=g;
}
vector<pii> res(m);
for(int i=0;i<m;i++)
{
int d=getdis(++x[i],++y[i]),j=y[i];
res[i]=t[j].fi<=1ll*d*t[j].se?t[j]:mp((ll)d,1);
}
return res;
}
本文来自博客园,作者:peiwenjun,转载请注明原文链接:https://www.cnblogs.com/peiwenjun/p/22282558
浙公网安备 33010602011771号