PriorityQueue

PriorityQueue 源码分析

PriorityQueue 是数组上的小根堆。刷题里 Top K、合并 K 个有序链表、数据流中位数(双堆)全靠它。
整个类没有一个”树节点”,完全二叉树被编码进数组下标:queue[k] 的父亲是 queue[(k-1)/2],两个孩子是 queue[2k+1]queue[2k+2]
堆序不变量只有一条:任何父亲不大于它的孩子,所以 queue[0] 永远是最小值。

代码块JAVA · 58 行收起展开
// 基于本地 JDK 源码 (D:/1ForCode/JAVA_Source, java.base, 2024 版), java.util.PriorityQueue
public class PriorityQueue<E> extends AbstractQueue<E>
    implements java.io.Serializable {

    // ...

    transient Object[] queue;       // 完全二叉树按层序摊平进数组, queue[0] 是堆顶

    int size;

    private final Comparator<? super E> comparator;  // null 则退回元素自身的 Comparable

    transient int modCount;         // fail-fast 计数, 迭代中 offer/poll 会让迭代器抛 CME

    public boolean offer(E e) {
        if (e == null)
            throw new NullPointerException();   // null 进堆第一次 compareTo 就会炸, 干脆入口拦掉
        modCount++;
        int i = size;
        if (i >= queue.length)
            grow(i + 1);
        siftUp(i, e);                           // 放到最后一个空位, 然后一路上浮
        size = i + 1;
        return true;
    }

    public E peek() {
        return (E) queue[0];                    // 空堆返回 null
    }

    public E poll() {
        final Object[] es;
        final E result;
        if ((result = (E) ((es = queue)[0])) != null) {
            modCount++;
            final int n;
            final E x = (E) es[(n = --size)];   // 拿最后一个元素
            es[n] = null;
            if (n > 0) {
                final Comparator<? super E> cmp;
                if ((cmp = comparator) == null)
                    siftDownComparable(0, x, es, n);    // 末元素填到顶, 一路下沉
                else
                    siftDownUsingComparator(0, x, es, n, cmp);
            }
        }
        return result;
    }

    private void grow(int minCapacity) {
        int oldCapacity = queue.length;
        // 小堆翻倍(+2), 大堆加 50%
        int newCapacity = ArraysSupport.newLength(oldCapacity,
                minCapacity - oldCapacity,
                oldCapacity < 64 ? oldCapacity + 2 : oldCapacity >> 1);
        queue = Arrays.copyOf(queue, newCapacity);
    }
}

上浮和下沉是堆的全部算法。注意它俩都不做交换,而是”挪坑”:先把路径上的元素逐个拉下来,最后把目标一次放进坑里,每层省一半写操作:

代码块JAVA · 44 行收起展开
// 基于本地 JDK 源码 (D:/1ForCode/JAVA_Source, java.base, 2024 版), java.util.PriorityQueue
    private static <T> void siftUpComparable(int k, T x, Object[] es) {
        Comparable<? super T> key = (Comparable<? super T>) x;
        while (k > 0) {
            int parent = (k - 1) >>> 1;         // 父下标: (k-1)/2
            Object e = es[parent];
            if (key.compareTo((T) e) >= 0)      // 不小于父亲, 堆序已满足, 停
                break;
            es[k] = e;                          // 父亲被拉下来占坑, 不是 swap
            k = parent;
        }
        es[k] = key;                            // x 最终落座, 全程只写它一次
    }

    private static <T> void siftDownComparable(int k, T x, Object[] es, int n) {
        Comparable<? super T> key = (Comparable<? super T>)x;
        int half = n >>> 1;                     // 第一个叶子的下标; 叶子无需下沉
        while (k < half) {
            int child = (k << 1) + 1;           // 先假设左孩子更小
            Object c = es[child];
            int right = child + 1;
            if (right < n &&
                ((Comparable<? super T>) c).compareTo((T) es[right]) > 0)
                c = es[child = right];          // 右孩子更小就换成右
            if (key.compareTo((T) c) <= 0)      // 不大于较小的孩子, 停
                break;
            es[k] = c;                          // 小孩子上来占坑
            k = child;
        }
        es[k] = key;
    }

    // 从乱序数组原地建堆: 从最后一个非叶节点倒着逐个下沉, Floyd 算法, O(n)
    private void heapify() {
        final Object[] es = queue;
        int n = size, i = (n >>> 1) - 1;
        final Comparator<? super E> cmp;
        if ((cmp = comparator) == null)
            for (; i >= 0; i--)
                siftDownComparable(i, (E) es[i], es, n);
        else
            for (; i >= 0; i--)
                siftDownUsingComparator(i, (E) es[i], es, n, cmp);
    }

原理串讲

offer 的链路:新元素放进数组末尾(层序的下一个空位,保持完全二叉树形态),然后 siftUp 沿父链上浮到不再小于父亲为止,路径长度是树高 $O(\log n)$。
poll 反过来:堆顶就是答案,但挖走堆顶会留洞,于是把最后一个元素搬到顶上补洞(形态恢复完全二叉树),再 siftDown 沉到合适位置。“用末元素补顶”是堆的标准手筋,它同时保住了两个不变量:形态上仍是完全二叉树,顺序上只有根可能违规,一次下沉即可修复。

下沉时 int half = n >>> 1 很典型:下标从 n/2 起全是叶子,叶子没有孩子自然无需下沉,循环条件 k < half 直接把一半节点排除在外。
heapify 建堆也是同一个观察,从最后一个非叶节点 (n/2)-1 倒着往前逐个下沉。
虽然单次下沉最深 $O(\log n)$,但大量节点靠近叶子层、下沉距离很短,总和收敛到 $O(n)$,比逐个 offer 的 $O(n \log n)$ 更快。
所以 new PriorityQueue<>(collection) 比循环 offer 划算。

removeAt(i) 有个精巧的细节:用末元素补第 i 个洞后先 siftDown,如果元素纹丝没动(es[i] == moved)再补一次 siftUp
因为补洞的末元素和被删元素没有大小关系,它可能需要往下走,也可能需要往上走,两个方向都得试。
这个方法还会把”末元素跑到了 i 前面”这一情况返回给迭代器,iterator.remove() 靠它避免漏遍历元素。

比较逻辑全程二选一:构造时给了 comparator 就用它,没给就把元素强转成 ComparablecompareTo。往堆里塞既没实现 Comparable 又没配比较器的对象,第二个元素入堆时抛 ClassCastException,第一个元素反而能混进去(没人跟它比)。

设计取舍

  • 默认小根堆。要大根堆传 Comparator.reverseOrder()。Top K 大的用小根堆限制容量 k(堆顶是第 k 大,比它小的进不来),方向感别搞反。
  • 比较器别写 (a, b) -> a - b,极值相减溢出后大小关系翻转,用 Integer.compare(a, b)
  • 迭代顺序只是数组的层序,中序遍历堆数组得不到有序序列。想按序输出只能连续 poll,一次 $O(\log n)$。
  • remove(Object)indexOf 线性扫 + removeAt,$O(n)$。刷题需要”修改堆内元素优先级”(如 Dijkstra 的 decrease-key)时,标准替代方案:直接把新状态再 offer 一份,poll 出来时校验是否过期,过期就丢弃。
  • 无界队列,只会 grow 不会拒绝;容量策略和 ArrayDeque 一样是小于 64 翻倍、否则加 50%。
  • 线程不安全。并发场景的对应物是 PriorityBlockingQueue(一把锁)和定时线程池里的 DelayedWorkQueue(同款堆 + leader-follower)。

延伸阅读