KM 算法学习笔记

· · 算法·理论

AI 使用声明:文章由作者自行完成。使用 AI 查找并更正了一些错误。

前言

个人的感觉是 KM 算法模板题的题解区中,少有可以给我讲透这个算法的题解,因此写了这一篇。

因为原题已经无法提交题解了,所以交算法理论。

回顾增广路理论

让我们回顾二分图的最大匹配这个题。

在不带权的最大匹配中,你面对着一个二分图,其中已经有一些边的两端被匹配。

你希望调整匹配,使得匹配数增加 1

一种平凡的情况是找一个两端都没有被匹配的边,然后匹配它的两个端点。

但是有时候不一定存在这样的边。这时你注意到这样的情况。

左1 - 右1
左2 - 右1(匹配边)
左2 - 右2

你发现这个时候,如果把 左2右1 解绑,就可以分别匹配 左1 - 右1左2 - 右2

同理地,情况可能是这样。

左1 - 右1
左2 - 右1(匹配边)
左2 - 右2
左3 - 右2(匹配边)
左3 - 右3

这类情况的特点是,存在一个简单路径,依次经过非匹配边 \to 匹配边 \to 非匹配边 \to 匹配边 \to \dots \to 非匹配边,且这条简单路径的两端都是未被匹配的点。

出现这种路径,你可以反转每条边的匹配状态,得到一个新的匹配,而且新匹配的匹配数比原匹配大 1

这就是增广路理论

定理 1 对于二分图的一个匹配,其不存在增广路当且仅当其是最大匹配。

证明 充分性:最大匹配如果有增广路,则反转增广路中边的状态就会得到更大的匹配,与当前匹配是最大匹配矛盾。

必要性:假如存在一个没有增广路且非最大的匹配 M,那么找到一个比它边数多的匹配 M'

考虑一个新的二分图 G'。对于 M 中的每条边,在 G' 上连接该边;对于 M' 中的每条边,在 G' 上连接该边。可能会有重边,这不重要。

因为在一个匹配中,每个点的度数不超过 1,所以 G' 中每个点的度数不超过 2

这意味着它的每个连通块都是链或者环。读者可以试着手玩一下,任何非链非环的连通块一定会出现度数大于 2 的点。

对于两条 G' 上相邻的边,一定一个来自 M,一个来自 M',因为一个匹配中的边一定是不相邻的。

考虑环,因为图是二分图,所以每个环都有偶数条边,显然 MM' 占据这个环的边数相等。

