交叉连接 - Problem - QOJ.ac

Statement

从 \(n\times m\) 的网格中选出格子的一个子集,使得:

  • 子集中的格子形成一个连通块。
  • 每个边界上都至少有一个格子存在于子集中。

求选出子集的最小点权和。

Analysis

容易想到图论建模跑最短路。若选一个中心格算四个方向的最短路,不仅会算重,而且忽略了一条路有可能绕路与另一条路会合然后走公共路径的情况。

那么最大的问题就是“会合”。所以如果能够找出答案中路径“会合”的关键点,问题似乎就解决了。

简单思考会发现,最终答案总可以找到两个关键点,且两个关键点都连向两个方向的边界,特别地允许关键点重叠。

那么考虑处理出每个点连向任意两个边界的最短长度。具体地,对于每个恰有两个 \(1\) 的 \(mask\in[0,16)\),计算出所有 \(dis[mask][u]\)。实现时可以先计算出恰有一个 \(1\) 的 \(mask\) 的答案,然后合并为有两个 \(1\) 的 \(mask\) 的答案。

然后需要找出两个关键点,使得 \(dis[mask][x] + dist(x,y)+dis[15\oplus mask][y]\) 最小。

首先枚举 \(mask\)。若固定 \(y\),则 \(dis[mask][x] + dist(x, y)\) 可以看作起点具有特殊权值的单源最短路。但每次从 \(y\) 开始跑最短路显然是劣的,我们希望只跑一次最短路就可以得到所有 \(y\) 的答案。

观察 \(dis[mask][x] + dist(x,y)\),发现起点 \(x\) 有可能在任意一处,但以它为起点的路径总是固定的。换句话说,一个 \(y\) 的路径继续走下去还可以得到其他 \(y\) 的路径,每次都重新跑是浪费的。

所以建立超级源点 \(s\) 连向所有网格点 \(x\),边权为 \(dist[mask][x] + a[x]\),网格点间边 \(u \to v\) 的权为 \(a_v\)。从 \(s\) 开始跑单源最短路,便可得到任意 \(y\) 的最小的 \(dis[mask][x] + dist(x, y)\),加上 \(dis[15\oplus mask][y]\) 即可用来更新最终答案。

复杂度只有 Dijkstra 的 \(O(nm\log nm)\),但常数极大。

Optimization

注意到点权很小,考虑采用桶来代替 Dijkstra 的优先队列,即 Dial 算法。

观察 \(dis\) 和 \(dist\) 的大小,发现最坏情况下会走 \(n+m\) 步,总权值为 \(35(n+m)\);\(dist\) 由于加上了 \(dis\) 总权值最大为 \(70(n+m)\)。

具体地,考虑维护当前遍历到的桶 \(cur\) 和已处理点数 \(p\):

  • 若 \(p\) 大于等于总点数,说明算法结束,直接退出。
  • 否则从 \(cur\) 开始找到往后第一个非空的桶,取出桶中的点,然后进行松弛操作。

容易发现往桶中加入的路径长是单调不降的,故时间复杂度为 \(O(nm+k(n+m))\),其中 \(k=70\),可以通过。

Code

