用 9 位掩码维护每格候选数字,结合数独约束与大小关系反复传播,再按候选最少的格子搜索。
OJ: luogu
题目 ID: P4573
难度:提高+/省选-
标签:搜索回溯位运算约束传播数独
日期: 2026-06-20 22:48
题意
给出一个 9 x 9 的新数独。
除了要满足普通数独规则:
- 每行是
1..9的排列 - 每列是
1..9的排列 - 每个
3 x 3宫也是1..9的排列
还额外给出了若干相邻格子的大小关系。
要求输出一组满足所有规则的完整填法。
思路
先看最基础、最容易理解的版本:
cpp
#include <bits/stdc++.h>
using namespace std;
const int FULL_MASK = (1 << 9) - 1;
char h_rel[9][8];
char v_rel[8][9];
int board[9][9];
int row_mask[9], col_mask[9], box_mask[9];
int get_box_id(int x, int y) {
return (x / 3) * 3 + (y / 3);
}
int bit_count(int x) {
return __builtin_popcount((unsigned int)x);
}
int lowbit_to_digit(int bit) {
for (int d = 1; d <= 9; d++) {
if (bit == (1 << (d - 1))) {
return d;
}
}
return 0;
}
int get_candidate_mask(int x, int y) {
int used = row_mask[x] | col_mask[y] | box_mask[get_box_id(x, y)];
int mask = FULL_MASK ^ used;
// 根据左右不等式删掉不可能的数字。
if (y < 8 && h_rel[x][y] != 0 && board[x][y + 1] != 0) {
int right_val = board[x][y + 1];
int new_mask = 0;
for (int d = 1; d <= 9; d++) {
if ((mask & (1 << (d - 1))) == 0) {
continue;
}
if ((h_rel[x][y] == '<' && d < right_val) || (h_rel[x][y] == '>' && d > right_val)) {
new_mask |= 1 << (d - 1);
}
}
mask = new_mask;
}
if (y > 0 && h_rel[x][y - 1] != 0 && board[x][y - 1] != 0) {
int left_val = board[x][y - 1];
int new_mask = 0;
for (int d = 1; d <= 9; d++) {
if ((mask & (1 << (d - 1))) == 0) {
continue;
}
if ((h_rel[x][y - 1] == '<' && left_val < d) || (h_rel[x][y - 1] == '>' && left_val > d)) {
new_mask |= 1 << (d - 1);
}
}
mask = new_mask;
}
// 根据上下不等式删掉不可能的数字。
if (x < 8 && v_rel[x][y] != 0 && board[x + 1][y] != 0) {
int down_val = board[x + 1][y];
int new_mask = 0;
for (int d = 1; d <= 9; d++) {
if ((mask & (1 << (d - 1))) == 0) {
continue;
}
if ((v_rel[x][y] == '^' && d < down_val) || (v_rel[x][y] == 'v' && d > down_val)) {
new_mask |= 1 << (d - 1);
}
}
mask = new_mask;
}
if (x > 0 && v_rel[x - 1][y] != 0 && board[x - 1][y] != 0) {
int up_val = board[x - 1][y];
int new_mask = 0;
for (int d = 1; d <= 9; d++) {
if ((mask & (1 << (d - 1))) == 0) {
continue;
}
if ((v_rel[x - 1][y] == '^' && up_val < d) || (v_rel[x - 1][y] == 'v' && up_val > d)) {
new_mask |= 1 << (d - 1);
}
}
mask = new_mask;
}
return mask;
}
bool dfs() {
int best_x = -1;
int best_y = -1;
int best_mask = 0;
int best_cnt = 10;
for (int i = 0; i < 9; i++) {
for (int j = 0; j < 9; j++) {
if (board[i][j] != 0) {
continue;
}
int mask = get_candidate_mask(i, j);
int cnt = bit_count(mask);
if (cnt == 0) {
return false;
}
if (cnt < best_cnt) {
best_cnt = cnt;
best_x = i;
best_y = j;
best_mask = mask;
}
}
}
if (best_x == -1) {
return true;
}
while (best_mask != 0) {
int lowbit = best_mask & (-best_mask);
best_mask -= lowbit;
int digit = lowbit_to_digit(lowbit);
int box_id = get_box_id(best_x, best_y);
board[best_x][best_y] = digit;
row_mask[best_x] |= lowbit;
col_mask[best_y] |= lowbit;
box_mask[box_id] |= lowbit;
if (dfs()) {
return true;
}
box_mask[box_id] ^= lowbit;
col_mask[best_y] ^= lowbit;
row_mask[best_x] ^= lowbit;
board[best_x][best_y] = 0;
}
return false;
}
void read_relations() {
vector<string> lines;
string line;
while ((int)lines.size() < 15 && getline(cin, line)) {
if (line.empty()) {
continue;
}
lines.push_back(line);
}
memset(h_rel, 0, sizeof(h_rel));
memset(v_rel, 0, sizeof(v_rel));
for (int idx = 0; idx < 15; idx++) {
vector<char> arr;
for (int i = 0; i < (int)lines[idx].size(); i++) {
char ch = lines[idx][i];
if (ch == '<' || ch == '>' || ch == '^' || ch == 'v') {
arr.push_back(ch);
}
}
if ((int)arr.size() == 6) {
int row = (idx / 5) * 3 + ((idx % 5) / 2);
int cols[6] = {0, 1, 3, 4, 6, 7};
for (int i = 0; i < 6; i++) {
h_rel[row][cols[i]] = arr[i];
}
} else {
int row = (idx / 5) * 3 + ((idx % 5 - 1) / 2);
for (int i = 0; i < 9; i++) {
v_rel[row][i] = arr[i];
}
}
}
}
int main() {
ios::sync_with_stdio(false);
cin.tie(nullptr);
memset(board, 0, sizeof(board));
memset(row_mask, 0, sizeof(row_mask));
memset(col_mask, 0, sizeof(col_mask));
memset(box_mask, 0, sizeof(box_mask));
read_relations();
dfs();
for (int i = 0; i < 9; i++) {
for (int j = 0; j < 9; j++) {
if (j) {
cout << ' ';
}
cout << board[i][j];
}
cout << '\n';
}
return 0;
}brute.cpp 每次找一个空格,现算它当前能填哪些数字,然后回溯尝试。
这个写法适合帮助理解题目规则,但如果只做裸搜,分支会偏大。
正式解的关键是:把“搜索”和“传播”分开。
对每个格子,我们用一个 9 位二进制掩码表示它当前还能填哪些数字:
- 第
d位是1,表示数字d还能填 - 是
0,表示这个数字已经被排除
然后在 DFS 前,反复做三类传播:
-
单点定值传播
如果某个格子只剩一个候选数字,那么同行、同列、同宫其它格子都不能再取这个数字。
-
不等式传播
如果某两个相邻格子满足
u < v,那么:u中所有“找不到更大配对值”的候选要删掉;v中所有“找不到更小配对值”的候选也要删掉。
-
隐藏单点
在某一行、某一列、某一宫里,如果某个数字只剩一个位置能放,就把那个位置直接定下来。
传播结束后:
- 如果出现空候选集合,当前分支无解;
- 如果所有格子都只剩一个候选,说明已经找到解;
- 否则选“候选数最少”的格子继续分支。
这就是典型的 MRV 剪枝:优先处理最受限制的变量。
这里还有一个很容易写错的输入细节:
- 题目并不是给出所有左右相邻格子的关系
- 而是每个
3x3宫内部那12对相邻格子的关系
所以跨宫边界的相邻格子之间,其实没有额外大小关系,代码里不能误加约束。
代码
cpp
#include <bits/stdc++.h>
using namespace std;
const int FULL_MASK = (1 << 9) - 1;
char h_rel[9][8]; // 同一行里相邻格子的左右大小关系
char v_rel[8][9]; // 同一列里相邻格子的上下大小关系
int peers[81][24];
int peer_cnt[81];
struct Arc {
int u, v;
char rel; // '<' 表示 u < v,'>' 表示 u > v
};
Arc arcs[200];
int arc_cnt;
int unit_cells[27][9];
int answer_mask[81];
int bit_count(int x) {
return __builtin_popcount((unsigned int)x);
}
int lowbit_to_digit(int bit) {
for (int d = 1; d <= 9; d++) {
if (bit == (1 << (d - 1))) {
return d;
}
}
return 0;
}
void add_peer(int u, int v) {
for (int i = 0; i < peer_cnt[u]; i++) {
if (peers[u][i] == v) {
return;
}
}
peers[u][peer_cnt[u]++] = v;
}
void build_units_and_peers() {
memset(peer_cnt, 0, sizeof(peer_cnt));
int id = 0;
for (int r = 0; r < 9; r++) {
for (int c = 0; c < 9; c++) {
unit_cells[id][c] = r * 9 + c;
}
id++;
}
for (int c = 0; c < 9; c++) {
for (int r = 0; r < 9; r++) {
unit_cells[id][r] = r * 9 + c;
}
id++;
}
for (int br = 0; br < 9; br += 3) {
for (int bc = 0; bc < 9; bc += 3) {
int pos = 0;
for (int r = br; r < br + 3; r++) {
for (int c = bc; c < bc + 3; c++) {
unit_cells[id][pos++] = r * 9 + c;
}
}
id++;
}
}
for (int u = 0; u < 81; u++) {
for (int t = 0; t < 27; t++) {
int found = -1;
for (int j = 0; j < 9; j++) {
if (unit_cells[t][j] == u) {
found = j;
break;
}
}
if (found == -1) {
continue;
}
for (int j = 0; j < 9; j++) {
int v = unit_cells[t][j];
if (v != u) {
add_peer(u, v);
}
}
}
}
}
void build_arcs() {
arc_cnt = 0;
for (int r = 0; r < 9; r++) {
for (int c = 0; c < 8; c++) {
if (h_rel[r][c] == 0) {
continue;
}
int u = r * 9 + c;
int v = r * 9 + c + 1;
arcs[arc_cnt++] = {u, v, h_rel[r][c]};
}
}
for (int r = 0; r < 8; r++) {
for (int c = 0; c < 9; c++) {
if (v_rel[r][c] == 0) {
continue;
}
int u = r * 9 + c;
int v = (r + 1) * 9 + c;
if (v_rel[r][c] == '^') {
arcs[arc_cnt++] = {u, v, '<'};
} else {
arcs[arc_cnt++] = {u, v, '>'};
}
}
}
}
// 根据不等式关系,删去不可能的候选数字。
bool revise_arc(int mask_u, int mask_v, char rel, int &new_u, int &new_v) {
new_u = 0;
new_v = 0;
for (int du = 1; du <= 9; du++) {
int bit_u = 1 << (du - 1);
if ((mask_u & bit_u) == 0) {
continue;
}
bool ok = false;
for (int dv = 1; dv <= 9; dv++) {
int bit_v = 1 << (dv - 1);
if ((mask_v & bit_v) == 0) {
continue;
}
if ((rel == '<' && du < dv) || (rel == '>' && du > dv)) {
ok = true;
break;
}
}
if (ok) {
new_u |= bit_u;
}
}
for (int dv = 1; dv <= 9; dv++) {
int bit_v = 1 << (dv - 1);
if ((mask_v & bit_v) == 0) {
continue;
}
bool ok = false;
for (int du = 1; du <= 9; du++) {
int bit_u = 1 << (du - 1);
if ((mask_u & bit_u) == 0) {
continue;
}
if ((rel == '<' && du < dv) || (rel == '>' && du > dv)) {
ok = true;
break;
}
}
if (ok) {
new_v |= bit_v;
}
}
return new_u != 0 && new_v != 0;
}
// 反复做三类传播:
// 1. 单点定值后删同行/同列/同宫候选
// 2. 按不等式删候选
// 3. 行列宫里的隐藏单点
bool propagate(int dom[]) {
while (true) {
bool changed = false;
for (int u = 0; u < 81; u++) {
if (dom[u] == 0) {
return false;
}
if (bit_count(dom[u]) != 1) {
continue;
}
int bit = dom[u];
for (int i = 0; i < peer_cnt[u]; i++) {
int v = peers[u][i];
if ((dom[v] & bit) == 0) {
continue;
}
int new_mask = dom[v] ^ bit;
if (new_mask == 0) {
return false;
}
if (new_mask != dom[v]) {
dom[v] = new_mask;
changed = true;
}
}
}
for (int i = 0; i < arc_cnt; i++) {
int u = arcs[i].u;
int v = arcs[i].v;
int new_u, new_v;
if (!revise_arc(dom[u], dom[v], arcs[i].rel, new_u, new_v)) {
return false;
}
if (new_u != dom[u]) {
dom[u] = new_u;
changed = true;
}
if (new_v != dom[v]) {
dom[v] = new_v;
changed = true;
}
}
for (int t = 0; t < 27; t++) {
for (int digit = 1; digit <= 9; digit++) {
int bit = 1 << (digit - 1);
int pos = -1;
int cnt = 0;
for (int j = 0; j < 9; j++) {
int u = unit_cells[t][j];
if (dom[u] & bit) {
cnt++;
pos = u;
}
}
if (cnt == 0) {
return false;
}
if (cnt == 1 && dom[pos] != bit) {
dom[pos] = bit;
changed = true;
}
}
}
if (!changed) {
return true;
}
}
}
bool dfs(int dom[]) {
int cur[81];
for (int i = 0; i < 81; i++) {
cur[i] = dom[i];
}
if (!propagate(cur)) {
return false;
}
int best = -1;
int best_cnt = 10;
for (int i = 0; i < 81; i++) {
int cnt = bit_count(cur[i]);
if (cnt > 1 && cnt < best_cnt) {
best_cnt = cnt;
best = i;
}
}
if (best == -1) {
for (int i = 0; i < 81; i++) {
answer_mask[i] = cur[i];
}
return true;
}
int mask = cur[best];
while (mask != 0) {
int bit = mask & (-mask);
mask -= bit;
int next_dom[81];
for (int i = 0; i < 81; i++) {
next_dom[i] = cur[i];
}
next_dom[best] = bit;
if (dfs(next_dom)) {
return true;
}
}
return false;
}
void read_relations() {
vector<string> token_lines;
string line;
while ((int)token_lines.size() < 15 && getline(cin, line)) {
if (line.empty()) {
continue;
}
token_lines.push_back(line);
}
for (int i = 0; i < 9; i++) {
for (int j = 0; j < 8; j++) {
h_rel[i][j] = 0;
}
}
for (int i = 0; i < 8; i++) {
for (int j = 0; j < 9; j++) {
v_rel[i][j] = 0;
}
}
for (int idx = 0; idx < 15; idx++) {
vector<char> arr;
for (int p = 0; p < (int)token_lines[idx].size(); p++) {
char ch = token_lines[idx][p];
if (ch == '<' || ch == '>' || ch == '^' || ch == 'v') {
arr.push_back(ch);
}
}
if ((int)arr.size() == 6) {
int row = (idx / 5) * 3 + ((idx % 5) / 2);
int cols[6] = {0, 1, 3, 4, 6, 7};
for (int i = 0; i < 6; i++) {
h_rel[row][cols[i]] = arr[i];
}
} else {
int row = (idx / 5) * 3 + ((idx % 5 - 1) / 2);
for (int i = 0; i < 9; i++) {
v_rel[row][i] = arr[i];
}
}
}
}
int main() {
ios::sync_with_stdio(false);
cin.tie(nullptr);
read_relations();
build_units_and_peers();
build_arcs();
int dom[81];
for (int i = 0; i < 81; i++) {
dom[i] = FULL_MASK;
}
if (!dfs(dom)) {
return 0;
}
for (int r = 0; r < 9; r++) {
for (int c = 0; c < 9; c++) {
if (c) {
cout << ' ';
}
cout << lowbit_to_digit(answer_mask[r * 9 + c]);
}
cout << '\n';
}
return 0;
}复杂度
这题本质上还是搜索题,理论最坏复杂度难以简单写成闭式,最坏仍然是指数级。
但由于:
- 棋盘固定只有
81格 - 每格候选最多
9个 - 传播与 MRV 剪枝都很强
实际运行足够快。
空间复杂度是
总结
这题的重点不是“把数独暴力搜出来”,而是要主动利用约束传播:
- 行列宫唯一性会删候选
- 大小关系也会删候选
- 某些数字在局部范围里会被逼成唯一位置
把这些传播和“候选数最少优先分支”结合起来,搜索树会小很多,代码也比较贴近普通竞赛风格的数独写法。
一图流解析
这张图把本题的建模、关键转移、实现检查和训练方法压缩到一页,适合读完正文后复盘。


