12.堆.md 12 KB

12. 堆

概念

堆(Heap) 是一种特殊的完全二叉树,并且满足堆序性

  • 最大堆(大顶堆):每个结点的值 不小于 其孩子的值,因此堆顶(根)是最大值
  • 最小堆(小顶堆):每个结点的值 不大于 其孩子的值,因此堆顶是最小值

用数组存储:因为堆是完全二叉树,可以紧凑地用数组存储,无需指针。若数组下标从 0 开始,对下标为 i 的结点:

  • 父结点下标:(i - 1) / 2
  • 左孩子下标:2 * i + 1
  • 右孩子下标:2 * i + 2

(若下标从 1 开始,则父为 i/2,左子 2i,右子 2i+1。本文统一使用从 0 开始。)

性质:堆只保证堆顶是最大/最小值,并不保证整个序列有序;兄弟结点之间没有大小关系约束。堆序性只沿"祖先-子孙"路径成立。

适用场景

  • 优先队列(Priority Queue):每次取最值,插入任意值,都是 O(log n)。
  • 堆排序:利用堆顶最值,反复取出堆顶排序,O(n log n)。
  • TopK 问题:如求 n 个数中最大的 k 个,用大小为 k 的最小堆维护,扫描一遍即可。

核心操作

  • 建堆(自底向上):从最后一个非叶结点(下标 n/2 - 1)开始,向前对每个结点做"下沉"(sift down),最终整个数组满足堆序。时间复杂度 O(n)
  • 插入(上浮,sift up):把新元素放到数组末尾,然后不断与父结点比较,若违反堆序则与父交换,向上浮动到合适位置。
  • 删除堆顶(下沉):把堆顶与最后一个元素交换,删除末尾,然后对新的堆顶做"下沉"(sift down),恢复堆序。通常配合返回被删的最大/最小值。
  • 取堆顶:直接返回数组下标 0 的元素,O(1)。

复杂度分析

操作 时间复杂度 空间复杂度
建堆(自底向上) O(n) O(1)(原地)
插入(上浮) O(log n) O(1)
删除堆顶(下沉) O(log n) O(1)
取堆顶 O(1) O(1)

为什么

  • 建堆时从最后一个非叶结点自底向上做下沉。每个结点下沉的代价正比于它下降的高度,但大部分结点位于树的下层、下沉次数少,累加起来总和为 O(n)(这是与逐次插入 O(n log n) 的关键区别)。
  • 插入上浮、删除下沉都只沿"根到叶子"的一条路径移动,路径长度 = 树高 = O(log n),每步交换 O(1),故 O(log n)。
  • 取堆顶只读数组首元素,O(1)。所有操作均为原地、无需额外数组,故空间 O(1)(递归下沉栈空间不计入)。

语言实现

下面是 4 种语言的完整最大堆实现,均演示:建堆、插入、删除最大、取最大。Python 额外提及内置的 heapq 模块(其默认实现为最小堆)。

C

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

// 最大堆结构体
typedef struct {
    int *arr;    // 数组
    int size;    // 当前元素个数
    int cap;     // 容量
} MaxHeap;

// 从 0 开始的父子下标
int parent(int i)  { return (i - 1) / 2; }
int left(int i)    { return 2 * i + 1; }
int right(int i)   { return 2 * i + 2; }

MaxHeap *createHeap(int cap) {
    MaxHeap *h = (MaxHeap *)malloc(sizeof(MaxHeap));
    h->arr = (int *)malloc(sizeof(int) * cap);
    h->size = 0;
    h->cap = cap;
    return h;
}

void swap(int *a, int *b) { int t = *a; *a = *b; *b = t; }

// 下沉:把下标 i 的元素下移到合适位置
void siftDown(MaxHeap *h, int i) {
    int largest = i;
    int l = left(i), r = right(i);
    if (l < h->size && h->arr[l] > h->arr[largest]) largest = l;
    if (r < h->size && h->arr[r] > h->arr[largest]) largest = r;
    if (largest != i) {
        swap(&h->arr[i], &h->arr[largest]);
        siftDown(h, largest);
    }
}

