GENIA

基础数据结构

2022-12-20

基础数据结构(Python)

数组(Array)

一组连续的内存空间,用来存储同类型的数据。数组通常用于存储固定大小的序列数据。

初始化:

法一:使用 [] 运算符

使用 [] 运算符创建一个空列表,然后使用 append() 方法向列表中添加元素。

1
2
3
4
5
6
7
# 创建一个空列表
arr = []
# 向列表中添加元素
arr.append(1)
arr.append(2)
arr.append(3)
print(arr) # [1, 2, 3]

法二:使用 list() 函数

使用 list() 函数创建一个列表,并在括号内指定列表中的元素。

1
2
3
# 使用 list() 函数创建列表
arr = list([1, 2, 3])
print(arr) # [1, 2, 3]

法三:使用列表推导式

可以使用列表推导式快速创建一个列表,语法如下:

1
2
3
4
#[表达式 for 变量 in 序列]
#例如,可以使用列表推导式创建一个从 1 到 10 的数字列表。
arr = [i for i in range(1, 11)]
print(arr) # [1, 2, 3, 4, 5, 6, 7, 8, 9, 10]

法四:使用 numpy 库

使用 import numpy as np 命令导入numpy 库,并使用 np.array() 函数创建数组。

1
2
3
4
5
6
7
8
9
10
11
12
#例如,可以使用 numpy 库创建一个从 1 到 10 的数字数组。
import numpy as np
arr = np.array([1, 2, 3, 4, 5, 6, 7, 8, 9, 10])
print(arr) # [ 1 2 3 4 5 6 7 8 9 10]
#可以使用 np.zeros() 函数创建一个指定大小的全 0 数组
#使用 np.ones() 函数创建一个指定大小的全 1 数组。
# 创建一个 3x3 的全 0 数组
arr = np.zeros((3, 3))
print(arr)
# 创建一个 3x3 的全 1 数组
arr = np.ones((3, 3))
print(arr)

链表(Linked List)

由若干个节点组成的数据结构,每个节点包含数据和指向下一个节点的指针。链表的优点在于可以动态地增加或删除节点,不需要预先确定数据的个数。

链表的优缺点:

链表的优点在于可以快速地在序列的开头或结尾添加或删除元素,不必对整个序列进行重新排序。

链表的缺点在于无法快速地访问序列中的特定位置,如果需要访问序列中间的元素,需要从头开始逐个遍历。因此,在查找和访问序列中的元素时,链表的性能较差。

如果需要快速访问序列中的元素,可以使用数组或者其他数据结构,例如树或哈希表。

使用类来实现:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
class Node:
def __init__(self, data=None, next=None):
self.data = data
self.next = next
class LinkedList:
def __init__(self):
self.head = None
def append(self, data):
new_node = Node(data)
if self.head is None:
self.head = new_node
return
current = self.head
while current.next:
current = current.next
current.next = new_node
def print_list(self):
current = self.head
while current:
print(current.data)
current = current.next

使用这个链表类,可以创建一个新的链表,并向其中添加节点。

1
2
3
4
5
6
7
8
# 创建一个新的链表
llist = LinkedList()
# 向链表中添加节点
llist.append("A")
llist.append("B")
llist.append("C")
# 打印链表中的所有节点
llist.print_list() # A B C

在这个示例中,我们创建了两个类:Node 类表示单个节点,LinkedList 类表示链表。

Node 类包含两个属性:data 和 next。data 属性用于存储节点的数据,next 属性用于存储下一个节点的引用。

LinkedList 类包含一个属性 head,用于存储链表的头节点。它还包含两个方法:append() 方法用于向链表中添加新的节点,print_list() 方法用于打印链表中的所有节点。

为链表添加一些功能:

可以使用以下方法来添加查找、删除、反转链表的方法:

1.查找节点:可以添加一个 find() 方法,该方法接收一个数据值作为参数,并在链表中查找与该数据值相同的节点。如果找到了该节点,返回该节点;如果没有找到,返回 None。

1
2
3
4
5
def find(self, data):
current = self.head
while current and current.data != data:
current = current.next
return current

2.删除节点:可以添加一个 delete() 方法,该方法接收一个数据值作为参数,并在链表中查找与该数据值相同的节点,然后删除该节点。

1
2
3
4
5
6
7
8
9
10
11
12
def delete(self, data):
current = self.head
previous = None
while current and current.data != data:
previous = current
current = current.next
if current is None:
return
if previous is None:
self.head = current.next
else:
previous.next = current.next

