CF2194F1 思路分享(状压 dp)

https://codeforces.com/problemset/problem/2194/F1

题意

给定一棵节点数为 \(n\) 的无根树,\(a_i\) 表示节点 \(i\) 的点权.

给定长度为 \(k\) 的序列 \(b\),将树划分成若干连通分量,每个连通分量的异或和为 \(b\) 中某个值,求合法划分的数量,模 \(10^9+7\).

\(2\le n \le 10^5\)\(1\le k \le 4\)\(0\le a_i,b_i \le 2^{30}\).

思路

考虑 \(dp\),需要维护子树内未封闭节点点权的异或和,记为 \(S\),无法直接维护,考虑转化.

\(mask\)\(b\) 中元素的状态掩码,令 \(V_{mask} = \bigoplus_{mask>>i\And 1}{b_i}\).

\(xr_u\)\(u\) 及子树的 \(a_i\) 异或和.

\(mask\) 表示 \(u\) 子树中已封闭部分状态的异或和,因此 \(S=xr_u \oplus V_{mask}\),而 \(mask\)\(2^k\) 量级的,可以通过维护 \(mask\) 间接维护 \(S\).

转移时,讨论是否封闭某个子树.

  • 不封闭子树 \(v\)\(dp[u][i\oplus j] \leftarrow dp[u][i] \times dp[v][j]\).

  • 封闭子树 \(v\)\(S_v\) 必须是 \(b\) 中某个元素,假设为 \(b_t\)\(dp[u][i\oplus j \oplus (1<<t)] \leftarrow dp[u][i] \times dp[v][j]\).

时间复杂度 \(\mathcal{O}(4^k \cdot n)\).

代码

//author:kzssCCC

#include <bits/stdc++.h>
using namespace std;
using ll = long long;

const int MOD = 1e9+7;

void solve(){
	int n,k;
	cin >> n >> k;

	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);
	}

	vector<int> a(n+1);
	for (int i=1;i<=n;i++){
		cin >> a[i];
	}

	map<int,int> mp;
	vector<int> b(k);
	for (int i=0;i<k;i++){
		cin >> b[i];
		mp[b[i]] = 1<<i;
	}
	
	vector<int>	V(1<<k);
	for (int u=0;u<1<<k;u++){
		for (int j=0;j<k;j++){
			if (u>>j&1){
				V[u]^=b[j];
			}
		}
	}

	vector<int> xr(n+1);
	vector<vector<ll>> dp(n+1,vector<ll>(1<<k));
	function<void(int,int)> dfs = [&](int u,int par){
		xr[u] = a[u];
		dp[u][0] = 1;
		for (auto& v:adj[u]){
			if (v==par) continue;
			dfs(v,u);
			xr[u]^=xr[v];			

			vector<ll> ndp(1<<k);
			for (int i=0;i<1<<k;i++){
				if (dp[u][i]==0) continue;
				for (int j=0;j<1<<k;j++){
					if (dp[v][j]==0) continue;

					ndp[i^j]= (ndp[i^j]+dp[u][i]*dp[v][j]%MOD)%MOD;
					if (mp.count(xr[v]^V[j])){
						ndp[i^mp[xr[v]^V[j]]^j] = (ndp[i^mp[xr[v]^V[j]]^j]+dp[u][i]*dp[v][j]%MOD)%MOD;
					}			
				}
			}
			dp[u] = ndp;
		}
	};
 	dfs(1,-1);

 	ll res = 0;
 	for (int u=0;u<1<<k;u++){
 		if (mp.count(xr[1]^V[u])){
 			res = (res+dp[1][u])%MOD;
 		}
 	}
 	cout << res << '\n';
}

int main(){
	ios::sync_with_stdio(false);
	cin.tie(0);
	
	int t = 1;
	cin >> t;
	while (t--) solve();

	return 0;
}
posted @ 2026-07-05 14:26  kzssCCC  阅读(3)  评论(0)    收藏  举报