15.图-最小生成树.md 15 KB

15. 图-最小生成树

概念

生成树(Spanning Tree):连通图 G 的一个极小连通子图,它包含 G 的全部 n 个顶点,但只有 n-1 条边,并且连通(无环)。一个连通图可以有多棵不同的生成树。

最小生成树(Minimum Spanning Tree, MST):在无向带权图(连通网)中,权值之和最小的那棵生成树。它在保证连通所有顶点的前提下,使总边权最小,常用于解决"成本最低的连通方案"问题(如铺设电缆、修路、组网)。

求 MST 有两个经典算法,都基于同一个贪心结论——切分定理:任意一个割(把顶点分成两部分),横跨该割的最短边必然属于某棵最小生成树。

  • Prim 算法(加点法):从任意一个顶点出发,维护"已在树中"的顶点集合,每一轮选择一条连接"树内顶点"与"树外顶点"的最短边,把该树外顶点拉进树中,直到所有顶点都进树。始终是一棵"逐渐长大的树"。
  • Kruskal 算法(加边法):把所有边按权值从小到大排序,依次考察每条边,若它的两个端点不在同一个连通分量里(加入后不会成环),就选它;否则丢弃。用并查集(Union-Find)高效判断"是否成环"。

两个算法的结果一定是同一棵(唯一时)或同总权值的最小生成树,只是选边的顺序不同。

并查集(Union-Find):一种支持"合并两个集合"和"查询两个元素是否同集合"的数据结构。常用数组实现:parent[x] 指向 x 的父结点,配合路径压缩按秩合并,使单次 find/union 近似 O(1)(反阿克曼函数)。

核心操作 / 算法

Prim(邻接矩阵,O(n²))

  1. 任选起点 s,dist[s] = 0,其余 dist[i] = ∞,全部未访问。
  2. 重复 n 次: a. 在未访问顶点中找 dist 最小者 u,标记访问,total += dist[u]; b. 用 u 松弛所有未访问邻点:dist[v] = min(dist[v], w(u,v))
  3. 返回 total 与选边(用 parent 数组记录每条边)。

Kruskal(O(e log e))

  1. 把所有边按权值升序排序。
  2. 初始化并查集(每个顶点自成一个集合)。
  3. 依次取边 (u,v,w):若 find(u) ≠ find(v),则选中该边并 union(u,v);否则成环丢弃。
  4. 选满 n-1 条边即得 MST。

复杂度分析

算法 时间复杂度 空间复杂度 说明
Prim(邻接矩阵) O(n²) O(n) 每轮扫一遍找最小 dist
Prim(二叉堆/优先队列优化) O((n+e)·log n) O(n+e) 用堆取最小,每条边可能触发一次堆更新
Kruskal(边排序) O(e·log e) O(n)(并查集) 主要开销在排序

为什么

  • Prim 每轮用 O(n) 扫描选最小 dist,共 n 轮 → O(n²);堆优化后取最小 O(log n),加边 n-1 次,每条边至多入堆一次并更新 → O((n+e)·log n)。
  • Kruskal 先对所有 e 条边排序 → O(e·log e);随后每条边一次 find/union,带路径压缩与按秩合并且近似 O(1),可忽略 → 总 O(e·log e)。适用于边较少的稀疏图;Prim 邻接矩阵版适用于稠密图

语言实现

下面 4 种语言的实现演示相同的操作:对同一个无向带权图(6 个顶点、10 条边)分别运行 PrimKruskal,打印所选最小生成树边及总权值。Kruskal 中各自实现一个简单的并查集

C

#include <stdio.h>
#include <stdlib.h>

#define MAX 20
#define INF 0x3f3f3f3f

int n;
int graph[MAX][MAX];      // 邻接矩阵,无向带权图

// ---------- 并查集(供 Kruskal 判环) ----------
int parent[MAX], rnk[MAX];

int find(int x) {
    if (parent[x] != x)
        parent[x] = find(parent[x]);   // 路径压缩
    return parent[x];
}

void unionSet(int a, int b) {          // 按秩合并
    a = find(a); b = find(b);
    if (a == b) return;
    if (rnk[a] < rnk[b]) { int t = a; a = b; b = t; }
    parent[b] = a;
    if (rnk[a] == rnk[b]) rnk[a]++;
}