3.反转链表:可以添加一个 reverse() 方法,该方法用于将链表中的节点反转。

1
2
3
4
5
6
7
8
9
def reverse(self):
current = self.head
previous = None
while current:
next = current.next
current.next = previous
previous = current
current = next
self.head = previous

栈(Stack)

一种后进先出(Last In First Out, LIFO)的数据结构,常用于实现程序的调用和返回、表达式求值等场景。

栈可以用于存储程序的执行过程中的临时数据,也可以用于实现递归算法。

栈的操作非常简单但是性能较差,如果需要快速访问序列中的元素,可以使用数组或队列等数据结构。

用列表实现:

1
2
3
4
5
6
7
8
9
stack = []
# 入栈
stack.append(1)
stack.append(2)
stack.append(3)
# 出栈
print(stack.pop()) # 3
print(stack.pop()) # 2
print(stack.pop()) # 1

用元组实现:

1
2
3
4
5
6
7
8
9
10
11
stack = ()
# 入栈
stack += (1,)
stack += (2,)
stack += (3,)
# 出栈
print(stack[-1]) # 3
stack = stack[:-1]
print(stack[-1]) # 2
stack = stack[:-1]
print(stack[-1]) # 1

元组的第一个元素是栈底元素,最后一个元素是栈顶元素。使用分片可以删除栈顶元素。

用类实现:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
class Stack:
def __init__(self):
self.items = []
def push(self, item):
self.items.append(item)
def pop(self):
return self.items.pop()
def is_empty(self):
return not self.items
stack = Stack()
# 入栈
stack.push(1)
stack.push(2)
stack.push(3)
# 出栈
print(stack.pop()) # 3
print(stack.pop()) # 2
print(stack.pop()) # 1

定义了一个 Stack 类,包含三个方法:push() 方法用于向栈中添加新的元素,pop() 方法用于弹出栈顶元素,is_empty() 方法用于判断栈是否为空。

队列(Queue)

一种先进先出(First In First Out, FIFO)的数据结构,常用于模拟排队等场景。队列可以用于存储程序的执行过程中的临时数据,也可以用于实现广度优先搜索算法。

队列的性能较差,如果需要快速在序列的开头或结尾添加或删除元素,可以使用栈或链表等数据结构。

用列表实现:

1
2
3
4
5
6
7
8
9
queue = []
# 入队
queue.append(1)
queue.append(2)
queue.append(3)
# 出队
print(queue.pop(0)) # 1
print(queue.pop(0)) # 2
print(queue.pop(0)) # 3

用元组实现:

1
2
3
4
5
6
7
8
9
10
11
queue = ()
# 入队
queue += (1,)
queue += (2,)
queue += (3,)
# 出队
print(queue[0]) # 1
queue = queue[1:]
print(queue[0]) # 2
queue = queue[1:]
print(queue[0]) # 3

用类实现:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
class Queue:
def __init__(self):
self.items = []
def enqueue(self, item):
self.items.append(item)
def dequeue(self):
return self.items.pop(0)
def is_empty(self):
return not self.items
# 使用类实现队列
queue = Queue()
# 入队
queue.enqueue(1)
queue.enqueue(2)
queue.enqueue(3)
# 出队
print(queue.dequeue()) # 1
print(queue.dequeue()) # 2
print(queue.dequeue()) # 3

定义了一个 Queue 类,包含三个方法:enqueue() 方法用于向队列中添加新的元素,dequeue() 方法用于弹出队列的第一个元素,is_empty() 方法用于判断队列是否为空。

通过python内置模块实现:

Python 提供内置的队列模块 queue,可以用来实现多线程编程中的同步机制。这个模块提供了 Queue 类、LifoQueue 类和 PriorityQueue 类等,可以用于实现不同的队列类型。

例如,可以使用 Queue 类来实现多线程的生产者-消费者模型:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
import threading
import queue
# 创建队列
q = queue.Queue()
# 创建生产者线程
def producer():
for i in range(10):
q.put(i)
print(f'生产了一个数字:{i}')
# 创建消费者线程
def consumer():
while True:
item = q.get()
print(f'消费了一个数字:{item}')
q.task_done()
# 创建并启动线程
t1 = threading.Thread(target=producer)
t2 = threading.Thread(target=consumer)
t1.start()
t2.start()
# 等待队列为空
q.join()
print('完成!')