// 上浮:把下标 i 的元素上移到合适位置
void siftUp(MaxHeap *h, int i) {
    while (i > 0 && h->arr[parent(i)] < h->arr[i]) {
        swap(&h->arr[parent(i)], &h->arr[i]);
        i = parent(i);
    }
}

// 建堆:从最后一个非叶结点开始自底向上下沉
void buildHeap(MaxHeap *h, int *a, int n) {
    for (int i = 0; i < n; i++) h->arr[i] = a[i];
    h->size = n;
    for (int i = n / 2 - 1; i >= 0; i--)
        siftDown(h, i);
}

// 插入
void insert(MaxHeap *h, int val) {
    if (h->size >= h->cap) return;
    h->arr[h->size] = val;
    siftUp(h, h->size);
    h->size++;
}

// 取最大(堆顶)
int peek(MaxHeap *h) {
    return h->arr[0];
}

// 删除最大并返回
int extractMax(MaxHeap *h) {
    int max = h->arr[0];
    h->arr[0] = h->arr[h->size - 1];
    h->size--;
    siftDown(h, 0);
    return max;
}

void printHeap(MaxHeap *h) {
    for (int i = 0; i < h->size; i++) printf("%d ", h->arr[i]);
    printf("\n");
}

int main() {
    int a[] = {3, 1, 6, 5, 2, 4};
    int n = 6;
    MaxHeap *h = createHeap(20);
    buildHeap(h, a, n);            // 建堆
    printf("建堆后: ");
    printHeap(h);

    insert(h, 10);                 // 插入
    printf("插入 10 后: ");
    printHeap(h);

    printf("堆顶(最大): %d\n", peek(h));

    printf("删除最大 %d 后: ", extractMax(h));
    printHeap(h);
    printf("再次删除最大 %d 后: ", extractMax(h));
    printHeap(h);
    return 0;
}

C++

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

// 最大堆(基于 vector)
class MaxHeap {
private:
    vector<int> arr;
    int parent(int i) { return (i - 1) / 2; }
    int left(int i)   { return 2 * i + 1; }
    int right(int i)  { return 2 * i + 2; }

    // 下沉
    void siftDown(int i) {
        int n = arr.size();
        int largest = i;
        int l = left(i), r = right(i);
        if (l < n && arr[l] > arr[largest]) largest = l;
        if (r < n && arr[r] > arr[largest]) largest = r;
        if (largest != i) {
            swap(arr[i], arr[largest]);
            siftDown(largest);
        }
    }

    // 上浮
    void siftUp(int i) {
        while (i > 0 && arr[parent(i)] < arr[i]) {
            swap(arr[parent(i)], arr[i]);
            i = parent(i);
        }
    }

public:
    MaxHeap() {}
    // 用数组建堆
    MaxHeap(const vector<int>& a) {
        arr = a;
        for (int i = (int)arr.size() / 2 - 1; i >= 0; i--)
            siftDown(i);
    }
    void insert(int val) {
        arr.push_back(val);
        siftUp(arr.size() - 1);
    }
    int peek() const { return arr[0]; }
    int extractMax() {
        int maxv = arr[0];
        arr[0] = arr.back();
        arr.pop_back();
        siftDown(0);
        return maxv;
    }
    void print() const {
        for (int x : arr) cout << x << " ";
        cout << endl;
    }
};

int main() {
    MaxHeap h({3, 1, 6, 5, 2, 4});   // 建堆
    cout << "建堆后: ";
    h.print();

    h.insert(10);                     // 插入
    cout << "插入 10 后: ";
    h.print();

    cout << "堆顶(最大): " << h.peek() << endl;

    cout << "删除最大 " << h.extractMax() << " 后: ";
    h.print();
    cout << "再次删除最大 " << h.extractMax() << " 后: ";
    h.print();
    return 0;
}

Java

import java.util.*;

public class MaxHeap {
    private int[] arr;
    private int size;

    public MaxHeap(int cap) { arr = new int[cap]; size = 0; }

    // 从数组建堆
    public MaxHeap(int[] a) {
        arr = Arrays.copyOf(a, Math.max(a.length, 16));
        size = a.length;
        for (int i = size / 2 - 1; i >= 0; i--) siftDown(i);
    }

    private int parent(int i) { return (i - 1) / 2; }
    private int left(int i)   { return 2 * i + 1; }
    private int right(int i)  { return 2 * i + 2; }

