普通的prim算法:
class Solution:
def minCostConnectPoints(self, points: List[List[int]]) -> int:
#构建邻接矩阵
#prim算法:逐步把点加入到树中,每次都加入离树最近的点
n = len(points)
matrix = [[0] * n for i in range(n)]
for i in range(n):
for j in range(1, n):
dist = abs(points[i][0] - points[j][0]) + abs(points[i][1] - points[j][1])
matrix[i][j] = dist
matrix[j][i] = dist
#print(matrix)
mindist = [float('inf')] * n #未加入树中的点到已经形成的树的最短距离
selected = [False] * n #是否被加入树中了
#先把0加入树,作为根节点
mindist[0] = 0
for i in range(1, n):
mindist[i] = matrix[i][0]
selected[0] = True
ans = 0
#遍历剩余n-1个节点
while not all(selected):
#找出距离树最近的那个点,记录距离和坐标
closest = float('inf')
closest_idx = -1
#for i in range(n): #i遍历的是树中的点 ???为何要遍历树中的点???,mindist存的已经是树外的点到树的最小距离了
# if not selected[i]: continue
# for j in range(n):#j遍历的是树外的点
# if selected[j]: continue
# if mindist[j] < closest:
# closest = mindist[j]
# closest_idx = j
for j in range(n):#j遍历的是树外的点
if selected[j]: continue
if mindist[j] < closest:
closest = mindist[j]
closest_idx = j
#将索引为j的点加入树中
ans += closest
mindist[closest_idx] = 0
selected[closest_idx] = True
for k in range(n):
if not selected[k] and matrix[k][closest_idx] < mindist[k]:
mindist[k] = matrix[k][closest_idx]
return ans
带堆优化的prim算法:主要是用堆优化 寻找树外的点到树的最小值 这个步骤,直接用堆找最小值更省时
class Solution:
def minCostConnectPoints(self, points: List[List[int]]) -> int:
#堆优化的prim算法
import heapq
n = len(points)
matrix = [[0] * n for i in range(n)]
for i in range(n):
for j in range(1, n):
dist = abs(points[i][0] - points[j][0]) + abs(points[i][1] - points[j][1])
matrix[i][j] = dist
matrix[j][i] = dist
selected = [False] * n
mindist = [[float('inf'), i] for i in range(n)] #树外节点距离树的距离,节点号
#一开始想用tuple,但是tuple是不可变对象
ans = 0
#将第0个节点放入
selected[0] = True
mindist[0][0] = 0
for j in range(1, n):
mindist[j][0] = matrix[0][j]
mindist.pop(0)
heapq.heapify(mindist)
while not all(selected):
closest, closest_idx = heapq.heappop(mindist)
ans += closest
selected[closest_idx] = True
#更新距离
for m in mindist:
if matrix[closest_idx][m[1]] < m[0]:
m[0] = matrix[closest_idx][m[1]]
heapq.heapify(mindist)
return ans
Kruskal算法,n个点之间有n(n-1)/2条边,留下n-1条。依次从最小的边开始添加,并查集用来检查添加当前边有没有导致形成环。导致成环的边就跳过。
class Solution:
def minCostConnectPoints(self, points: List[List[int]]) -> int:
#Kruskal算法,n个点之间有n(n-1)/2条边,留下n-1条
#依次从最小的边开始添加,并查集用来检查添加当前边有没有导致形成环
#初始化每条边,(起点,终点,长度)
n = len(points)
edges = []
parents = {i:i for i in range(n)}
for i in range(n):
for j in range(i + 1, n):
edges.append((i, j, abs(points[i][0] - points[j][0]) + abs(points[i][1] - points[j][1])))
edges.sort(key = lambda x : x[2])
#print(edges)
def find(x):
root = x
while x != parents[x]:
x = parents[x]
while parents[root] != x:
tmp = parents[root]
parents[root] = x
root = tmp
return x
added_count = 0
ans = 0
for x, y, lenth in edges: #有选择地merge,符合条件的(不成环的)才merge
px = find(x)
py = find(y)
if px != py:
parents[px] = py
added_count += 1
ans += lenth
if added_count == n-1:
break
return ans
&spm=1001.2101.3001.5002&articleId=112967172&d=1&t=3&u=1cb0922c4d624a549ef51563ae916522)
854