(你可能在想 MM' 是否会有公共边。它们的公共边必然会形成二元环,因此不影响证明。)

对于链,MM' 占据某条链的边数之差,可能是 01-1

因为 M' 的边数多于 M,一定存在一条链,M' 占据的边数比 M 的边数多 1。那么这条链上的边分别属于 M',M,M',M,\dots,M'

考虑原图上的这条链!链上的点中,起点和终点必然不是 M 的匹配点,其他的点必然是 M 的匹配点。且对于一条链,属于 M' 的部分,不可能是 M 的匹配边。于是这条链其实是 M 的一条增广路!与 M 没有增广路矛盾。

证毕。

回顾匈牙利算法

匈牙利算法的思路是寻找增广路

放一个匈牙利算法的代码,可能不少人已经背下来了。

#include<bits/stdc++.h>
#define int long long
using namespace std;
int n,m,e;
vector<int> E[1000005];
int match[1000005];
int vis[1000005];
int dfs(int u){
    for (int v:E[u]){
        if (vis[v])continue;
        vis[v]=1;
        if (!match[v]||dfs(match[v])){
            match[v]=u;return 1;
        }
    }
    return 0;
}
signed main(){
    cin>>n>>m>>e;
    for (int i=1;i<=e;i++){
        int u,v;cin>>u>>v;
        E[u].push_back(v);
    }
    int ans=0;
    for (int i=1;i<=n;i++){
        for (int j=1;j<=m;j++)vis[j]=0;
        if (dfs(i))ans++;
    }
    cout<<ans;
}

算法流程是这样的:

每次找到一条增广路并处理,会新增被匹配的左部点和右部点各一个。已经匹配过的左部点或右部点不会变得未匹配。因此按照某个顺序枚举左部点各一次就可以了。

Q:从 u 出发没有找到增广路,以后就再也没有从 u 出发的增广路了吗?

A:是的。考虑它连接的所有点都被 u 之前的左部点抢占了,在它之后的左部点处理时,这些和 u 连接的右部点再也不会解放出来和 u 配对。

Q:运行完这个算法之后,就再也没有新的增广路了吗?

A:是的。考虑增广路必须从未匹配的点出发,而根据前一个问题,这些点出发都没有增广路。

Q:为什么有 vis 数组?

A:增广路必须是简单路径。必须避免找增广路时多次访问某个节点。此外,对于这份代码的实现,多次访问同一个节点会导致死循环。

带权二分图匹配问题

为了简化问题,我们在处理带权的二分图匹配的时候,一般给所有原图中没有的边连上边权为 - \infty 的新边,然后要求完美匹配,也就是点数不多于另一侧的那一侧的所有点必须要全部匹配上。

如果求出的最大权完美匹配包含边权为 -\infty 的边,说明不存在完美匹配。

Q:我求权值和最大增广路,好吗?

A:不。

这本质是费用流算法。实现出来是 O(n^4) 的,效率没有 KM 高。很难做到更好。这个思路不再延伸。

顶标理论

考虑另一个问题。

对于一个完全二分图,你需要给每个点赋一个整数作为权值,称为顶标

你需要设置顶标,使得对于二分图中的任何一条边,其两个端点的顶标之和不小于边的权值。这样的一组顶标称为可行顶标

你需要最小化可行顶标的所有顶标之和。

为了方便阐述,令满足权值等于其两个端点的顶标之和的边称为相等边

定理 2 若相等边构成的子图有完美匹配,那么此时的顶标之和等于原图的最大权完美匹配。

证明 考虑原图的任何一个完美匹配。它有 n 条边,因为每条边的权值不大于它的两个端点的顶标之和,而每个点恰好作为完美匹配中一条边的端点,所以这个匹配的权值不大于所有顶标之和。

考虑相等边构成的子图的一个完美匹配,它的权值显然是所有顶标之和,得证。

根据定理 2,我们想要求出原图的最大权完美匹配,产生了一个新的方向,就是构造一个满足条件的可行顶标。

KM 算法的核心思路

先令左部点的顶标为它连接的权值最大边,右部点的顶标为 0

这时不一定有完美匹配。我们需要在失配时尝试调整顶标,使得匹配数增加。

接下来的描述中,相等边构成的子图称为相等子图

考虑我们进行匈牙利算法,dfs 在当前相等子图寻找从 u 这个左部点出发的增广路时,发现找不到这样的增广路。

令这次遍历遍历到的左部点集合为 S,右部点集合为 T。令 S'S 以外的左部点,T'T 以外的右部点。

观察

我们想要扩展 T。这意味着我们需要调整顶标,使得 ST' 之间的某条边变成相等边。

此外,我们不能影响 ST 之间的边,否则很容易就影响到已有的匹配了。

你发现:

Q:S'T 之间的某些边被踢出相等边,就算不会影响 u 的匹配,但是不会影响 u 之后的点的匹配吗?

A:不会,后面的点如果想用掉这些被踢出的边,还会用同样的操作把它补回来。

选取一个恰当的 d 的数值,使得有边被添加进相等边,且依然满足所有边的端点顶标之和大于权值。

这样的 d 是唯一的,具体地,他的取值是 d_0 = \min_{i \in S,j \in T'} (a_i +b_j - f_{i,j}),其中 a_i 是左部点 i 的当前顶标,b_j 是右部点 j 的当前顶标,f_{i,j} 是边 (i,j) 的权值。

解释:

选择这个 d,进行操作之后,式子去到最小值的边会变成相等边,从而 ST 都会扩展。

Q:得出的 d 可能会是 0,并导致算法卡住吗?

A:不会。如果会,这条卡住算法的边其实是相等边,与 (S,T') 没有相等边矛盾。

Q:如何维护这个 d

A:对于每个右部点 v,维护 slack_v = \min_{u \in S}(a_u + b_v - f_{u,v})。这样 d 就是 \min_{v \in T'} slack_v。在 dfs 的过程中,可以顺便求出这样一个 slack 数组,每次找增广路前都清空。

在给 u 找增广路的时候,最多扩展 O(n) 次,就会成功找到一个未匹配的右部点在 T 中,从而成功找到增广路并匹配。

于是你成功得到了一个求二分图最大权匹配的算法!

让我们尝试分析它的复杂度。

乘起来,是 O(n^4)

::::info[代码]

#include<bits/stdc++.h>
#define int long long
using namespace std;
int n,m,a[505],b[505],f[505][505];
int match[505];
int s[505],t[505];
bool dfs(int u){
    s[u]=1;
    for (int v=1;v<=n;v++){
        if (t[v])continue;
        if (a[u]+b[v]==f[u][v]){
            t[v]=1;
            if (!match[v]||dfs(match[v])){
                match[v]=u;
                return 1;
            }
        }
    }
    return 0;
}
signed main(){
    cin>>n>>m;
    for (int i=1;i<=n;i++){
        for (int j=1;j<=n;j++)f[i][j]=-1e18;
        a[i]=-1e18,b[i]=0;
    }
    for (int i=1;i<=m;i++){
        int u,v;cin>>u>>v;
        cin>>f[u][v];
        a[u]=max(a[u],f[u][v]);
    }
    for (int i=1;i<=n;i++){
        for (int j=1;j<=n;j++)s[j]=t[j]=0;
        while (!dfs(i)){
            int d=1e18;
            for (int j=1;j<=n;j++){
                if (!s[j])continue;
                for (int k=1;k<=n;k++){
                    if (t[k])continue;
                    d=min(d,a[j]+b[k]-f[j][k]);
                }
            }
            for (int j=1;j<=n;j++){
                if (s[j])a[j]-=d;
                if (t[j])b[j]+=d;
            }
            for (int j=1;j<=n;j++)s[j]=t[j]=0;
        }
    }
    int ans=0;
    for (int i=1;i<=n;i++){
        ans+=a[i]+b[i];
    }
    cout<<ans<<"\n";
    for (int i=1;i<=n;i++){
        cout<<match[i]<<" ";
    }
    return 0;
}

::::

优化

你觉得这太劣了。每更新一次相等子图,就要把 ST 里面本来就有的点重新跑一遍,非常不优雅。

你心想,只要对新加入 S 的左部点跑增广路就可以了。然后你发现实现起来很困难,很多细节不好处理。例如回溯的时候更新 match,这里的回溯是残缺的无法更新。

思路是对的,但是我们不能利用回溯来维护当前的增广路了。

既然不能利用回溯,那么 dfs 的优势就消失了,我们改为使用 bfs。

两个问题:

先解决第一个问题。

维护 matchx_umatchy_v 分别表示与左部点 u 匹配的边和与右部点 v 匹配的边。

维护 pre_v,表示上一次使 slack_v 更新的左部点。

定理 3 在增广路 (u_0,v_0,u_1,v_1,\dots,u_k,v_k) 中,u_i = pre_{v_i}

证明 对于 i \not = k,搜索过程中就已经有 slack_{v_i}=0 了,因此 v_i 被访问到时,pre_{v_i} 被设置为使它入队的左部点,也就是交错路径的前一个点 u_i

对于 i=k,考虑两种方式得到的相等边。

第一种是搜索的时候搜到了 T 中的未匹配的 v_k。这种情况与 i \not = k 类似。

第二种是更新顶标之后发现 slack_{v_k}=0 了,从而找到了增广路。这时,因为 d 取得是最小 slack,而 pre_{v_k} 就是和 v_k 连接着最接近相等边的点,更新顶标之后,(pre_{v_k},v_k) 一定是相等边,而且由于它刚出现,它是非匹配边。而由于搜索过程的缘故,存在从 u_0pre_{v_k} 的交错路径,得证。

根据定理 3,当我们找到 v_k 之后,我们可以直接回溯得到整条增广路。因为有 u_i = pre_{v_i}v_i = matchx_{u_{i+1}}

像更新链表一样,在回溯时计算新的 matchxmatchy 即可。这样第一个问题解决了。

然后解决第二个问题。第二个问题其实就是,将 S 中的所有点顶标 -dT 中的点顶标 +d 之后,slack 会发生怎样的变化。

答案很简单,对于 v \in T'slack_v 减少 d

证明方式还是考虑 (S,T)(S',T')(S,T')(S',T) 这四类边。

理由是,(S',T') 根本不贡献答案。所以对于 v \in T'slack_v 本质就是 (S,T') 的贡献。而这一类边的贡献全都减少了 d

所有问题都解决了。于是算法成立。

这时,对于找 u 出发的增广路,由于每条边只会跑一次,所以找增广路这一部分复杂度是 O(n^2);而更新顶标和 slack 的部分,因为只会扩展 O(n) 次,每次扩展需要花 O(n) 时间更新,所以也是 O(n^2)

总复杂度是 O(n^3)

代码实现

如果对实现还是没有头绪,可以对照上面给出的算法过程观看代码。

本代码可以通过二分图最大权完美匹配模板题。

::::info[代码]

#include<bits/stdc++.h>
#define int long long
using namespace std;
int n,m,a[505],b[505],f[505][505];
int matchx[505],matchy[505],slack[505];
int s[505],t[505];
int pre[505];
void aug(int u){
    while (u){
        int t=matchx[pre[u]];
        matchx[pre[u]]=u;
        matchy[u]=pre[u];
        u=t;
    }
}
void bfs(int U){
    for (int i=1;i<=n;i++)s[i]=t[i]=0,slack[i]=1e18;
    queue<int> q;
    q.push(U);
    while (1){
        while (!q.empty()){
            int u=q.front();q.pop();
            s[u]=1;
            for (int v=1;v<=n;v++){
                if (a[u]+b[v]-f[u][v]<slack[v]){
                    slack[v]=a[u]+b[v]-f[u][v];
                    pre[v]=u;
                }
                if (t[v])continue;
                if (a[u]+b[v]==f[u][v]){
                    pre[v]=u;
                    t[v]=1;
                    if (!matchy[v]){
                        aug(v);
                        return;
                    }
                    else q.push(matchy[v]);
                }
            }
        }
        int d=1e18;
        for (int i=1;i<=n;i++){
            if (!t[i])d=min(d,slack[i]);
        }
        for (int i=1;i<=n;i++){
            if (s[i])a[i]-=d;
            if (t[i])b[i]+=d;
            if (!t[i])slack[i]-=d;
        }
        for (int i=1;i<=n;i++){
            if (!t[i]){
                if (slack[i]==0){
                    t[i]=1;
                    if (!matchy[i]){
                        aug(i);
                        return;
                    }
                    else q.push(matchy[i]);
                }
            }
        }
    }
}
signed main(){
    cin>>n>>m;
    for (int i=1;i<=n;i++){
        for (int j=1;j<=n;j++)f[i][j]=-1e18;
        a[i]=-1e18,b[i]=0;
    }
    for (int i=1;i<=m;i++){
        int u,v;cin>>u>>v;
        cin>>f[u][v];
        a[u]=max(a[u],f[u][v]);
    }
    for (int i=1;i<=n;i++){
        bfs(i);
    }
    int ans=0;
    for (int i=1;i<=n;i++){
        ans+=a[i]+b[i];
    }
    cout<<ans<<"\n";
    for (int i=1;i<=n;i++){
        cout<<matchy[i]<<" ";
    }
    return 0;
}

::::