前言

在 C++ STL 中,priority_queue 是一个非常常用的容器适配器。它可以让我们快速获取当前优先级最高的元素。默认情况下,STL 的 priority_queue 是一个 大堆,也就是每次 top() 取到的都是最大值。本文我们不直接使用 STL 的 priority_queue,而是自己手写一个简化版本,理解它背后的核心原理:堆。

一、priority_queue 的本质

priority_queue 的底层通常使用 来实现。
堆是一棵逻辑上的完全二叉树,但在代码中一般用数组或 vector 存储。
对于下标为 i 的节点:

parent = (i - 1) / 2;
leftChild = i * 2 + 1;
rightChild = i * 2 + 2;

比如数组:

[66, 30, 2, 4, 3]

可以看成这样一棵树:

        66
      /    \
    30      2
   /  \
  4    3

如果每个父亲节点都比孩子节点大,这就是 大堆。

如果每个父亲节点都比孩子节点小,这就是 小堆。

二、比较器设计

为了让我们的优先级队列既支持大堆,也支持小堆,可以使用仿函数作为比较器。

template<class T>
struct Less
{
    bool operator()(const T& x, const T& y) const {
        return x < y;
    }
};

template<class T>
struct Greater
{
    bool operator()(const T& x, const T& y) const {
        return x > y;
    }
};

这里的设计和 STL 类似。

  • 如果使用 Less,表示当父亲节点小于孩子节点时,需要交换,于是堆顶会变成最大值,也就是 大堆。

  • 如果使用 Greater,表示当父亲节点大于孩子节点时,需要交换,于是堆顶会变成最小值,也就是 小堆。

三、priority_queue 的基本框架

我们可以把容器类型和比较器都设计成模板参数:

template <class T, class Container = vector<T>, class Compare = Less<T>>
class priority_queue
{
public:
    priority_queue() = default;

private:
    Container _con;
};

这里:

Container = vector<T>

表示默认底层容器是 vector。

Compare = Less<T>

表示默认是大堆。

bit::priority_queue<int>//完整类名大致如下:

如果想要小堆,可以这样写:

bit::priority_queue<int, vector<int>, bit::Greater<int>> pq;

四、向上调整 adjust_up

当我们插入一个新元素时,新元素会先被尾插到 vector 的末尾。
但是尾插之后,它可能破坏堆的结构,所以需要向上调整。

void adjust_up(int child)
{
    Compare com;
    int parent = (child - 1) / 2;

    while (child > 0)
    {
        if (com(_con[parent], _con[child]))
        {
            swap(_con[child], _con[parent]);
            child = parent;
            parent = (child - 1) / 2;
        }
        else
        {
            break;
        }
    }
}

以大堆为例:

if (_con[parent] < _con[child])

说明孩子比父亲大,不符合大堆规则,所以交换。

插入接口就很简单:

void push(const T& x)
{
    _con.push_back(x);
    adjust_up(_con.size() - 1);
}

所以 push 的核心流程是:

尾插元素 -> 向上调整

时间复杂度是:O(logN)

五、向下调整 adjust_down

删除堆顶元素时,不能直接删除 _con[0],因为这样会破坏数组结构。

通常做法是:

  1. 交换堆顶元素和最后一个元素;
  2. 交换堆顶元素和最后一个元素;
  3. 从堆顶开始向下调整。
void pop()
{
    swap(_con[0], _con[_con.size() - 1]);
    _con.pop_back();
    adjust_down(0);
}

向下调整代码如下:

void adjust_down(int parent)
{
    Compare com;
    int child = parent * 2 + 1;

    while (child < _con.size())
    {
        if (child + 1 < _con.size() && com(_con[child], _con[child + 1]))
        {
            ++child;
        }

        if (com(_con[parent], _con[child]))
        {
            swap(_con[child], _con[parent]);
            parent = child;
            child = parent * 2 + 1;
        }
        else
        {
            break;
        }
    }
}

这里有两个关键点。

第一个关键点是选择左右孩子中优先级更高的那个:

if (child + 1 < _con.size() && com(_con[child], _con[child + 1]))
{
    ++child;
}

对于大堆来说,这一步是在左右孩子中选更大的。

对于小堆来说,这一步是在左右孩子中选更小的。

第二个关键点是判断父亲和孩子是否需要交换:

if (com(_con[parent], _con[child]))
{
    swap(_con[child], _con[parent]);
}

