CF2127E 思路分享(线段树合并,线段树上二分,构造)
https://codeforces.com/problemset/problem/2127/E
题意概述
给定一棵根为 \(1\) 的有根树,每个节点有权值 \(w_i\) 和 颜色 \(c_i\),有一些节点没有颜色,即 \(c_i=0\)。
称节点 \(u\) 为 \(cutie\) 当且仅当存在两个节点 \(a,b\) 满足:
-
\(lca(a,b)=u\)
-
\(c_a = c_b\)
-
\(c_a \ne c_u\)
树的花费定义为所有 \(cutie\) 节点的权值之和。
给所有没有颜色的节点分配一个 \(1\) 到 \(k\) 的颜色,使得花费最小。
\(3\le n \le 2\times 10^5\),\(2 \le k \le n\)。
思路
对于每个节点 \(u\),记 \(cnt_i\) 为不同子树中颜色 \(i\) 出现的次数(一棵子树中出现多次算一次)。
如果存在两种及以上颜色 \(cnt_i \ge 2\),那么 \(u\) 一定是 \(cuite\)。
如果只有一种颜色 \(cnt_i \ge 2\),如果 \(c_u = i\),\(u\) 不是 \(cuite\);或者 \(c_u = 0\),可以将 \(u\) 染成 \(i\),使得 \(u\) 不是 \(cuite\)。除了这种情况以外,如果 \(u\) 没被染色,将 \(u\) 随便染成一种子树中出现过的颜色即可,如果子树没有出现过任何颜色,可以先不染这个点,等最后把没染的点染成父节点的颜色即可。需要特判颜色全 \(0\) 的情况。
分析可知,这样的贪心是正确的。
可以使用线段树合并维护子树内的颜色信息,维护一种颜色出现的最大次数即可。使用线段树上二分来寻找 \(cnt_i \ge 2\) 的颜色,通过找第一个和最后一个来判断有没有至少两种这样的颜色。
处理完当前节点,需要将每种颜色出现次数限制为 \(1\),使用 \(lazy\) 标记实现即可。
时间复杂度 \(\mathcal{O}(n\log k)\)。
代码
//author:kzssCCC
#include <bits/stdc++.h>
using namespace std;
using ll = long long;
class node{
public:
int mx=0,left=-1,right=-1;
bool lazy = false;
};
void solve(){
int n,k;
cin >> n >> k;
vector<ll> W(n+1);
for (int i=1;i<=n;i++){
cin >> W[i];
}
vector<int> C(n+1);
for (int i=1;i<=n;i++){
cin >> C[i];
}
vector<vector<int>> adj(n+1);
for (int i=0;i<n-1;i++){
int u,v;
cin >> u >> v;
adj[u].push_back(v);
adj[v].push_back(u);
}
if (accumulate(C.begin()+1,C.end(),0ll)==0){
cout << 0 << '\n';
for (int i=1;i<=n;i++){
cout << 1 << ' ';
}
cout << '\n';
return;
}
vector<node> seg{{}};
vector<int> root(n+1);
for (int i=1;i<=n;i++){
seg.push_back({});
root[i] = seg.size()-1;
}
ll res = 0;
function<void(int)> push_down = [&](int rt){
if (!seg[rt].lazy) return;
if (seg[rt].left!=-1){
seg[seg[rt].left].mx = min(seg[seg[rt].left].mx,1);
seg[seg[rt].left].lazy = true;
}
if (seg[rt].right!=-1){
seg[seg[rt].right].mx = min(seg[seg[rt].right].mx,1);
seg[seg[rt].right].lazy = true;
}
seg[rt].lazy = false;
};
function<int(int,int,int,int)> unite = [&](int rt1,int rt2,int l,int r){
if (rt1==-1 || rt2==-1){
return rt1==-1?rt2:rt1;
}
if (l==r){
seg[rt1].mx += seg[rt2].mx;
return rt1;
}
push_down(rt1);
push_down(rt2);
int mid = l+r >> 1;
seg[rt1].left = unite(seg[rt1].left,seg[rt2].left,l,mid);
seg[rt1].right = unite(seg[rt1].right,seg[rt2].right,mid+1,r);
seg[rt1].mx = max(seg[rt1].left!=-1?seg[seg[rt1].left].mx:0,seg[rt1].right!=-1?seg[seg[rt1].right].mx:0);
return rt1;
};
function<void(int,int,int,int)> update = [&](int rt,int l,int r,int pos){
if (l==r){
seg[rt].mx++;
return;
}
push_down(rt);
int mid = l+r >> 1;
if (pos<=mid){
if (seg[rt].left==-1){
seg.push_back({});
seg[rt].left = seg.size()-1;
}
update(seg[rt].left,l,mid,pos);
}
else{
if (seg[rt].right==-1){
seg.push_back({});
seg[rt].right = seg.size()-1;
}
update(seg[rt].right,mid+1,r,pos);
}
seg[rt].mx = max(seg[rt].left!=-1?seg[seg[rt].left].mx:0,seg[rt].right!=-1?seg[seg[rt].right].mx:0);
};
function<int(int,int,int)> get = [&](int rt,int l,int r){
if (rt==-1 || seg[rt].mx==0) return -1;
if (l==r){
return l;
}
push_down(rt);
int mid = l+r >> 1;
int res = get(seg[rt].left,l,mid);
if (res!=-1) return res;
return get(seg[rt].right,mid+1,r);
};
function<int(int,int,int)> fv = [&](int rt,int l,int r){
if (rt==-1 || seg[rt].mx<2) return -1;
if (l==r){
return l;
}
push_down(rt);
int mid = l+r >> 1;
int res = fv(seg[rt].left,l,mid);
if (res!=-1) return res;
return fv(seg[rt].right,mid+1,r);
};
function<int(int,int,int)> lv = [&](int rt,int l,int r){
if (rt==-1 || seg[rt].mx<2) return -1;
if (l==r){
return l;
}
push_down(rt);
int mid = l+r >> 1;
int res = lv(seg[rt].right,mid+1,r);
if (res!=-1) return res;
return lv(seg[rt].left,l,mid);
};
function<void(int,int)> dfs = [&](int u,int par){
for (auto& v:adj[u]){
if (v==par) continue;
dfs(v,u);
unite(root[u],root[v],1,k);
}
int p1 = fv(root[u],1,k);
int p2 = lv(root[u],1,k);
if (p1!=-1 && p1!=p2 || p1!=-1 && C[u]!=0 && C[u]!=p1){
res += W[u];
}
if (C[u]==0){
if (p1!=-1 && p1==p2){
C[u] = p1;
}
else if (seg[root[u]].mx!=0){
C[u] = get(root[u],1,k);
}
}
else update(root[u],1,k,C[u]);
seg[root[u]].lazy = true;
seg[root[u]].mx = min(seg[root[u]].mx,1);
};
dfs(1,-1);
function<void(int,int)> dfs2 = [&](int u,int par){
for (auto& v:adj[u]){
if (v==par) continue;
if (C[v]==0){
C[v] = C[u];
}
dfs2(v,u);
}
};
dfs2(1,-1);
cout << res << '\n';
for (int i=1;i<=n;i++){
cout << C[i] << ' ';
}
cout << '\n';
}
int main(){
ios::sync_with_stdio(false);
cin.tie(0);
int t = 1;
cin >> t;
while (t--) solve();
return 0;
}

浙公网安备 33010602011771号