题解:P16459 [UOI 2026] Tree Subsets

可以直接 ntt 全局平衡二叉树做到两个老哥虽然常数比较大不一定能过,但是看到这个题就反应出这个也是没救了。

在一年前我吹牛时胡出这个东西时发表暴论要把所有可以这样优化的树上背包出成两个老哥。

题意:一棵树,点有 \(0/1\) 的点权,要求你选 \(k\) 个点,要求一个点如果被选入,其子树也得被选入,询问是否可以在选 \(k\) 个点的时候凑出所有点异或和为 \(0/1\),对 \(k\in[1,n]\) 都要求。\(n\le 4\times 10^5\)

做法:

假设你并不知道这个 trick,该怎么想出来这个东西呢?

首先树上选若干个点这显然是树上背包,我们可以朴素的写出来一个平方的 dp,比较简单就不赘述了。

然后观察一下部分分,发现比较特别,sub 7 和 8 给的叫,这个图是把一挖掉之后若干条链。发现链非常好做的一点是他在选若干个点时非常简单,我们只需要把每一条链选 \(x\) 个的结果记为 \(F(x)\),我们把 \(F(x)\) 合并起来就可以了。

现在考虑怎么样快速地合并 \(F(x)\),有一个无脑的做法叫,我分别记两个 \(F_0(x),F_1(x)\) 代表第 \(i\) 项是否可以是 \(0/1\),直接卷积加法即可,可以做到一个老哥。但是这太魔怔了,明明每一项只有 \(0/1\) 我们却需要一个卷积显然是啥子做法,我们考虑不分离 \(F\) 直接做,每一个数的值就是题目中的 \(0,1,2\),那么假设我需要计算 \(F(x)G(x)[x^i]\) 的值,我们可以计算出来其要求在 \(F(x)\) 的有效范围和 \(G(x)\) 中的,注意我们这个 \(F(x)\) 的系数 \(0\) 并不是没有,所以如果在幂次外面的是不可以被选入有效范围内的。首先如果这两个范围中有 \(2\) 那么显然我这个数也是 \(2\),否则我考虑得是 \(0,1\) 都有才行,那么这个东西应该是把 \(G(x)\) 的有效区间翻转,和 \(F(x)\) 的有效区间对位异或,发现如果两个区间完全相等或者刚好是完全不相等,那么答案为对应异或值,否则为 \(2\),这个东西就可以直接采用哈希判断。这样我们的复杂度就是线性的,两个函数长度之和。

然后考虑怎么合并多个函数,因为我的代价是两个长度之和,就不能用启发式来做了,所幸我们知道正经的多项式合并也有这个问题,解决方法是直接分治合并即可,那么 sub 7 和 8 就做到了一个 \(\log\)

那么考虑怎么解决满分,考虑还是沿用上面的思路,我们对点 \(u\) 的子树这么做一下,但是直接这么做显然不太对,这样合并是 \(O(sz_u)\) 的,我们考虑上一个重链剖分,我们把轻儿子先分治合并,然后再把重链分治合并,这样就可以做到两个老哥,用全局平衡二叉树可以做到一个老哥。大概思路就这么简单了。

一些实现的细节:我们合并重链时要求这样一个柿子:

\[((\cdots((F_1(x)+x^{a_1})F_2(x) +x^{a_2})\cdots)F_n(x)+x^{a_n}) \]

这个东西我们考虑维护一对多项式,记为 \((F(x),G(x))\),那么合并时复合两个多项式就变成:\((F_l(x)F_r(x),G_l(x)F_r(x)+G_r(x))\),答案就是这一对之和。

但是这里还有点小问题,我们发现 \(a_i\) 之和没有保证,且你不好确认 \(G\) 的定义域,因为我们只有定义域内的 \(0,1,2\) 才是有意义的,但是经过一些简单的分讨会发现 \(G\) 的定义域其实是连续的,所以我们记录定义域,并且只保留定义域内的部分去做发现对于 \(a_i\) 他的大小就是正确的。

还不太懂可以看一下代码,感觉写的还是很清晰的。

代码:

#include <bits/stdc++.h>
using namespace std;
const int maxn = 4e5 + 5; const unsigned long long bs = 131;
vector<int> e[maxn]; int n, a[maxn], p[maxn];
unsigned long long pw[maxn];
struct Hash {
	unsigned long long res; int len;
	friend Hash operator+(Hash x, Hash y) {
		Hash ans;
		ans.res = (x.res * pw[y.len] + y.res);
		ans.len = x.len + y.len;
		return ans;
	}
	friend Hash operator-(Hash x, Hash y) {
		Hash ans; ans.len = x.len - y.len;
		ans.res = (x.res - y.res * pw[x.len - y.len]);
		return ans;
	}
	Hash() {
		len = 0, res = 0;
	}
	Hash(unsigned long long V) {
		len = 1, res = V;
	}
	friend bool operator!=(Hash x, Hash y) {
		return (x.res != y.res || x.len != y.len);
	}
};
mt19937 rnd(time(0));
unsigned long long val[2] = {12, 31};
unsigned long long st[2][maxn];
struct Poly {
	vector<int> a, s; int l, r;
	vector<Hash> pre, suf;
	Poly() {
		l = r = 0;
		a.clear(), s.clear(), pre.clear(), suf.clear();
	}
	int& operator[](int x) {
		return a[x];
	}
	int size() {
		return a.size();
	}
	void resize(int N) {
		a.resize(N), pre.resize(N), suf.resize(N), s.resize(N);
	}
	void get_Hash() {
		int n = size();
		for (int i = 0; i < n; i++)
			s[i] = (i ? s[i - 1] : 0) + (a[i] == 2);
		for (int i = 0; i < n; i++)
			pre[i] = (i ? pre[i - 1] : Hash()) + Hash(val[a[i]]);
		for (int i = n - 1; i >= 0; i--)	
			suf[i] = (i == n - 1 ? Hash() : suf[i + 1]) + Hash(val[a[i]]);
	}
	Hash get_z(int l, int r) {
		return pre[r] - (l == 0 ? Hash() : pre[l - 1]);
	}
	Hash get_r(int l, int r) {
		return suf[l] - (r == size() - 1 ? Hash() : suf[r + 1]);
	}
	int get_s(int l, int r) {
		return s[r] - (l == 0 ? 0 : s[l - 1]);
	}
	friend Poly operator*(Poly f, Poly g) {
		if(!f.size())
			return g;
		if(!g.size())
			return f;
//		cout << f.size() << " " << g.size() << " " << f.l << " " << f.r << " " << g.l << " " << g.r << endl;
//		for (int i = 0; i < f.size(); i++)
//			cout << f[i] << " ";
//		cout << endl;
//		for (int i = 0; i < g.size(); i++)
//			cout << g[i] << " ";
//		cout << endl;
		f.get_Hash(), g.get_Hash();
//		for (int i = 0; i < f.size(); i++)
//			cout << f.pre[0][i].res[0] << " ";
//		cout << endl;
		Poly ans; ans.l = f.l + g.l, ans.r = f.r + g.r;
		ans.resize(ans.r - ans.l + 1);
		for (int i = ans.l; i <= ans.r; i++) {
			int lxf = 0, rxf = i, lxg = 0, rxg = i;
			rxf = min(rxf, f.r), rxg = min(rxg, g.r);
			lxf = max(lxf, f.l), lxg = max(lxg, g.l);
			lxf = max(lxf, i - rxg), lxg = max(lxg, i - rxf);
			rxg = min(rxg, i - lxf), rxf = min(rxf, i - lxg);
			unsigned long long x = f.get_z(lxf - f.l, rxf - f.l).res, y = g.get_r(lxg - g.l, rxg - g.l).res;
//			cout << i << " " << lxf << " " << rxf << " " << lxg << " " << rxg << " " << x << " " << y << endl;
			if(f.get_s(lxf - f.l, rxf - f.l) || g.get_s(lxg - g.l, rxg - g.l) || (x != y && st[0][rxf - lxf + 1] - x != y))
				ans[i - ans.l] = 2;
			else 
				ans[i - ans.l] = f[lxf - f.l] ^ g[rxg - g.l];
		}
//		cout << ans.l << " " << ans.r << endl;
//		for (int i = 0; i < ans.size(); i++)
//			cout << ans[i] << " ";
//		cout << endl;
//		cout << "BOMB" << endl;
		return ans;
	}
	friend Poly operator+(Poly f, Poly g) {
		if(!f.size())
			return g;
		if(!g.size())
			return f;
		Poly ans; ans.l = min(f.l, g.l), ans.r = max(f.r, g.r);
		ans.resize(ans.r - ans.l + 1);
//		cout << g.l << " " << g.r << " " << g.size() << endl;
		for (int i = ans.l; i <= ans.r; i++) {
			int x = (f.l <= i && f.r >= i ? f[i - f.l] : -1), y = (g.l <= i && i <= g.r ? g[i - g.l] : -1);
			if(x == -1)
				ans[i - ans.l] = y;
			else if(y == -1)
				ans[i - ans.l] = x;
			else {
				if(x == 2 || y == 2)
					ans[i - ans.l] = 2;
				else
					ans[i - ans.l] = (x == y ? x : 2);
			}
		}
//		cout << ans.l << " " << ans.r << endl;
		return ans;
	}
	void get_to(int lx, int rx) {
		while(r < rx)
			a.push_back(0), r++;
//		cout << a.size() << endl;
		reverse(a.begin(), a.end());
		while(l > lx)
			a.push_back(0), l--;
		reverse(a.begin(), a.end());
		resize(a.size());
//		cout << a.size() << "asdf" << endl;
	}
} ;
Poly f[maxn], g[maxn], rest[maxn];
int sz[maxn], res[maxn], tot;
Poly solve_son(int l, int r) {
	if(l > r) {
		Poly f;
		return f;
	}
	if(l == r)
		return f[l];
	int mid = l + r >> 1;
	return solve_son(l, mid) * solve_son(mid + 1, r);
}
pair<Poly, Poly> solve_chain(int l, int r, vector<int> &pos) {
	if(l == r) {
		Poly f1 = g[l], f2; f2.resize(1);
//		cout << pos[l] << " " << l << " " << sz[pos[l]] << " " << endl;
		f2[0] = res[pos[l]]; f2.l = sz[pos[l]], f2.r = sz[pos[l]];
		return make_pair(f1, f2);
	}
	int mid = l + r >> 1;
	pair<Poly, Poly> v1 = solve_chain(l, mid, pos), v2 = solve_chain(mid + 1, r, pos);
	pair<Poly, Poly> res; 
	res.first = v1.first * v2.first;
//	cout << l << " adsg" << r << endl;
	res.second = v1.second * v2.first;
//	cout << l << " " << r << "adsfa;lkg" << " " << res.first.size() << " " << res.second.size() << endl;
	res.second = res.second + v2.second;
//	for (int i = 0; i < res.second.size(); i++)
//		cout << res.second[i] << " ";
//	cout << endl;
	return res;
}
int son[maxn];
void dfs1(int u, int fa) {
	sz[u] = 1, res[u] = a[u]; son[u] = 0;
//	cout << u << endl;
	for (int i = 0; i < e[u].size(); i++) {
		int v = e[u][i];
		if(v == fa)
			continue;
		dfs1(v, u);
		sz[u] += sz[v], res[u] ^= res[v];
		if(sz[son[u]] < sz[v])
			son[u] = v;
	}
}
int s = 0;
void dfs2(int u, int fa) {
	vector<int> pos;
	for (int p = u; p; p = son[p]) {
		pos.push_back(p);
		for (int i = 0; i < e[p].size(); i++) {
			if(e[p][i] == son[p])
				continue;
			dfs2(e[p][i], p);
		}
	}
	int cnt = 0;
//	cout << u << endl;
	reverse(pos.begin(), pos.end());
	for (int j = 0; j < pos.size(); j++) {
		tot = 0; int p = pos[j];
		for (int i = 0; i < e[p].size(); i++) {
			if(e[p][i] == son[p])
				continue;
			f[++tot] = rest[e[p][i]];
		}
		g[cnt++] = solve_son(1, tot);
		if(g[cnt - 1].size())
			g[cnt - 1].get_to(0, (int)g[cnt - 1].size() - 1);
//		g[cnt - 1][0] = 0;
//		if(g[cnt - 1].size())
//			cout << g[cnt - 1][1] << " " << pos[j] << endl;
	}
//	cout << "start" << endl;
	pair<Poly, Poly> rt = solve_chain(0, cnt - 1, pos);
	rest[u] = rt.first + rt.second;
//	cout << "u aslkdjg" << endl;
	rest[u].get_to(0, sz[u]);
	s += sz[u];
//	cout << u << "sadlkj" << endl;
//	for (int i = 0; i < rest[u].size(); i++)
//		cout << rest[u][i] << " ";
//	cout << endl;
}
void solve() {
	cin >> n;
	pw[0] = 1;
	for (int i = 1; i <= n; i++)
		pw[i] = pw[i - 1] * bs, 
		st[0][i] = st[0][i - 1] + pw[i - 1] * (val[0] + val[1]);
	for (int i = 1; i <= n; i++)
		cin >> a[i], e[i].clear();
	for (int i = 2; i <= n; i++)
		cin >> p[i], e[p[i]].push_back(i);
	dfs1(1, 0);
	dfs2(1, 0);
	for (int i = 1; i <= n; i++)
		cout << rest[1][i] << " ";
	cout << endl;
}
signed main() {
//	freopen("test.in", "r", stdin);
//	freopen("std.out", "w", stdout);
	ios::sync_with_stdio(false);
	int T; cin >> T;
	while(T--)	
		solve();
	return 0;
}
/*
1
5
1 0 1 1 0
1 2 1 4
*/
posted @ 2026-07-04 16:26  LUlululu1616  阅读(37)  评论(0)    收藏  举报