P4242 树上的毒瘤
这题太累人了,就算是我这种马蜂紧凑的也爆写了将近 5 KB 的代码。
观察一下题目,先是维护树上颜色段,想到树剖+线段树;再是若干关键点之间两两求贡献,想到虚树+点分治。写点分治时,我们固定一个 LCA,不难发现两条链合并时 LCA 处颜色相等,那贡献就是颜色段和减 $1$。那我们建虚树时就直接把两个关键点在原图上的链的颜色段数减一作为边权,然后统计答案时就加上一个 $1$ 即可。
史代码:
C++
#include<bits/stdc++.h>
using namespace std;
#define ls p<<1
#define rs p<<1|1
const int N=1e5+5;
int n,q,op,a,b,c,idx,sum,rt,tot,sumd;
int col[N],dot[N],nd[N],tag[N];
int fa[N],dep[N],dfn[N],top[N],son[N],inv[N];
int sz[N],vis[N],ans[N],dp[N],siz[N],pdis[N],sumv[N];
vector<int> g[N],sonv[N];
vector<pair<int,int>> vt[N];
struct node{
int sum,lc,rc,tag;
}tr[N<<2];
bool cmp(int x,int y){
return dfn[x]<dfn[y];
}
void dfs1(int u,int f){
dep[u]=dep[f]+1,fa[u]=f,siz[u]=1;
for(int v:g[u]){
if(v==f) continue;
dfs1(v,u);
siz[u]+=siz[v];
if(siz[v]>siz[son[u]]) son[u]=v;
}
}
void dfs2(int u,int tf){
dfn[u]=++idx,inv[idx]=u,top[u]=tf;
if(son[u]) dfs2(son[u],tf);
for(int v:g[u]){
if(v==fa[u]||v==son[u]) continue;
dfs2(v,v);
}
}
int lca(int x,int y){
while(top[x]!=top[y]){
if(dep[top[x]]<dep[top[y]]) swap(x,y);
x=fa[top[x]];
}
return dep[x]<dep[y]?x:y;
}
node merge(node x,node y){
if(!x.sum) return y;
return {x.sum+y.sum-(x.rc==y.lc),x.lc,y.rc,0};
}
void pushup(int p){
tr[p]=merge(tr[ls],tr[rs]);
}
void pushdown(int p){
if(tr[p].tag){
tr[ls]={1,tr[p].tag,tr[p].tag,tr[p].tag};
tr[rs]={1,tr[p].tag,tr[p].tag,tr[p].tag};
tr[p].tag=0;
}
}
void build(int p,int l,int r){
if(l==r){
tr[p]={1,col[inv[l]],col[inv[l]],0};
return;
}
int mid=(l+r)>>1;
build(ls,l,mid);
build(rs,mid+1,r);
pushup(p);
}
void change(int p,int l,int r,int x,int y,int k){
if(l>=x&&r<=y){
tr[p]={1,k,k,k};
return;
}
int mid=(l+r)>>1;
pushdown(p);
if(x<=mid) change(ls,l,mid,x,y,k);
if(y>mid) change(rs,mid+1,r,x,y,k);
pushup(p);
}
node get(int p,int l,int r,int x,int y){
if(l>=x&&r<=y) return tr[p];
int mid=(l+r)>>1;
pushdown(p);
if(y<=mid) return get(ls,l,mid,x,y);
if(x>mid) return get(rs,mid+1,r,x,y);
return merge(get(ls,l,mid,x,y),get(rs,mid+1,r,x,y));
}
void pchange(int x,int y,int k){
while(top[x]!=top[y]){
if(dep[top[x]]<dep[top[y]]) swap(x,y);
change(1,1,n,dfn[top[x]],dfn[x],k);
x=fa[top[x]];
}
if(dep[x]>dep[y]) swap(x,y);
change(1,1,n,dfn[x],dfn[y],k);
}
int pquery(int x,int y){
node res={0,0,0,0};
while(top[x]!=top[y]){
res=merge(get(1,1,n,dfn[top[x]],dfn[x]),res);
x=fa[top[x]];
}
res=merge(get(1,1,n,dfn[y],dfn[x]),res);
return res.sum;
}
void buildvt(){
sort(dot+1,dot+1+a,cmp);
for(int i=1;i<=a;i++){
nd[++sum]=dot[i];
if(i<a) nd[++sum]=lca(dot[i],dot[i+1]);
}
sort(nd+1,nd+1+sum,cmp);
sum=unique(nd+1,nd+1+sum)-nd-1;
for(int i=1;i<sum;i++){
int f=lca(nd[i],nd[i+1]),wis=pquery(nd[i+1],f);
vt[f].push_back({nd[i+1],wis-1});
vt[nd[i+1]].push_back({f,wis-1});
}
}
void clearvt(){
for(int i=1;i<sum;i++){
int f=lca(nd[i],nd[i+1]);
vt[f].clear();
vt[nd[i+1]].clear();
vis[f]=vis[nd[i+1]]=ans[f]=ans[nd[i+1]]=0;
}
sum=tot=0;
}
void getrt(int u,int f){
siz[u]=1,dp[u]=0;
for(auto[v,w]:vt[u]){
if(v==f||vis[v]) continue;
getrt(v,u);
siz[u]+=siz[v];
dp[u]=max(dp[u],siz[v]);
}
dp[u]=max(dp[u],tot-siz[u]);
if(dp[u]<dp[rt]) rt=u;
}
void getdis(int u,int f,int dis,int rot){
if(tag[u]){
pdis[u]=dis,sumd+=dis;
sonv[rot].push_back(u);
}
for(auto[v,w]:vt[u]){
if(v==f||vis[v]) continue;
getdis(v,u,dis+w,rot);
}
}
void calc(int u){
int totd=0;
sumd=0;
for(auto[v,w]:vt[u]){
if(vis[v]) continue;
sumv[v]=sumd;
getdis(v,u,w,v);
sumv[v]=sumd-sumv[v],totd+=sonv[v].size();
}
for(auto[v,w]:vt[u]){
if(vis[v]) continue;
for(auto x:sonv[v]){
ans[x]+=(pdis[x]+1)*(totd-sonv[v].size())+(sumd-sumv[v]);
if(tag[u]) ans[x]+=(pdis[x]+1),ans[u]+=(pdis[x]+1);
}
sonv[v].clear();
}
}
void solve(int u){
vis[u]=1;
calc(u);
for(auto[v,w]:vt[u]){
if(vis[v]) continue;
tot=siz[v],dp[v]=1e9,rt=0;
getrt(v,0);
solve(rt);
}
}
int main(){
cin>>n>>q;
for(int i=1;i<=n;i++) cin>>col[i];
for(int i=1,u,v;i<n;i++){
cin>>u>>v;
g[u].push_back(v);
g[v].push_back(u);
}
dfs1(1,0);
dfs2(1,1);
build(1,1,n);
while(q--){
cin>>op;
if(op==1){
cin>>a>>b>>c;
pchange(a,b,c);
}
else{
cin>>a;
for(int i=1;i<=a;i++) cin>>dot[i],tag[dot[i]]=1;
buildvt();
getrt(nd[1],0);
solve(nd[1]);
for(int i=1;i<=a;i++){
cout<<ans[dot[i]]<<" ";
tag[dot[i]]=0;
}
cout<<"\n";
clearvt();
}
}
return 0;
}


Comments NOTHING