对于大堆来说,如果父亲小于孩子,就交换。

对于小堆来说,如果父亲大于孩子,就交换。

pop 的时间复杂度也是:O(logN)

六、用区间构造堆

除了一个一个 push,我们还可以用一段区间直接构造优先级队列。

template <class InputIterator>
priority_queue(InputIterator first, InputIterator last)
    : _con(first, last)
{
    for (int i = (_con.size() - 2) / 2; i >= 0; i--)
    {
        adjust_down(i);
    }
}

这段代码的核心思想是:从最后一个非叶子节点开始,依次向下调整
最后一个非叶子节点的下标是:

(_con.size() - 2) / 2

这种建堆方式的时间复杂度是:O(N)

比一个一个 push 的 O(NlogN) 更高效。

七、完整代码

#pragma once
#include <vector>
#include <functional>
#include <iostream>
using namespace std;

namespace mao
{
    template<class T>
    struct Less
    {
        bool operator()(const T& x, const T& y) const {
            return x < y;
        }
    };

    template<class T>
    struct Greater
    {
        bool operator()(const T& x, const T& y) const {
            return x > y;
        }
    };

    template <class T, class Container = vector<T>, class Compare = Less<T>>
    class priority_queue
    {
    public:
        template <class InputIterator>
        priority_queue(InputIterator first, InputIterator last)
            : _con(first, last)
        {
            for (int i = (_con.size() - 2) / 2; i >= 0; i--)
            {
                adjust_down(i);
            }
        }

        priority_queue() = default;

        void adjust_up(int child)
        {
            Compare com;
            int parent = (child - 1) / 2;

            while (child > 0)
            {
                if (com(_con[parent], _con[child]))
                {
                    swap(_con[child], _con[parent]);
                    child = parent;
                    parent = (child - 1) / 2;
                }
                else
                {
                    break;
                }
            }
        }

        void adjust_down(int parent)
        {
            Compare com;
            int child = parent * 2 + 1;

            while (child < _con.size())
            {
                if (child + 1 < _con.size() && com(_con[child], _con[child + 1]))
                {
                    ++child;
                }

                if (com(_con[parent], _con[child]))
                {
                    swap(_con[child], _con[parent]);
                    parent = child;
                    child = parent * 2 + 1;
                }
                else
                {
                    break;
                }
            }
        }

        void push(const T& x)
        {
            _con.push_back(x);
            adjust_up(_con.size() - 1);
        }

        void pop()
        {
            swap(_con[0], _con[_con.size() - 1]);
            _con.pop_back();
            adjust_down(0);
        }

        const T& top() const
        {
            return _con[0];
        }

        bool empty() const
        {
            return _con.empty();
        }

        size_t size() const
        {
            return _con.size();
        }

    private:
        Container _con;
    };
}

八、测试代码

#include "priority_queue.h"

int main()
{
    int a[] = { 30, 4, 2, 66, 3 };

    mao::priority_queue<int, vector<int>, mao::Greater<int>> pq(a, a + 5);

    pq.push(3);
    pq.push(1);
    pq.push(5);
    pq.push(7);
    pq.push(2);

    while (!pq.empty())
    {
        cout << pq.top() << " ";
        pq.pop();
    }

    cout << endl;

    return 0;
}

九、几个需要注意的问题

1.pop 前最好判断是否为空

当前代码中:

void pop()
{
    swap(_con[0], _con[_con.size() - 1]);
    _con.pop_back();
    adjust_down(0);
}

如果队列为空时调用 pop(),会出现越界访问。更严谨的写法可以加断言:

void pop()
{
    assert(!_con.empty());

    swap(_con[0], _con[_con.size() - 1]);
    _con.pop_back();

    if (!_con.empty())
    {
        adjust_down(0);
    }
}

2.top 前也最好判断是否为空

当前代码:

const T& top() const
{
    return _con[0];
}

如果空队列调用 top(),也会越界。
可以改成:

const T& top() const
{
    assert(!_con.empty());
    return _con[0];
}

3. 区间构造函数在空区间时可能有隐患

for (int i = (_con.size() - 2) / 2; i >= 0; i--)
{
    adjust_down(i);
}

如果 _con.size() 是 0,由于 size() 返回的是无符号类型 size_t,表达式 _con.size() - 2 可能发生无符号整数下溢。

for (int i = (int)(_con.size() - 2) / 2; i >= 0; i--)
{
    adjust_down(i);
}

更多推荐