// ---------- Prim:加点法,邻接矩阵,O(n^2) ----------
void prim() {
    int dist[MAX], mstParent[MAX], visited[MAX];
    for (int i = 0; i < n; i++) {
        dist[i] = INF;
        visited[i] = 0;
        mstParent[i] = -1;
    }
    dist[0] = 0;                        // 从顶点 0 出发
    int total = 0;
    for (int k = 0; k < n; k++) {
        // 选未访问且 dist 最小的顶点
        int u = -1, min = INF;
        for (int i = 0; i < n; i++)
            if (!visited[i] && dist[i] < min) { min = dist[i]; u = i; }
        if (u == -1) break;             // 图不连通
        visited[u] = 1;
        total += dist[u];
        if (mstParent[u] != -1)
            printf("(%d, %d) weight %d\n", mstParent[u], u, dist[u]);
        // 用 u 松弛未访问邻点
        for (int v = 0; v < n; v++)
            if (!visited[v] && graph[u][v] < dist[v]) {
                dist[v] = graph[u][v];
                mstParent[v] = u;
            }
    }
    printf("Prim 总权重: %d\n", total);
}

// ---------- 边结构体(供 Kruskal) ----------
typedef struct {
    int u, v, w;
} Edge;
Edge edges[MAX * MAX];
int edgeCount;

int cmpEdge(const void *a, const void *b) {
    return ((Edge *)a)->w - ((Edge *)b)->w;
}

// ---------- Kruskal:加边法,排序 + 并查集,O(e log e) ----------
void kruskal() {
    qsort(edges, edgeCount, sizeof(Edge), cmpEdge);
    for (int i = 0; i < n; i++) {
        parent[i] = i;
        rnk[i] = 0;
    }
    int total = 0, cnt = 0;
    for (int i = 0; i < edgeCount && cnt < n - 1; i++) {
        int fu = find(edges[i].u), fv = find(edges[i].v);
        if (fu != fv) {                 // 不成环才选
            unionSet(fu, fv);
            printf("(%d, %d) weight %d\n", edges[i].u, edges[i].v, edges[i].w);
            total += edges[i].w;
            cnt++;
        }
    }
    printf("Kruskal 总权重: %d\n", total);
}

int main() {
    n = 6;
    for (int i = 0; i < n; i++)
        for (int j = 0; j < n; j++)
            graph[i][j] = (i == j) ? 0 : INF;

    // 无向带权图的边 (u, v, w)
    int raw[][3] = {{0,1,6},{0,2,1},{0,3,5},{1,2,5},{1,4,3},
                    {2,3,5},{2,4,6},{2,5,4},{3,5,2},{4,5,6}};
    edgeCount = 0;
    for (int i = 0; i < 10; i++) {
        int u = raw[i][0], v = raw[i][1], w = raw[i][2];
        graph[u][v] = graph[v][u] = w;
        edges[edgeCount].u = u;
        edges[edgeCount].v = v;
        edges[edgeCount].w = w;
        edgeCount++;
    }

    printf("=== Prim ===\n");
    prim();
    printf("=== Kruskal ===\n");
    kruskal();
    return 0;
}

C++

#include <iostream>
#include <vector>
#include <algorithm>
using namespace std;

const int INF = 0x3f3f3f3f;

struct Edge {
    int u, v, w;
};

// 并查集
class DSU {
    vector<int> parent, rnk;
public:
    DSU(int n) {
        parent.resize(n);
        rnk.resize(n, 0);
        for (int i = 0; i < n; i++) parent[i] = i;
    }
    int find(int x) {
        if (parent[x] != x) parent[x] = find(parent[x]);   // 路径压缩
        return parent[x];
    }
    void unite(int a, int b) {                             // 按秩合并
        a = find(a); b = find(b);
        if (a == b) return;
        if (rnk[a] < rnk[b]) swap(a, b);
        parent[b] = a;
        if (rnk[a] == rnk[b]) rnk[a]++;
    }
};

