P5984 [PA 2019] Podatki drogowe 题解
首先我们需要能对两个序列进行比较,我们可以将每一个权值看作一个 \(n\) 进制的数,我们发现我们实际上根本不会进位,所以可以直接对于每一个 \(n^i\) 维护其个数,然后比较就可以直接使用哈希+二分来比较。
然后我们要想一个办法来比较任意两条路径的大小,不难发现我们可以先进行点分治,然后对于每一个点,将其从该点到根的路径上的所有边权建一棵权值线段树来维护每一种边权出现数量的哈希值,由于每个点从 fa 到该点的增量很少,所以可以考虑使用主席树。然后单次比较就可以在主席树上二分做到单 log 比较。
然后我们可以考虑对于一条路径的权值来确定其排名,我们发现我们可以先对于分治到的每一个节点 \(u\) 连通块中的节点 \(v\) 按照到路径 \(u\) 到 \(v\) 的权值大小排序,然后确定排名时可以对于每一个连通块使用双指针找到所有小当前路径权值的路径个数。排序的时间复杂度为 \(O(n\log^3n)\) ,双指针时间复杂度是 \(O(n\log^2 n)\) 。
由于我们要找排名为 \(k\) 的路径,所以可以想到使用二分,但由于我们的值域非常大,所以无法对值域进行二分,我们可以考虑一种做法叫做随机二分,就是每一次在当前合法路径集合中随机选择一条路径的权值作为 \(mid\) ,然后按照正常二分过程进行二分,期望二分次数为 \(O(\log n)\) 。乘上每一次确定排名的时间复杂度 \(O(n\log^3 n)\) ,变成了 \(O(n\log^4 n)\) ,不是很能通过,但是我们发现瓶颈在于排序,但排序完全没有必要在二分里面进行,可以直接移动到循环外面,时间复杂度就变为了 \(O(n\log^3 n)\) ,可以通过。
Tips:
- 我们发现一个分治节点 \(u\) 和当前连通块中的 \(v\) 形成了一个点对,每次查找这个点对过于麻烦,所以可以直接将每一个点对对影成一个数字方便维护。
- 对于当前合法路径集合,我们发现在同一个分治节点下,一个节点所对应的另一个节点一定在排好序后的区间中连续,所以可以直接在排好序的序列中维护当前合法区间。
- 二分时候要注意对于权值相同的路径的写法,我们可以考虑处理出来小于当前权值的个数 \(cnt1\) 和小于等于当前权值的个数 \(cnt2\) ,然后当 \(k\le cnt1\) 时往小找, \(cnt1<k\le cnt2\) 时可以直接确定答案, \(k>cnt2\) 时往大找,这样可以保证时间复杂度正确。
代码:
#include<bits/stdc++.h>
using namespace std;
#define il inline
//#define int long long
#define ll long long
bool St;
struct IO
{
static const int Size=(1<<21);
char buf[Size],*p1,*p2;
int st[105],Top;
~IO(){clear();}
il void clear(){fwrite(buf,1,Top,stdout);Top=0;}
il char gc(){return p1==p2&&(p2=(p1=buf)+fread(buf,1,Size,stdin),p1==p2)?EOF:*p1++;}
il void pc(const char c){Top==Size&&(clear(),0);buf[Top++]=c;}
il IO& operator >>(char& c){while(c=gc(),c==' ' || c=='\n' || c=='\r');return *this;}
template<typename T>il IO& operator >>(T& x)
{
x=0;bool f=0;char c=gc();
while(!isdigit(c)){if(c=='-') f=1;c=gc();}
while(isdigit(c)){x=(x<<1)+(x<<3)+(c^48);c=gc();}
f?x=-x:0;
return *this;
}
il IO& operator >>(string& s)
{
s="";char c=gc();
while(c==' ' || c=='\n' || c=='\r') c=gc();
while(c!=' ' && c!='\n' && c!='\r' && c!=EOF) s+=c,c=gc();
return *this;
}
il IO& operator <<(const char c){pc(c);return *this;}
template<typename T> il IO& operator <<(T x)
{
if(x<0) pc('-'),x=-x;
do st[++st[0]]=x%10,x/=10;while(x);
while(st[0]) pc(st[st[0]--]+'0');
return *this;
}
il IO& operator <<(const string s){for(auto c:s) pc(c);return *this;}
il IO& operator <<(const char* c){for(int i=0;c[i];i++) pc(c[i]);return *this;}
} fin,fout;
const int p=13331,MOD=998244353;
mt19937_64 rnd(time(0)*13331);
const int N=25010,mod=1e9+7;
int totedge=1,h[N],n,k,vis[N],ctr,mn,sum,siz[N],tot,rt[N<<4],l[N],r[N],col[N<<5],id[N<<5],tmp[N<<5],nowu,nowv,p1[N<<5],p2[N<<5],Len[N<<5],ss;;
ll powp[N],fpow[N];
int liml[N<<5],limr[N<<5];
unordered_map<int,int> mp[N];
vector<int> idx[N<<5];
il int qpow(ll x,int y,int pmod)
{
ll ret=1;
while(y)
{
if(y&1) ret=ret*x%pmod;
x=x*x%pmod;y>>=1;
}
return ret;
}
struct edge
{
int to,nxt,w;
} e[N<<1];
struct Tree
{
struct node
{
int lid,rid;
ll sum;
} t[N<<8];
int tot=0;
il void push_up(int x)
{
t[x].sum=(t[t[x].lid].sum+t[t[x].rid].sum)%MOD;
}
il int add(int x,int L,int R,int iii)
{
int ID=++tot;
t[ID]=t[x];
if(L==R)
{
t[ID].sum=(t[ID].sum+powp[iii])%MOD;
return ID;
}
int mid=(L+R)>>1;
if(iii<=mid) t[ID].lid=add(t[ID].lid,L,mid,iii);
else t[ID].rid=add(t[ID].rid,mid+1,R,iii);
push_up(ID);
return ID;
}
il bool compare(int u,int v,int L,int R)
{
if(L==R) return t[u].sum*fpow[L]%MOD<t[v].sum*fpow[L]%MOD;
int mid=(L+R)>>1;
if(t[t[u].rid].sum!=t[t[v].rid].sum) return compare(t[u].rid,t[v].rid,mid+1,R);
else return compare(t[u].lid,t[v].lid,L,mid);
}
il bool compare2(int u1,int u2,int v1,int v2,int L,int R)
{
if(L==R) return (t[u1].sum+t[u2].sum)*fpow[L]%MOD<(t[v1].sum+t[v2].sum)*fpow[L]%MOD;
int mid=(L+R)>>1;
if((t[t[u1].rid].sum+t[t[u2].rid].sum)%MOD!=(t[t[v1].rid].sum+t[t[v2].rid].sum)%MOD) return compare2(t[u1].rid,t[u2].rid,t[v1].rid,t[v2].rid,mid+1,R);
else return compare2(t[u1].lid,t[u2].lid,t[v1].lid,t[v2].lid,L,mid);
}
il int query(int u,int v,int L,int R)
{
if(L==R)
{
return ((t[u].sum+t[v].sum)*fpow[L]%MOD)*qpow(n,L,mod)%mod;
}
int mid=(L+R)>>1;
return (query(t[u].lid,t[v].lid,L,mid)+query(t[u].rid,t[v].rid,mid+1,R))%mod;
}
} T;
il void add_edge(int u,int v,int w)
{
totedge++;
e[totedge].to=v;
e[totedge].nxt=h[u];
e[totedge].w=w;
h[u]=totedge;
}
il void get_ctr(int u,int fa)
{
siz[u]=1;
int mx=0;
for(int i=h[u];i;i=e[i].nxt)
{
int v=e[i].to;
if(v==fa || vis[v]) continue;
get_ctr(v,u);
siz[u]+=siz[v];
mx=max(mx,siz[v]);
}
mx=max(mx,sum-siz[u]);
if(mx<mn) ctr=u,mn=mx;
}
il void get_mp(int u,int fa)
{
mp[ctr][u]=++tot;
for(int i=h[u];i;i=e[i].nxt)
{
int v=e[i].to;
if(v==fa || vis[v]) continue;
get_mp(v,u);
}
}
il void init_mp(int u)
{
ss+=siz[u];
vis[u]=1;
ctr=u;
l[ctr]=tot+1;
get_mp(u,0);
r[ctr]=tot;
for(int i=h[u];i;i=e[i].nxt)
{
int v=e[i].to;
if(vis[v]) continue;
sum=mn=siz[v];
get_ctr(v,0);
get_ctr(ctr,0);
init_mp(ctr);
}
}
il void dfs(int u,int fa)
{
int idu=mp[ctr][u];
for(int i=h[u];i;i=e[i].nxt)
{
int v=e[i].to,idv=mp[ctr][v];
if(v==fa || vis[v]) continue;
rt[idv]=T.add(rt[idu],1,n,e[i].w);
if(!fa) col[idv]=idv;
else col[idv]=col[idu];
idx[col[idv]].push_back(idv);
dfs(v,u);
}
}
il void calc(int u)
{
vis[u]=1;
ctr=u;
int idu=mp[ctr][ctr];
col[idu]=idu;
dfs(u,0);
sort(id+l[ctr],id+r[ctr]+1,[](int x,int y){return T.compare(rt[x],rt[y],1,n);});
for(int i=l[ctr];i<=r[ctr];i++) tmp[id[i]]=i;
for(int i=l[ctr];i<=r[ctr];i++)
{
for(int& j:idx[i]) j=tmp[j];
sort(idx[i].begin(),idx[i].end());
liml[i]=l[ctr],limr[i]=i-1;
}
for(int i=h[u];i;i=e[i].nxt)
{
int v=e[i].to;
if(vis[v]) continue;
sum=mn=siz[v];
get_ctr(v,0);
get_ctr(ctr,0);
calc(ctr);
}
}
il void solve()
{
ll s=0;
for(int i=1;i<=tot;i++)
{
Len[i]=0;
if(liml[i]<=limr[i])
{
Len[i]+=limr[i]-liml[i]+1;
Len[i]-=upper_bound(idx[col[i]].begin(),idx[col[i]].end(),limr[i])-lower_bound(idx[col[i]].begin(),idx[col[i]].end(),liml[i]);
}
s+=Len[i];
}
ll val=rnd()%s+1;s=0;
for(int i=1;i<=tot;i++)
{
if(s+Len[i]>=val)
{
nowu=id[i];
int ttt=limr[i];
while(s<val)
{
while(col[ttt]==col[i]) ttt--;
ttt--;
s++;
}
ttt++;
nowv=id[ttt];
break;
}
else s+=Len[i];
}
ll cnt1=0,cnt2=0;
for(int i=1;i<=n;i++)
{
int pos=r[i],pos2=r[i];
for(int j=l[i];j<=r[i];j++)
{
while(pos>=l[i] && !T.compare2(rt[id[pos]],rt[id[j]],rt[nowu],rt[nowv],1,n)) pos--;
while(pos2>=l[i] && T.compare2(rt[nowu],rt[nowv],rt[id[pos2]],rt[id[j]],1,n)) pos2--;
p1[j]=min(limr[j],pos),p2[j]=min(pos2,limr[j]);
if(p1[j]>=liml[j])
{
cnt1+=p1[j]-liml[j]+1;
cnt1-=upper_bound(idx[col[j]].begin(),idx[col[j]].end(),p1[j])-lower_bound(idx[col[j]].begin(),idx[col[j]].end(),liml[j]);
}
if(p2[j]>=liml[j])
{
cnt2+=p2[j]-liml[j]+1;
cnt2-=upper_bound(idx[col[j]].begin(),idx[col[j]].end(),p2[j])-lower_bound(idx[col[j]].begin(),idx[col[j]].end(),liml[j]);
}
}
}
if(k<=cnt1) for(int i=1;i<=tot;i++) limr[i]=p1[i];
else if(cnt1<k && k<=cnt2) {fout<<T.query(rt[nowu],rt[nowv],1,n);exit(0);}
else {k-=cnt2;for(int i=1;i<=tot;i++) liml[i]=p2[i]+1;}
}
bool Ed;
signed main()
{
cerr<<(&St-&Ed)/1024.0/1024.0<<"\n";
fin>>n>>k;
powp[0]=fpow[0]=1;
for(int i=1;i<=n;i++) powp[i]=1ll*powp[i-1]*p%MOD,fpow[i]=qpow(powp[i],MOD-2,MOD);
for(int i=1;i<n;i++)
{
int u,v,w;fin>>u>>v>>w;
add_edge(u,v,w);
add_edge(v,u,w);
}
mn=sum=n;
get_ctr(1,0);
int rrtt=ctr;
get_ctr(rrtt,0);
init_mp(rrtt);
for(int i=1;i<=n;i++) vis[i]=0;
for(int i=1;i<=tot;i++) id[i]=i;
mn=sum=n;
get_ctr(rrtt,0);
calc(rrtt);
for(int i=1;i<=tot;i++) tmp[i]=col[i];
for(int i=1;i<=tot;i++) col[i]=tmp[id[i]];
while(1) solve();
return 0;
}
浙公网安备 33010602011771号