C++ 手写 priority_queue
文章目录
前言
在 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],因为这样会破坏数组结构。
通常做法是:
- 交换堆顶元素和最后一个元素;
- 交换堆顶元素和最后一个元素;
- 从堆顶开始向下调整。
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);
}
更多推荐
所有评论(0)