// Prim:加点法,邻接矩阵,O(n^2)
void prim(const vector<vector<int>>& g, int n) {
    vector<int> dist(n, INF), mstParent(n, -1);
    vector<bool> visited(n, false);
    dist[0] = 0;                        // 从顶点 0 出发
    int total = 0;
    for (int k = 0; k < n; k++) {
        // 选未访问且 dist 最小的顶点
        int u = -1;
        for (int i = 0; i < n; i++)
            if (!visited[i] && (u == -1 || dist[i] < dist[u])) u = i;
        if (u == -1) break;             // 图不连通
        visited[u] = true;
        total += dist[u];
        if (mstParent[u] != -1)
            cout << "(" << mstParent[u] << ", " << u << ") weight " << dist[u] << endl;
        // 用 u 松弛未访问邻点
        for (int v = 0; v < n; v++)
            if (!visited[v] && g[u][v] < dist[v]) {
                dist[v] = g[u][v];
                mstParent[v] = u;
            }
    }
    cout << "Prim 总权重: " << total << endl;
}

// Kruskal:加边法,排序 + 并查集,O(e log e)
void kruskal(vector<Edge>& edges, int n) {
    sort(edges.begin(), edges.end(), [](const Edge& a, const Edge& b) {
        return a.w < b.w;
    });
    DSU dsu(n);
    int total = 0, cnt = 0;
    for (const Edge& e : edges) {
        if (dsu.find(e.u) != dsu.find(e.v)) {   // 不成环才选
            dsu.unite(e.u, e.v);
            cout << "(" << e.u << ", " << e.v << ") weight " << e.w << endl;
            total += e.w;
            if (++cnt == n - 1) break;
        }
    }
    cout << "Kruskal 总权重: " << total << endl;
}

int main() {
    int n = 6;
    vector<vector<int>> g(n, vector<int>(n, INF));
    for (int i = 0; i < n; i++) g[i][i] = 0;

    // 无向带权图的边 (u, v, w)
    vector<Edge> edges = {
        {0,1,6},{0,2,1},{0,3,5},{1,2,5},{1,4,3},
        {2,3,5},{2,4,6},{2,5,4},{3,5,2},{4,5,6}
    };
    for (const Edge& e : edges)
        g[e.u][e.v] = g[e.v][e.u] = e.w;

    cout << "=== Prim ===" << endl;
    prim(g, n);
    cout << "=== Kruskal ===" << endl;
    kruskal(edges, n);
    return 0;
}

Java

import java.util.*;

public class MST {

    // 边结构体
    static class Edge implements Comparable<Edge> {
        int u, v, w;
        Edge(int u, int v, int w) { this.u = u; this.v = v; this.w = w; }
        public int compareTo(Edge o) { return this.w - o.w; }
    }

    // 并查集
    static class DSU {
        int[] parent, rank;
        DSU(int n) {
            parent = new int[n];
            rank = new int[n];
            for (int i = 0; i < n; i++) parent[i] = i;
        }
        int find(int x) {
            if (parent[x] != x) parent[x] = find(parent[x]);   // 路径压缩
            return parent[x];
        }
        void unite(int a, int b) {                             // 按秩合并
            a = find(a); b = find(b);
            if (a == b) return;
            if (rank[a] < rank[b]) { int t = a; a = b; b = t; }
            parent[b] = a;
            if (rank[a] == rank[b]) rank[a]++;
        }
    }

    // Prim:加点法,邻接矩阵,O(n^2)
    static void prim(int[][] g, int n) {
        int[] dist = new int[n];
        int[] mstParent = new int[n];
        boolean[] visited = new boolean[n];
        Arrays.fill(dist, Integer.MAX_VALUE);
        Arrays.fill(mstParent, -1);
        dist[0] = 0;                        // 从顶点 0 出发
        int total = 0;
        for (int k = 0; k < n; k++) {
            // 选未访问且 dist 最小的顶点
            int u = -1;
            for (int i = 0; i < n; i++)
                if (!visited[i] && (u == -1 || dist[i] < dist[u])) u = i;
            if (u == -1) break;             // 图不连通
            visited[u] = true;
            total += dist[u];
            if (mstParent[u] != -1)
                System.out.println("(" + mstParent[u] + ", " + u + ") weight " + dist[u]);
            // 用 u 松弛未访问邻点
            for (int v = 0; v < n; v++)
                if (!visited[v] && g[u][v] < dist[v]) {
                    dist[v] = g[u][v];
                    mstParent[v] = u;
                }
        }
        System.out.println("Prim 总权重: " + total);
    }