创建了两个线程:一个生产者线程和一个消费者线程。生产者线程向队列中添加数字,消费者线程从队列中取出数字并打印。

树(Tree)

由若干个节点组成的数据结构,每个节点有零个或多个子节点。树通常用于表示层级关系、存储有序数据等。常见的树包括二叉树、平衡树、Trie树等。

树的不同遍历方式:

主要包括深度优先遍历和广度优先遍历。

深度优先遍历是指从根节点开始,先递归遍历左子树,再递归遍历右子树。深度优先遍历的常用方法有先序遍历、中序遍历、后序遍历。

1.先序遍历(Pre-order Traversal): 先访问根节点,再递归地访问左子树,最后递归地访问右子树。

2.中序遍历(In-order Traversal): 先递归地访问左子树,再访问根节点,最后递归地访问右子树。

3.后序遍历(Post-order Traversal): 先递归地访问左子树,再递归地访问右子树,最后访问根节点。

广度优先遍历是指从根节点开始,按照层级顺序依次遍历每个节点。常用的广度优先遍历方法是层序遍历。

4.层序遍历(Level-order Traversal): 按照树的层级顺序,从上到下、从左到右地访问每个节点。

例如:

1
2
3
4
5
6
7
8
9
10
11
   1
/ \
2 3
/ / \
4 5 6
'''
先序遍历:1, 2, 4, 3, 5, 6
中序遍历:4, 2, 1, 5, 3, 6
后序遍历:4, 2, 5, 6, 3, 1
层序遍历:1, 2, 3, 4, 5, 6
'''

不同的遍历方式在不同的应用场景中有着不同的用途。例如,在二叉搜索树中,中序遍历可以用来得到有序的节点序列。在二叉树的某些应用中,后序遍历可以用来释放子树的资源,以便在释放根节点之前可以访问它的子节点。

使用类来实现:

1
2
3
4
5
6
7
8
9
10
11
12
13
class TreeNode:
def __init__(self, val=None):
self.val = val
self.left = None
self.right = None
# 使用类实现树
root = TreeNode(1)
root.left = TreeNode(2)
root.right = TreeNode(3)
root.left.left = TreeNode(4)
root.left.right = TreeNode(5)
root.right.left = TreeNode(6)
root.right.right = TreeNode(7)

定义了一个 TreeNode 类,表示树的节点。每个节点都有一个值(val)和两个子节点(left 和 right)。

使用类实现树的好处在于可以在类中添加方法,用于实现树的遍历、查找、插入、删除等操作。

1.添加 preorder_traversal() 方法来实现树的先序遍历:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
    def preorder_traversal(self):
res = []
res.append(self.val)
if self.left:
res += self.left.preorder_traversal()
if self.right:
res += self.right.preorder_traversal()
return res
# 使用类实现树的先序遍历
root = TreeNode(1)
root.left = TreeNode(2)
root.right = TreeNode(3)
root.left.left = TreeNode(4)
root.left.right = TreeNode(5)
root.right.left = TreeNode(6)
root.right.right = TreeNode(7)
print(root.preorder_traversal()) # [1, 2, 4, 5, 3, 6, 7]

这个方法递归遍历每个节点的左子树和右子树,并将每个节点的值存储在结果列表中。

2.添加 inorder_traversal() 方法来实现树的中序遍历:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
    def inorder_traversal(self):
res = []
if self.left:
res += self.left.inorder_traversal()
res.append(self.val)
if self.right:
res += self.right.inorder_traversal()
return res
# 使用类实现树的中序遍历
root = TreeNode(1)
root.left = TreeNode(2)
root.right = TreeNode(3)
root.left.left = TreeNode(4)
root.left.right = TreeNode(5)
root.right.left = TreeNode(6)
root.right.right = TreeNode(7)
print(root.inorder_traversal()) # [4, 2, 5, 1, 6, 3, 7]

递归遍历每个节点的左子树,然后处理当前节点的值,最后遍历每个节点的右子树。

3.添加postorder_traversal() 方法来实现树的后序遍历

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
    def postorder_traversal(self):
