Python分治最近点对算法≥7输入报TypeError的修复方法
问题修复:最近点对分治算法类型错误
错误原因
报错的核心是函数返回值顺序不统一:
brute_force函数返回(最近点对, 最小距离)(元组在前,浮点数在后)strip_closest函数返回(最小距离, 最近点对)(浮点数在前,元组在后)
当点数量≥7时,递归深度增加,会混合两种不同结构的返回值,导致执行min(min_left[1], min_right[1])时,一个操作数是距离(浮点数),另一个是点对(元组),触发类型比较错误。
另外还有潜在bug:递归处理右分支时,传入的长度参数错误(用mid而非右分支实际长度),会导致奇数个点时右分支长度计算错误。
修复步骤
- 统一返回值顺序:将
strip_closest的返回值顺序调整为和brute_force一致,即(最近点对, 最小距离)。 - 修正右分支递归参数:递归调用右分支时,传入右分支实际长度
len(right)而非mid。
修复后的完整代码
import math import sys class Point(): def __init__(self, x, y): self.x = x self.y = y def __repr__(self): return f"({self.x}, {self.y})" """ Returns the Euclidean distance between two points """ def distance(point1, point2): return math.sqrt((point1.x - point2.x)**2 + (point1.y - point2.y)**2) def sort_points(Points): points_sorted = sorted(Points, key=lambda p: [p.x, p.y]) return points_sorted def brute_force(Points, n): min_dist = sys.float_info.max closest_pair = (Points[0], Points[1]) for i in range(n): for j in range(i+1, n): dist = distance(Points[i], Points[j]) if dist < min_dist: min_dist = dist closest_pair = (Points[i], Points[j]) return closest_pair, min_dist def recursion(Points, n): points_sorted = sort_points(Points) # Base Case if n <= 3: return brute_force(points_sorted, n) # Find mid-point mid = n//2 mid_point = points_sorted[mid] # Split the points into two branches left = points_sorted[:mid] right = points_sorted[mid:] # Recursively find the smallest distance on the left and right min_left = recursion(left, len(left)) min_right = recursion(right, len(right)) # Find the closest pair of the two sides min_distance = min(min_left[1], min_right[1]) if min_left[1] <= min_right[1]: min_pair = min_left else: min_pair = min_right # Build the strip array to find points smaller than delta delta = min_distance strip = [] for i in range(n): if abs(points_sorted[i].x - mid_point.x) < min_distance: strip.append(points_sorted[i]) # Return closest pair or even closer if found in the strip return strip_closest(strip, min_pair, min_distance) def strip_closest(strip, min_pair, min_distance): strip_min_dist = min_distance strip_min_pair = min_pair # This loop will run at most 6 times for i in range(len(strip)): for j in range(i+1, min(i+7, len(strip))): dist = distance(strip[i], strip[j]) if dist < strip_min_dist: strip_min_dist = dist strip_min_pair = (strip[i], strip[j]) # 统一返回顺序:(点对, 距离) return strip_min_pair, strip_min_dist # Driver code Points = [Point(15, -37), Point(-45, -36), Point(19, -18), Point(-76, 64), Point(0, -30), Point(-47, -33), Point(7, 8), Point(0, 8)] n = len(Points) print(Points) print(recursion(Points, n))
额外优化说明
- 给
brute_force初始化closest_pair,避免n=2时可能的未定义问题 - 调整
sort_points中的lambda参数名(用p代替Point,避免和类名冲突) - 修正
strip_closest外层循环范围(原range(len(strip)-1)会漏掉最后一个点的比较,改为range(len(strip))更合理)
内容的提问来源于stack exchange,提问作者elliot999
相关产品推荐
相关产品推荐

