[CQOI2013] 新数独

GitHub跳转原题关系图返回列表

用 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 前,反复做三类传播:

  1. 单点定值传播

    如果某个格子只剩一个候选数字,那么同行、同列、同宫其它格子都不能再取这个数字。

  2. 不等式传播

    如果某两个相邻格子满足 u < v,那么:

    • u 中所有“找不到更大配对值”的候选要删掉;
    • v 中所有“找不到更小配对值”的候选也要删掉。
  3. 隐藏单点

    在某一行、某一列、某一宫里,如果某个数字只剩一个位置能放,就把那个位置直接定下来。

传播结束后:

  • 如果出现空候选集合,当前分支无解;
  • 如果所有格子都只剩一个候选,说明已经找到解;
  • 否则选“候选数最少”的格子继续分支。

这就是典型的 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 剪枝都很强

实际运行足够快。

空间复杂度是 O(81)O(81) 量级。

总结

这题的重点不是“把数独暴力搜出来”,而是要主动利用约束传播:

  • 行列宫唯一性会删候选
  • 大小关系也会删候选
  • 某些数字在局部范围里会被逼成唯一位置

把这些传播和“候选数最少优先分支”结合起来,搜索树会小很多,代码也比较贴近普通竞赛风格的数独写法。

一图流解析

这张图把本题的建模、关键转移、实现检查和训练方法压缩到一页,适合读完正文后复盘。

一图流解析