res = []
if self.left:
res += self.left.postorder_traversal()
if self.right:
res += self.right.postorder_traversal()
res.append(self.val)
return res
# 使用类实现树的后序遍历
root = TreeNode(1)
root.left = TreeNode(2)
root.right = TreeNode(3)
root.left.left = TreeNode(4)
root.left.right = TreeNode(5)
root.right.left = TreeNode(6)
root.right.right = TreeNode(7)
print(root.postorder_traversal()) # [4, 5, 2, 6, 7, 3, 1]

递归遍历每个节点的左子树和右子树,最后处理当前节点的值。

4.添加level_order_traversal() 方法来实现树的层序遍历

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
    def level_order_traversal(self):
res = []
q = [self]
while q:
node = q.pop(0)
res.append(node.val)
if node.left:
q.append(node.left)
if node.right:
q.append(node.right)
return res
# 使用类实现树的层序遍历
root = TreeNode(1)
root.left = TreeNode(2)
root.right = TreeNode(3)
root.left.left = TreeNode(4)
root.left.right = TreeNode(5)
root.right.left = TreeNode(6)
root.right.right = TreeNode(7)
print(root.level_order_traversal()) # [1, 2, 3, 4, 5, 6, 7]

使用一个队列来存储每一层的节点,每次取出队列中的第一个节点并将其子节点加入队列。

哈希表(Hash Table)

是一种用于快速查找的数据结构。它通过计算数据的哈希值,将数据存储在表中的某个位置。哈希表常用于实现字典、集合等数据类型。

哈希表的基本原理是:使用哈希函数将数据映射到表中的某个位置,并使用链表来解决冲突。哈希函数计算出的值称为哈希值,哈希表中的每个位置称为桶。

哈希表的优点在于查找的时间复杂度为O(1),但是如果哈希函数不良或者数据分布不均匀,哈希表的性能可能会降低。

哈希值及计算

哈希值是指将数据通过哈希函数转换成固定长度的哈希码的过程,这个过程称为哈希。哈希值可以用来比较数据之间的相似度或者标识数据,常用于数据存储、检索和加密等方面。

Python 语言提供了内置函数 hash() 来计算哈希值。可以使用如下代码计算任意数据的哈希值:

1
hash_value = hash(data)

由于哈希函数是把任意长度的数据转换成固定长度的哈希码,因此不同的数据可能会得到相同的哈希值,这种现象称为哈希冲突。为了尽量减少哈希冲突的发生,Python 中的 hash() 函数针对不同的数据类型使用了不同的哈希算法。

除了使用 Python 中的内置函数 hash() 来计算哈希值之外,你还可以使用 Python 中的哈希函数库,例如 hashlib 模块。hashlib 模块提供了各种常见的哈希算法,例如 MD5、SHA-1 等。

1
2
3
4
5
6
7
8
9
10
11
12
13
import hashlib
#计算字符串的 MD5 哈希值
md5 = hashlib.md5()
md5.update('hello'.encode())
hash_value = md5.hexdigest()
#计算字符串的 SHA-1 哈希值
sha1 = hashlib.sha1()
sha1.update('hello'.encode())
hash_value = sha1.hexdigest()
#计算字符串的 SHA-256 哈希值
sha256 = hashlib.sha256()
sha256.update('hello'.encode())
hash_value = sha256.hexdigest()

注意:使用 hashlib 模块计算哈希值时,必须将字符串编码成二进制数据,然后使用 update() 方法更新哈希值,最后使用 hexdigest() 方法获取哈希值。

字典:现成的哈希表

1
2
3
4
5
6
#创建
scores = {'Alice': 95, 'Bob': 75, 'Charlie': 85}
#访问
score = scores['Alice'] # score 等于 95
#更新
scores['Alice'] = 99 # 将 Alice 的成绩更新为 99

字典的优点之一是允许使用任意的不可变类型作为键。因此可以使用字符串、数字、元组等类型作为键。

需要注意,字典并不是按照插入顺序来存储数据的。如果需要按照插入顺序来存储数据,可以使用 Python 的内置数据类型 OrderedDict。

用 set 实现

set 类型是一种无序且不重复的数据集合,可以使用 set 来快速查找数据是否存在,也可以使用 set 来去重。

1
2
3
4
5
6
7
8
9
10
11
12
13
#创建一个set
s = set([1, 2, 3])
#添加元素
s.add(4)
#删除元素
s.remove(2)
#求并集
s2 = set([3, 4, 5])
s3 = s1.union(s2) # s3 等于 set([1, 2, 3, 4, 5])
#求交集
s3 = s1.intersection(s2) # s3 等于 set([3])
#求差集
s3 = s1.difference(s2) # s3 等于 set([1, 2])