    // Kruskal:加边法,排序 + 并查集,O(e log e)
    static void kruskal(List<Edge> edges, int n) {
        Collections.sort(edges);
        DSU dsu = new DSU(n);
        int total = 0, cnt = 0;
        for (Edge e : edges) {
            if (dsu.find(e.u) != dsu.find(e.v)) {   // 不成环才选
                dsu.unite(e.u, e.v);
                System.out.println("(" + e.u + ", " + e.v + ") weight " + e.w);
                total += e.w;
                if (++cnt == n - 1) break;
            }
        }
        System.out.println("Kruskal 总权重: " + total);
    }

    public static void main(String[] args) {
        int n = 6;
        int[][] g = new int[n][n];
        for (int[] row : g) Arrays.fill(row, Integer.MAX_VALUE);
        for (int i = 0; i < n; i++) g[i][i] = 0;

        // 无向带权图的边 (u, v, w)
        int[][] raw = {{0,1,6},{0,2,1},{0,3,5},{1,2,5},{1,4,3},
                       {2,3,5},{2,4,6},{2,5,4},{3,5,2},{4,5,6}};
        List<Edge> edges = new ArrayList<>();
        for (int[] e : raw) {
            g[e[0]][e[1]] = g[e[1]][e[0]] = e[2];
            edges.add(new Edge(e[0], e[1], e[2]));
        }

        System.out.println("=== Prim ===");
        prim(g, n);
        System.out.println("=== Kruskal ===");
        kruskal(edges, n);
    }
}

Python

class DSU:
    """并查集:用于 Kruskal 判环"""

    def __init__(self, n):
        self.parent = list(range(n))
        self.rank = [0] * n

    def find(self, x):
        """路径压缩"""
        if self.parent[x] != x:
            self.parent[x] = self.find(self.parent[x])
        return self.parent[x]

    def unite(self, a, b):
        """按秩合并"""
        a, b = self.find(a), self.find(b)
        if a == b:
            return
        if self.rank[a] < self.rank[b]:
            a, b = b, a
        self.parent[b] = a
        if self.rank[a] == self.rank[b]:
            self.rank[a] += 1


INF = float("inf")

# 无向带权图的边 (u, v, w)
raw_edges = [(0,1,6),(0,2,1),(0,3,5),(1,2,5),(1,4,3),
             (2,3,5),(2,4,6),(2,5,4),(3,5,2),(4,5,6)]
n = 6


def prim(g, n, start=0):
    """Prim:加点法,邻接矩阵,O(n^2)"""
    dist = [INF] * n
    mst_parent = [-1] * n
    visited = [False] * n
    dist[start] = 0
    total = 0
    for _ in range(n):
        # 选未访问且 dist 最小的顶点
        u = min((v for v in range(n) if not visited[v]), key=lambda v: dist[v])
        if dist[u] == INF:          # 图不连通
            break
        visited[u] = True
        total += dist[u]
        if mst_parent[u] != -1:
            print(f"({mst_parent[u]}, {u}) weight {dist[u]}")
        # 用 u 松弛未访问邻点
        for v in range(n):
            if not visited[v] and g[u][v] < dist[v]:
                dist[v] = g[u][v]
                mst_parent[v] = u
    print("Prim 总权重:", total)


def kruskal(edges, n):
    """Kruskal:加边法,排序 + 并查集,O(e log e)"""
    edges = sorted(edges, key=lambda e: e[2])
    dsu = DSU(n)
    total = 0
    cnt = 0
    for u, v, w in edges:
        if dsu.find(u) != dsu.find(v):      # 不成环才选
            dsu.unite(u, v)
            print(f"({u}, {v}) weight {w}")
            total += w
            cnt += 1
            if cnt == n - 1:
                break
    print("Kruskal 总权重:", total)


if __name__ == "__main__":
    # 邻接矩阵
    g = [[INF] * n for _ in range(n)]
    for i in range(n):
        g[i][i] = 0
    for u, v, w in raw_edges:
        g[u][v] = g[v][u] = w

    print("=== Prim ===")
    prim(g, n)
    print("=== Kruskal ===")
    kruskal(raw_edges, n)