题解:CF2252F Spectral Components

唐唐题。

固定颜色 \(c\),设其出现次数为 \(cnt_c\)。对于一条边 \(e=(u,v)\),设断开这条边后两个连通块中 \(c\) 的出现次数分别为 \(x_e\)\(cnt_c-x_e\)。对于一个连通块 \(S\),若 \(e\in S\)\(e\) 对距离和没有贡献,否则贡献为 \(w_e=\min(x_e,cnt_c-x_e)\)

显然答案的下界是 \(\sum w_e\) 减去 \(w\) 的前 \(k_c-1\) 大之和。我们证明这个下界可以取到。

证明

对颜色 \(c\) 取带权重心 \(r\),以 \(r\) 为根。那么对于 \(e=(u,fa_u)\),我们有 \(w_e=sz_u\),那么显然在一条从根向下的路径上 \(w\) 单调递减。把所有边按权值从大到小排序,权值相同的按深度从小到大排序,那么此时直接选择前 \(k_c-1\) 条边得到的必然是一个包含 \(r\) 的连通块。\(\Box\)

于是我们只需要对每种颜色 \(c\) 维护出 \(\sum w_e\)\(w\) 的前 \(k_c-1\) 大之和。

考虑对所有颜色为 \(c\) 的点建出虚树,设 \(f_u\)\(u\) 子树内颜色为 \(c\) 的点数。那么对于一条虚树边 \((u,v)\),其中 \(u\)\(v\) 的祖先,这上面的 \(dep_v-dep_u\) 条边的 \(w\) 都是 \(\min(f_v,cnt_c-f_v)\)。记录所有 \((w,dep_v-dep_u)\) 二元组即可求出前 \(k_c-1\) 大之和。

时间复杂度为 \(\mathcal{O}(n\log{n})\)

代码
#include <bits/stdc++.h>

using namespace std;

using ll = long long;
using i128 = __int128;
using ui = unsigned int;
using ull = unsigned long long;
using u128 = unsigned __int128;
using ld = long double;
using pii = pair<int, int>;
const int MAXN = 2e5 + 5, LOGN = 18;

template<typename T> T lowbit(T x) { return x & -x; }
template<typename T> void chkMin(T &x, T y) { x = y < x ? y : x; }
template<typename T> void chkMax(T &x, T y) { x = x < y ? y : x; }
constexpr int lg2(ll x) { return 63 ^ __builtin_clzll(x); }
constexpr ll bitCeil(ll x) { return x == 1 ? 1ll : 1ll << lg2(x - 1) + 1; }

int tc, n, c[MAXN], k[MAXN];
ll ans[MAXN];
vector<int> T[MAXN], buc[MAXN], VT[MAXN], C;
int stmp, dfn[MAXN], dep[MAXN];
int top, stk[MAXN];
int cnt[MAXN];
bool mark[MAXN];

struct ST {
	int f[LOGN][MAXN];

	int get(int x, int y) {
		return dfn[x] < dfn[y] ? x : y;
	}

	void init() {
		for (int i = 1; (1 << i) <= n; ++i)
			for (int j = 1; j <= n - (1 << i) + 1; ++j)
				f[i][j] = get(f[i - 1][j], f[i - 1][j + (1 << i - 1)]);
	}

	int query(int l, int r) {
		int k = lg2(r - l + 1);
		return get(f[k][l], f[k][r - (1 << k) + 1]);
	}
} st;

int lca(int x, int y) {
	if (x == y) return x;
	if (dfn[x] > dfn[y]) swap(x, y);
	return st.query(dfn[x] + 1, dfn[y]);
}

void dfs(int u, int faU) {
	dfn[u] = ++stmp;
	st.f[0][stmp] = faU;
	for (int v : T[u]) {
		if (v == faU) continue;
		dep[v] = dep[u] + 1;
		dfs(v, u);
	}
}

int main() {
	ios::sync_with_stdio(false);
	cin.tie(nullptr);
	
	cin >> tc;
	while (tc--) {
		cin >> n;
		for (int i = 1; i <= n; ++i) buc[i].clear();
		for (int i = 1; i <= n; ++i) {
			cin >> c[i];
			buc[c[i]].emplace_back(i);
		}
		for (int i = 1; i <= n; ++i) cin >> k[i];
		for (int i = 1; i <= n; ++i) T[i].clear();
		for (int i = 1; i < n; ++i) {
			int u, v;
			cin >> u >> v;
			T[u].emplace_back(v);
			T[v].emplace_back(u);
		}

		stmp = 0;
		dfs(1, 0);
		st.init();

		auto ins = [&](int x, int y) {
			if (!mark[x]) {
				mark[x] = true;
				C.emplace_back(x);
			}
			if (!mark[y]) {
				mark[y] = true;
				C.emplace_back(y);
			}
			VT[x].emplace_back(y);
		};

		for (int col = 1; col <= n; ++col) {
			if (buc[col].empty()) {
				ans[col] = -1;
				continue;
			}

			sort(buc[col].begin(), buc[col].end(), [&](int x, int y) {
				return dfn[x] < dfn[y];
			});
			stk[top = 1] = buc[col][0];
			for (int i = 1; i < buc[col].size(); ++i) {
				int x = buc[col][i], d = lca(stk[top], x);
				while (top > 1 && dep[d] <= dep[stk[top - 1]]) {
					ins(stk[top - 1], stk[top]);
					--top;
				}
				if (d != stk[top]) {
					ins(d, stk[top]);
					stk[top] = d;
				}
				stk[++top] = x;
			}
			for (int i = top - 1; i; --i) ins(stk[i], stk[i + 1]);

			vector<pii> vec;
			ll sum = 0;
			auto dfs = [&](auto &&self, int u) -> void {
				cnt[u] = c[u] == col;
				for (int v : VT[u]) {
					self(self, v);
					cnt[u] += cnt[v];
					int w1 = min<int>(cnt[v], buc[col].size() - cnt[v]), w2 = dep[v] - dep[u];
					vec.emplace_back(w1, w2);
					sum += (ll)w1 * w2;
				}
			};
			dfs(dfs, stk[1]);

			sort(vec.begin(), vec.end(), greater<>());
			int r = k[col] - 1;
			for (auto [w1, w2] : vec) {
				if (!r) break;
				int v = min(r, w2);
				sum -= (ll)v * w1;
				r -= v;
			}
			ans[col] = sum;

			for (int x : C) {
				mark[x] = false;
				VT[x].clear();
			}
			C.clear();
		}

		for (int i = 1; i <= n; ++i) cout << ans[i] << " \n"[i == n];
	}
	return 0;
}
posted @ 2026-08-07 11:12  P2441M  阅读(12)  评论(0)    收藏  举报