如何在Codeforces Xor-Paths问题中应用Meet-in-the-Middle算法
解决F. Xor-Paths问题的Meet-in-the-Middle方法
问题描述
给定n×m的矩形网格,每个单元格(i,j)对应数字a[i][j]。需要计算从左上角(0,0)到右下角(n-1,m-1)的路径数量,满足两个约束:
- 仅能向右或向下移动,即从(i,j)可移动到(i,j+1)或(i+1,j),不能超出网格范围。
- 路径上所有数字的异或结果等于k(异或操作在Java/C++中用
^表示)。
当n=20、m=20时,暴力遍历所有路径会超时。我当前的暴力递归解法如下:
#include <iostream> #include <vector> #include <unordered_set> using namespace std; bool bfs(int i, int j, long long k, long long curr_xor, const vector<vector<long long>>& a, int& count, vector<vector<unordered_set<long long>>>& bad_cells) { if (j == a[0].size() || i == a.size()) return false; if (bad_cells[i][j].find(curr_xor) != bad_cells[i][j].end()) return false; long long new_xor = curr_xor ^ a[i][j]; if (i == a.size() - 1 && j == a[0].size() - 1) { if (new_xor == k) { ++count; return true; } return false; } bool right_valid = bfs(i, j + 1, k, new_xor, a, count, bad_cells); bool down_valid = bfs(i + 1, j, k, new_xor, a, count, bad_cells); if (!right_valid) bad_cells[i][j + 1].insert(new_xor); if (!down_valid) bad_cells[i + 1][j].insert(new_xor); return right_valid || down_valid; } long long solve(int n, int m, long long k, const vector<vector<long long>>& a) { int count = 0; vector<vector<unordered_set<long long>>> bad_cells(n + 1, vector<unordered_set<long long>>(m + 1)); bfs(0, 0, k, 0, a, count, bad_cells); return count; } int main() { ios_base::sync_with_stdio(false); cin.tie(nullptr); cout.tie(nullptr); int n, m; long long k; cin >> n >> m >> k; vector<vector<long long>> a(n, vector<long long>(m)); for (int i = 0; i < n; ++i) { for (int j = 0; j < m; ++j) { cin >> a[i][j]; } } cout << solve(n, m, k, a); }
该解法通过递归遍历所有路径,计算异或值并维护bad_cells记录无效异或值,但仍会超时。问题标签包含meet-in-the-middle,我考虑过拆分网格,但不确定如何合理拆分并合并结果,希望了解该算法的高效应用方式。
Meet-in-the-Middle解法思路
1. 路径拆分逻辑
从起点到终点的路径总共需要走n+m-2步。我们选择中间步数t = (n+m-2)/2,所有满足i+j = t的单元格构成路径的“中点层”——任何完整路径都会恰好经过该层中的一个单元格。这样可以把原问题拆成两个独立的子问题:
- 计算从起点到中点层所有单元格的路径,记录每个单元格对应的异或值出现次数。
- 计算从终点到中点层所有单元格的反向路径,结合前半部分的统计结果,累加符合条件的路径数量。
2. 前半路径统计
用DFS或BFS遍历从起点(0,0)出发的所有路径,直到走到中点层i+j = t。对于每个中点单元格(i,j),用哈希表记录每个异或值的出现次数(比如mid[i][j][xor_val]表示从起点到(i,j)异或值为xor_val的路径数)。
3. 后半路径匹配
从终点(n-1,m-1)出发,反向遍历(向左或向上移动)到中点层。对于每个到达的中点单元格(i,j),计算当前的异或值y(该值对应正向路径中从(i,j)到终点的异或结果,不含(i,j)本身的数值)。此时,我们需要前半路径中异或值x满足x ^ y = k(完整路径的异或结果为k),即x = k ^ y。查询前半部分哈希表中x的出现次数,累加到总结果中。
4. 实现代码示例
#include <iostream> #include <vector> #include <unordered_map> using namespace std; typedef long long ll; int n, m; ll k; vector<vector<ll>> a; vector<vector<unordered_map<ll, int>>> mid; // 前半部分DFS:遍历到中点层,记录异或值次数 void dfs1(int i, int j, ll xor_val, int t) { xor_val ^= a[i][j]; if (i + j == t) { mid[i][j][xor_val]++; return; } if (i + 1 < n) dfs1(i+1, j, xor_val, t); if (j + 1 < m) dfs1(i, j+1, xor_val, t); } // 后半部分DFS:反向遍历到中点层,统计符合条件的路径数 ll dfs2(int i, int j, ll xor_val, int t) { if (i + j == t) { return mid[i][j].count(k ^ xor_val) ? mid[i][j][k ^ xor_val] : 0; } ll res = 0; if (i - 1 >= 0) res += dfs2(i-1, j, xor_val ^ a[i][j], t); if (j - 1 >= 0) res += dfs2(i, j-1, xor_val ^ a[i][j], t); return res; } int main() { ios::sync_with_stdio(false); cin.tie(nullptr); cin >> n >> m >> k; a.resize(n, vector<ll>(m)); for (int i = 0; i < n; ++i) { for (int j = 0; j < m; ++j) { cin >> a[i][j]; } } int t = (n + m - 2) / 2; mid.resize(n, vector<unordered_map<ll, int>>(m)); dfs1(0, 0, 0, t); ll ans = dfs2(n-1, m-1, 0, t); cout << ans << endl; return 0; }
代码说明
dfs1负责遍历前半路径,将每个中点单元格的异或值出现次数存入哈希表。dfs2从终点反向遍历,计算到中点时的异或值,查询前半部分哈希表中满足条件的路径数并累加。- 时间复杂度为
O(2^((n+m)/2)),对于n=m=20的情况,仅需处理约5e5条路径,完全符合时间要求。
内容的提问来源于stack exchange,提问作者Szyszka947
相关产品推荐
相关产品推荐