注意:set 类型只能存储可哈希的数据类型,例如数字、字符串、元组等。不能使用 set 存储列表、字典或其他可变的数据类型,因为这些数据类型的哈希值是可变的。

1
2
3
4
5
6
7
8
#创建
scores = {('Alice', 95), ('Bob', 75), ('Charlie', 85)}
#查找
if ('Alice', 95) in scores:
print('Alice has score 95')
#更新
scores.remove(('Alice', 95))
scores.add(('Alice', 99))

相比于字典,set类型无法查找指定的键对应的值、无法按照顺序遍历哈希表中的数据等。

使用 Python 的内置模块 collections 实现

Python 中的 collections 模块提供了一系列的容器数据类型,其中包括哈希表类型。相比于字典,可以赋予哈希表一些特殊的功能。

1.使用 collections 模块中的字典子类型 defaultdict 实现

虽与字典类型类似,但是 defaultdict 可以设置默认值。如果在访问不存在的键时,defaultdict 会自动创建一个新的键值对,并将默认值赋给新的键。

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
from collections import defaultdict
#创建
d = defaultdict(int) # 创建一个默认值为 0 的 defaultdict
#访问
d['a'] = 1
print(d['a']) # 输出 1
print(d['b']) # 输出 0,b 键不存在,defaultdict 会自动创建一个新的键值对,并将默认值赋给 b 键
#删除
del d['a']
#遍历
for key, value in d.items():
print(key, value)
#查找
if 'a' in d:
print('a exists in d')
#修改
d['a'] = 2 # 修改 a 键对应的值为 2
#获取所有键值对
print(d.items()) # 输出 [('a', 2)]
#获取所有键
print(d.keys()) # 输出 ['a']
#获取所有值
print(d.values()) # 输出 [2]
2.使用 collections 模块中的字典子类型Counter实现

相比于字典,Counter 可以用来统计数据出现的次数。它还支持多种常用的统计方法,例如求和、求平均值、求标准差等。

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
from collections import Counter
#创建
c = Counter([1, 2, 2, 3, 3, 3])
#访问
print(c[1]) # 输出 1
print(c[2]) # 输出 2
print(c[3]) # 输出 3
#统计数据出现的次数
print(c.most_common())
# 输出 [(3, 3), (2, 2), (1, 1)],表示数字 3 出现了 3 次,数字 2 出现了 2 次,数字 1 出现了 1 次
#求和
print(sum(c.values())) # 输出 6,表示 1+2+2+3+3+3=6
#求平均值
print(sum(c.values()) / len(c)) # 输出 2.0,表示 (1+2+2+3+3+3)/6=2
#求标准差
import statistics
print(statistics.stdev(c.values())) # 输出 1.0,表示标准差为 1.0

手动实现

1
2
3
4
5
6
7
8
9
10
#存储
scores = [['Alice', 95], ['Bob', 75], ['Charlie', 85]]
#访问
for student, score in scores:
if student == 'Alice':
print(score) # 输出 95
#更新
for i, (student, score) in enumerate(scores):
if student == 'Alice':
scores[i] = ['Alice', 99] # 将 Alice 的成绩更新为 99

手动实现哈希表需要注意一些细节,例如哈希函数的设计、冲突的解决方法等。

图(Graph)

图由节点(称为顶点)和边组成。边表示两个节点之间的关系,并可以具有不同的权值或负载。

图可以有向或无向。有向图中,边有方向,即它从一个节点指向另一个节点。无向图中,边没有方向,即它们连接两个节点,但不指向任何一个节点。

图也可以是稠密的或稀疏的。稠密图中,大多数节点都是相互连接的,而稀疏图中,节点之间的连接相对较少。

图在许多领域中都很有用,例如社交网络分析、交通网络建模、计算机网络等。它们可以用于查找最短路径、求解最小生成树、模拟货运路线等问题。

图可以使用多种数据结构来存储,如邻接矩阵、邻接表、边集数组和边列表。其中邻接矩阵和邻接表是较常用的两种结构。

邻接矩阵是一个二维数组,其中第i行第j列的元素表示节点i和节点j之间是否存在边。如果存在,则值为1,否则为0。邻接矩阵适用于稠密图,其中大多数节点都有边相连。它可以很快地查找两个节点之间是否存在边,但需要较多的存储空间。