// https://qoj.ac/problem/17855/statement/zh_cn
#include<bits/stdc++.h>
#define ll long long
#define div cerr << "--------------------\n"
#define endl cout << '\n'
#define DEBUG(xxx) cerr << #xxx << " = " << xxx << '\n'
#define st first
#define nd second
#define pii pair<int, int>
#define pll pair<ll, ll>
#define pushb push_back
using namespace std;
constexpr int inf = 1e9, N = 1e3 + 10, M = N * N, d[4][2] = {1, 0, -1, 0, 0, 1, 0, -1}, K = 140 * N;
constexpr ll INF = 1e18;
int n, m, a[M], id[N][N];
vector<int> G[M];
vector<pll> T[M];
void link(int u, int v){
    G[u].pushb(v);
    T[u].pushb({v, a[v]});
}
ll dis[20][M], dist[M];
vector<int> buc[K];
bool vis[M];
void Dij1(int j){
    memset(vis, 0, sizeof(vis));
    int top = 35 * (n + m) + 1;
    for(int i = 0; i < top; i++) buc[i].clear();
    int s = n * m;
    if(j == 1) s += 1;
    else if(j == 2) s += 2;
    else if(j == 4) s += 3;
    else s += 4;
    dis[j][s] = 0;
    buc[0].pushb(s);
    int cur = 0, p = 0;
    while(p < n * m + 1){
        while(buc[cur].empty()) cur++;
        int u = buc[cur].back(); buc[cur].pop_back();
        if(vis[u]) continue;
        vis[u] = 1;
        p++;
        for(auto v : G[u])
            if(dis[j][v] > dis[j][u] + a[u]){
                dis[j][v] = dis[j][u] + a[u];
                buc[dis[j][v]].pushb(v);
            }
    }
}
void Dij2(int srt){
    memset(vis, 0, sizeof(vis));
    memset(dist, 0x3f, sizeof(dist));
    int top = 70 * (n + m) + 1;
    for(int i = 0; i < top; i++) buc[i].clear();
    dist[srt] = 0;
    buc[0].pushb(srt);
    int cur = 0, p = 0;
    while(p < n * m + 1){
        while(buc[cur].empty()) cur++;
        int u = buc[cur].back(); buc[cur].pop_back();
        if(vis[u]) continue;
        vis[u] = 1;
        p++;
        for(auto [v, w] : T[u])
            if(!vis[v] && dist[v] > dist[u] + w){
                dist[v] = dist[u] + w;
                buc[dist[v]].pushb(v);
            }
    }
}
int main(){
    // freopen("crosslink.in", "r", stdin);
    // freopen("crosslink.out", "w", stdout);
    ios :: sync_with_stdio(0), cin.tie(0), cout.tie(0);
    cin >> n >> m;
    for(int i = 1; i <= n; i++) for(int j = 1; j <= m; j++) id[i][j] = (i - 1) * m + j;
    for(int i = 1; i <= n; i++)
        for(int j = 1; j <= m; j++){
            char c; cin >> c;
            if('0' <= c && c <= '9') a[id[i][j]] = c - '0';
            else a[id[i][j]] = 10 + c - 'A';
        }
    for(int i = 1; i <= n; i++)
        for(int j = 1; j <= m; j++)
            for(int k = 0; k < 4; k++){
                int x = i + d[k][0], y = j + d[k][1];
                if(x < 1 || y < 1 || x > n || y > m) continue;
                link(id[i][j], id[x][y]), link(id[x][y], id[i][j]);
            }
    for(int j = 1; j <= m; j++) link(n * m + 1, id[1][j]), link(n * m + 3, id[n][j]);
    for(int i = 1; i <= n; i++) link(n * m + 2, id[i][1]), link(n * m + 4, id[i][m]);
    memset(dis, 0x3f, sizeof(dis));
    for(int j = 1; j < 16; j *= 2) Dij1(j);
    for(int j = 1; j < 16; j++){
        if(__builtin_popcount(j) != 2) continue;
        int p = -1, q = -1;
        for(int k = 0; k < 4; k++)
            if(j & (1 << k))
                if(p == -1) p = k;
                else q = k;
        for(int i = 1; i <= n * m; i++) dis[j][i] = dis[1 << p][i] + dis[1 << q][i];
    }
    ll ans = INF;
    for(int j = 1; j < 16; j++){
        if(__builtin_popcount(j) != 2) continue;
        T[0].clear();
        for(int i = 1; i <= n * m; i++) T[0].pushb({i, dis[j][i] + a[i]});
        Dij2(0);
        for(int i = 1; i <= n * m; i++) ans = min(ans, dist[i] + dis[0b1111 ^ j][i]);
    }
    cout << ans;
    return 0;
}