You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

为何基于堆的Dijkstra实现性能低于朴素Dijkstra实现?

为什么我的堆优化版Dijkstra算法比朴素版更慢?

我实现了朴素版Dijkstra算法与基于堆的优化版Dijkstra算法,但意外发现朴素版运行速度更快,调试代码后仍未找到问题所在,特此询问原因。

数据导入与处理代码

import time
with open("DijkstraTest2.txt", 'r') as input:
    lines = input.readlines()

lengths = {}
vertices = []
for line in lines:
    contents = line.split("\t")
    vertices.append(contents[0])
    for content in contents:
        content = content.replace('\n', '')
        if ',' in content:
            edge = contents[0] + '-' + content.split(',')[0]
            lengths[edge] = int(content.split(',')[1])

朴素版Dijkstra实现

def NaiveDijkstra(vertices, start_point, lengths):
    X = [start_point]
    shortest_paths = {}
    for vertex in vertices:
        if vertex == start_point:
            shortest_paths[vertex] = 0
        else:
            shortest_paths[vertex] = 999999999999
    subset = [key for key in lengths.keys() if start_point == key.split('-')[0]
              and key.split('-')[0] in X and key.split('-')[1] not in X]
    while len(subset) > 0:
        temp_min_dict = {}
        for edge in subset:
            temp_min = shortest_paths[edge.split('-')[0]] + lengths[edge]
            temp_min_dict[edge] = temp_min
        new_edge = min(temp_min_dict, key=temp_min_dict.get)
        X.append(new_edge.split('-')[1])
        shortest_paths[new_edge.split('-')[1]] = shortest_paths[new_edge.split('-')[0]] + lengths[new_edge]
        subset = []
        for key in lengths.keys():
            if key.split('-')[0] in X and key.split('-')[1] not in X:
                subset.append(key)
    return shortest_paths

start_time = time.time()
print(NaiveDijkstra(vertices = vertices, start_point = '1', lengths = lengths)['197'])
print(time.time() - start_time, "seconds")

堆优化版Dijkstra实现

class Heap:
    def __init__(self):
        self.size = 0
        self.lst = []

    def swap(self, a):
        if self.size == 1:
            return self.lst
        else:
            if a == 1:
                i = 1
            else:
                i = a // 2
            while i > 0:
                if i * 2 - 1 >= self.size:
                    break
                elif self.lst[i - 1][1] > self.lst[i * 2 - 1][1]:
                    temp = self.lst[i - 1]
                    self.lst[i - 1] = self.lst[i * 2 - 1]
                    self.lst[i * 2 - 1] = temp
                elif i * 2 >= self.size:
                    break
                elif self.lst[i - 1][1] > self.lst[i * 2][1]:
                    temp = self.lst[i - 1]
                    self.lst[i - 1] = self.lst[i * 2]
                    self.lst[i * 2] = temp
                i -= 1
        # print(f"output: {self.lst}")

    def insert(self, element):
        # print(f"input: {self.lst}")
        self.lst.append(element)
        self.size += 1
        self.swap(self.size)

    def extractmin(self):
        val = self.lst.pop(0)[0]
        self.size -= 1
        self.swap(self.size - 1)
        return val

    def delete(self, deleted):
        ix = self.lst.index(deleted)
        temp = self.lst[-1]
        self.lst[ix] = temp
        self.lst[-1] = deleted
        self.lst.pop(-1)
        self.size -= 1
        #self.swap(self.size)


def FastDijkstra(vertices, start_point, lengths):
    X = []
    h = Heap()
    width = {}
    shortest_paths = {}
    for vertex in vertices:
        if vertex == start_point:
            width[vertex] = 0
            h.insert((vertex, width[vertex]))
        else:
            width[vertex] = 999999999999
            h.insert((vertex, width[vertex]))
    while h.size > 0:
        w = h.extractmin()
        X.append(w)
        shortest_paths[w] = width[w]
        Y = set(vertices).difference(X)
        for x in X:
            for y in Y:
                key = f"{x}-{y}"
                if lengths.get(key) is not None:
                    h.delete((y, width[y]))
                    if width[y] > shortest_paths[x] + lengths[key]:
                        width[y] = shortest_paths[x] + lengths[key]
                    h.insert((y, width[y]))
    return shortest_paths


start_time = time.time()
print(FastDijkstra(vertices=vertices, start_point='1', lengths=lengths)['197'])
print(time.time() - start_time, "seconds")

问题分析

你的堆优化版比朴素版慢,核心原因是实现完全偏离了堆优化Dijkstra的正确逻辑,导致时间复杂度反而更高:

  1. 堆实现的致命低效

    • extractmin用lst.pop(0):Python列表头部删除是O(n)复杂度,完全丧失堆O(logn)取最小值的优势。正确做法是交换堆顶和最后一个元素,pop末尾元素后执行下沉操作。
    • delete用lst.index(deleted):遍历查找元素是O(n)复杂度,堆优化需要维护节点到堆索引的映射字典,快速定位元素位置。
    • swap逻辑混乱:没有正确实现堆的上浮(插入后)和下沉(取最小值后)操作,堆结构维护错误,做了大量无用功。
  2. 算法逻辑错误,时间复杂度爆炸
    FastDijkstra里的for x in X: for y in Y是O(V²)复杂度,和朴素版持平甚至更糟。堆优化版的正确逻辑是:取出堆顶节点后,只遍历该节点的所有邻接节点,而非遍历所有已处理/未处理节点的组合,这直接导致大量不必要的计算。

  3. 冗余操作过多
    每次更新距离时先delete再insert,这两个操作本身低效。堆优化版通常无需删除旧元素,允许堆中存在同一节点的多个距离记录,后续取出时若发现记录距离大于当前已知最短路径,直接跳过即可。

修复建议

  • 重构Heap类:实现正确的最小堆,包含O(logn)的extractmin、insert,用索引映射实现高效更新。
  • 修改FastDijkstra逻辑:取出节点后仅遍历其所有出边,更新邻接节点距离,若距离更新则插入新的(节点,距离)到堆中(无需删除旧记录)。
  • 优化数据结构:将lengths改为邻接表(如adj = {u: [(v, weight), ...]}),快速获取节点邻接边,避免字符串拼接查字典的开销。

内容的提问来源于stack exchange,提问作者Jonsi Billups

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.19 16:40:27