邻接表是一个数组,其中的每个元素都是一个链表,表示与该节点相连的所有节点。邻接表适用于稀疏图,其中节点之间的连接相对较少。它不需要像邻接矩阵那样多的存储空间,但查找特定边所需的时间可能会稍长。

如果需要对图进行更复杂的操作,可以使用 Python 的第三方库,例如 NetworkX。

使用字典(邻接表)存储

1
2
3
4
5
6
7
8
9
graph = {
'A': ['B', 'C'],
'B': ['C', 'D'],
'C': ['D'],
'D': ['C'],
'E': ['F'],
'F': ['C']
}
#图中有 6 个节点(A、B、C、D、E 和 F)和 7 条边

字典中的每个键都是一个节点,对应的值是连接到该节点的其他节点的列表。

使用列表(邻接矩阵)存储

1
2
3
4
5
6
7
8
9
graph = [
[0, 1, 1, 0, 0, 0],
[0, 0, 1, 1, 0, 0],
[0, 0, 0, 1, 0, 0],
[0, 0, 1, 0, 0, 0],
[0, 0, 0, 0, 0, 1],
[0, 0, 0, 0, 0, 0]
]
图中有 6 个节点(0、1、2、3、4 和 5)和 7 条边

堆(Heap)

堆是一种特殊的树形数据结构,它满足以下性质:

  • 堆是一颗完全二叉树,也就是说,除了最后一层,其他层的节点都是满的;最后一层的节点都靠左对齐。
  • 堆分为两种:最大堆和最小堆。最大堆的性质是,任意一个节点的值都大于等于其子节点的值;最小堆的性质是,任意一个节点的值都小于等于其子节点的值。

堆通常用数组来实现。在堆中,第 $i$ 个元素的左子节点的索引是 $2i+1$,右子节点的索引是 $2i+2$,父节点的索引是 $\lfloor \frac{i-1}{2} \rfloor$。

堆常用于优先队列,也常用于排序。例如,可以使用堆排序算法来对数组进行排序。堆排序的时间复杂度是 $O(nlogn)$。

使用堆来实现优先队列

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
class PriorityQueue:
def __init__(self):
self.heap = []
self.index = 0

def push(self, item, priority):
# 将新元素插入堆中
self.heap.append((priority, self.index, item))
self.index += 1
# 将新元素上浮到合适的位置
self._float_up(len(self.heap) - 1)

def pop(self):
# 将堆顶元素弹出
top_item = self.heap[0]
# 将堆最后一个元素移到堆顶,然后将其下沉到合适的位置
self.heap[0] = self.heap[-1]
self.heap.pop()
self._float_down(0)
return top_item

def _float_up(self, index):
# 如果父节点的优先级比当前节点低,则将父节点下沉,并将当前节点上浮
parent = (index - 1) // 2
if index > 0 and self.heap[index][0] > self.heap[parent][0]:
self.heap[index], self.heap[parent] = self.heap[parent], self.heap[index]
self._float_up(parent)

def _float_down(self, index):
# 如果子节点的优先级比当前节点高,则将子节点上浮,并将当前节点下沉
left = index * 2 + 1
right = index * 2 + 2
largest = index
if left < len(self.heap) and self.heap[left][0] > self.heap[largest][0]:
largest = left
if right < len(self.heap) and self.heap[right][0] > self.heap[largest][0]:
largest = right
if largest != index:
self.heap[index], self.heap[largest] = self.heap[largest], self.heap[index]
self._float_down(largest)

这个优先队列类有以下几个方法:

  • __init__ 方法:初始化堆和索引。
  • push 方法:将新元素插入堆中,并将其上浮到合适的位置。
  • pop 方法:将堆顶元素弹出,并将堆最后一个元素移到堆顶,然后将其下沉到合适的位置。
  • _float_up 方法:将节点上浮到合适的位置。
  • _float_down 方法:将节点下沉到合适的位置。

这个优先队列类使用了堆的性质,来保证优先级最高的元素总是位于堆顶。

使用方法类似于普通的队列,例如:

1
2
3
4
5
6
7
pq = PriorityQueue()
pq.push('A', 10)
pq.push('B', 5)
pq.push('C', 20)
print(pq.pop()) # 输出 ('C', 20)
print(pq.pop()) # 输出 ('A', 10)
print(pq.pop()) # 输出 ('B', 5)