    // 下沉
    private void siftDown(int i) {
        int largest = i;
        int l = left(i), r = right(i);
        if (l < size && arr[l] > arr[largest]) largest = l;
        if (r < size && arr[r] > arr[largest]) largest = r;
        if (largest != i) {
            int t = arr[i]; arr[i] = arr[largest]; arr[largest] = t;
            siftDown(largest);
        }
    }

    // 上浮
    private void siftUp(int i) {
        while (i > 0 && arr[parent(i)] < arr[i]) {
            int t = arr[parent(i)]; arr[parent(i)] = arr[i]; arr[i] = t;
            i = parent(i);
        }
    }

    public void insert(int val) {
        if (size == arr.length) arr = Arrays.copyOf(arr, arr.length * 2);
        arr[size] = val;
        siftUp(size);
        size++;
    }

    public int peek() { return arr[0]; }

    public int extractMax() {
        int max = arr[0];
        arr[0] = arr[size - 1];
        size--;
        siftDown(0);
        return max;
    }

    public void print() {
        for (int i = 0; i < size; i++) System.out.print(arr[i] + " ");
        System.out.println();
    }

    public static void main(String[] args) {
        MaxHeap h = new MaxHeap(new int[]{3, 1, 6, 5, 2, 4}); // 建堆
        System.out.print("建堆后: ");
        h.print();

        h.insert(10);                                          // 插入
        System.out.print("插入 10 后: ");
        h.print();

        System.out.println("堆顶(最大): " + h.peek());

        System.out.print("删除最大 " + h.extractMax() + " 后: ");
        h.print();
        System.out.print("再次删除最大 " + h.extractMax() + " 后: ");
        h.print();
    }
}

Python

class MaxHeap:
    """最大堆(基于列表实现)"""

    def __init__(self, a=None):
        self.arr = list(a) if a else []
        if a:
            self._build()

    @staticmethod
    def _parent(i): return (i - 1) // 2
    @staticmethod
    def _left(i):   return 2 * i + 1
    @staticmethod
    def _right(i):  return 2 * i + 2

    def _sift_down(self, i):
        """下沉"""
        n = len(self.arr)
        while True:
            largest = i
            l, r = self._left(i), self._right(i)
            if l < n and self.arr[l] > self.arr[largest]:
                largest = l
            if r < n and self.arr[r] > self.arr[largest]:
                largest = r
            if largest == i:
                break
            self.arr[i], self.arr[largest] = self.arr[largest], self.arr[i]
            i = largest

    def _sift_up(self, i):
        """上浮"""
        while i > 0 and self.arr[self._parent(i)] < self.arr[i]:
            self.arr[self._parent(i)], self.arr[i] = self.arr[i], self.arr[self._parent(i)]
            i = self._parent(i)

    def _build(self):
        """自底向上建堆"""
        for i in range(len(self.arr) // 2 - 1, -1, -1):
            self._sift_down(i)

    def insert(self, val):
        self.arr.append(val)
        self._sift_up(len(self.arr) - 1)

    def peek(self):
        return self.arr[0]

    def extract_max(self):
        """删除并返回最大元素"""
        maxv = self.arr[0]
        self.arr[0] = self.arr[-1]
        self.arr.pop()
        self._sift_down(0)
        return maxv

    def __repr__(self):
        return " ".join(map(str, self.arr))


if __name__ == "__main__":
    h = MaxHeap([3, 1, 6, 5, 2, 4])     # 建堆
    print("建堆后:", h)

    h.insert(10)                         # 插入
    print("插入 10 后:", h)

    print("堆顶(最大):", h.peek())
    print("删除最大", h.extract_max(), "后:", h)
    print("再次删除最大", h.extract_max(), "后:", h)

    # Python 内置 heapq 模块:默认是最小堆
    import heapq
    data = [3, 1, 6, 5, 2, 4]
    heapq.heapify(data)                 # 原地建最小堆
    print("heapq 建堆:", data)
    heapq.heappush(data, 0)             # 插入
    print("heapq 插入 0 后:", data)
    print("heapq 弹出最小:", heapq.heappop(data))
    # 如需最大堆,可把元素取负存入