优化重心细分算法时间复杂度:无法完成7次迭代请求
问题描述
我在研究中尝试复现一篇文章中的图表,但无法完成7次及以上迭代。作为编程新手,我发现代码时间复杂度为O(6^n)(重心细分将三角形分为6个,以此类推),6次迭代耗时约20秒,7次始终无法完成。我的conventional_form函数基于文章规则:给定三角形,先缩放至最长边为1,旋转使最长边水平,必要时翻转确保顶点位置符合规范。以下是我的代码:
import numpy as np import matplotlib.pyplot as plt def conventional_form(a, b, c): # Find longest side and scale to length 1 side1 = np.abs(a - b) side2 = np.abs(b - c) side3 = np.abs(c - a) if side1 >= side2 and side1 >= side3: longest_side = side1 v1, v2, v3 = a, b, c elif side2 >= side1 and side2 >= side3: longest_side = side2 v1, v2, v3 = b, c, a else: longest_side = side3 v1, v2, v3 = c, a, b scaling_factor = 1 / longest_side scaled_v1 = v1 * scaling_factor scaled_v2 = v2 * scaling_factor scaled_v3 = v3 * scaling_factor # Rotating scaled_vertices so the longest side is horizontal and the other vertex is above the horizontal line theta = np.angle(v2 - v1) rotation_factor = np.exp(-1j * theta) rotated_v1 = scaled_v1 * rotation_factor rotated_v2 = scaled_v2 * rotation_factor rotated_v3 = scaled_v3 * rotation_factor if rotated_v3.imag < rotated_v1.imag: rotated_v3 = complex(rotated_v3.real, 2 * rotated_v1.imag - rotated_v3.imag) # Translate rotated_vertices so v1 = 0 and v2 = 0 + 1j translation = -rotated_v1 a_new = rotated_v1 + translation b_new = rotated_v2 + translation c_new = rotated_v3 + translation # Moving top vertex to left-side if needed if c_new.real >= 0.50: c_new = complex(1 - c_new.real, c_new.imag) return a_new, b_new, c_new def plot_point(point, color='black', markersize=2): plt.scatter(point.real, point.imag, color=color, s=markersize) def barycentric_subdivision(a, b, c, subdivisions): if subdivisions == 0: return # Transform the vertices to the conventional form a_new, b_new, c_new = conventional_form(a, b, c) # Plot the c_new vertex plot_point(c_new) # Define the 6 new triangles ab_mid = (a + b) / 2 bc_mid = (b + c) / 2 ca_mid = (c + a) / 2 centroid = (a + b + c) / 3 triangles = [ (a, ab_mid, centroid), (ab_mid, b, centroid), (b, bc_mid, centroid), (bc_mid, c, centroid), (c, ca_mid, centroid), (ca_mid, a, centroid) ] # Recursively apply barycentric subdivision to each new triangle for tri in triangles: barycentric_subdivision(*tri, subdivisions - 1) def calculate_plotted_points(subdivisions): return 1 + 6 ** (subdivisions - 1) # Input a = complex(1, 1.5) b = complex(1.5, 2.5) c = complex(0.5, 2.5) subdivisions = 7 plt.figure() barycentric_subdivision(a, b, c, subdivisions) plt.xlim(0, 1) plt.ylim(0, 1) num_plotted_points = calculate_plotted_points(subdivisions) plt.title(f"Subdivisions: {subdivisions}\nPlotted Points: {num_plotted_points}") plt.show()
优化建议
1. 用迭代替代递归,消除栈开销
n=7时递归会产生6^7=279936次函数调用,Python的函数调用开销会被放大。改用队列存储待处理三角形的迭代方式,能大幅降低这部分开销。
2. 批量绘图,减少渲染次数
原代码每次绘制单个点都调用plt.scatter,IO和渲染成本极高。应该先收集所有需要绘制的点,最后一次性调用plt.scatter完成绘制。
3. 简化conventional_form运算逻辑
将复数转换为numpy实数数组,用矩阵运算合并旋转、平移操作,减少复数转换的额外开销;同时只返回需要绘制的顶点,不需要保留整个三角形的三个顶点,节省内存和计算量。
4. 利用numpy向量化加速
用numpy数组存储顶点坐标,替代单个复数,能利用numpy的底层优化加速运算。
修改后的代码示例
import numpy as np import matplotlib.pyplot as plt def conventional_form(a, b, c): # 转换为numpy实数坐标数组 pts = np.array([[a.real, a.imag], [b.real, b.imag], [c.real, c.imag]]) # 计算各边长度 side1 = np.linalg.norm(pts[0] - pts[1]) side2 = np.linalg.norm(pts[1] - pts[2]) side3 = np.linalg.norm(pts[2] - pts[0]) # 确定最长边与对应顶点顺序 max_side = max(side1, side2, side3) if max_side == side1: v1, v2, v3 = pts[0], pts[1], pts[2] elif max_side == side2: v1, v2, v3 = pts[1], pts[2], pts[0] else: v1, v2, v3 = pts[2], pts[0], pts[1] # 缩放至最长边为1 scaling_factor = 1 / max_side v1_scaled = v1 * scaling_factor v2_scaled = v2 * scaling_factor v3_scaled = v3 * scaling_factor # 旋转使最长边水平 dx = v2_scaled[0] - v1_scaled[0] dy = v2_scaled[1] - v1_scaled[1] theta = np.arctan2(dy, dx) rot_mat = np.array([[np.cos(-theta), -np.sin(-theta)], [np.sin(-theta), np.cos(-theta)]]) # 平移到原点后旋转 v3_rot = rot_mat @ (v3_scaled - v1_scaled) # 确保顶点在水平线上方 if v3_rot[1] < 0: v3_rot[1] = -v3_rot[1] # 平移到标准位置(v1=(0,0), v2=(1,0)) v3_new = v3_rot # 翻转到左半侧 if v3_new[0] >= 0.5: v3_new[0] = 1 - v3_new[0] return v3_new def barycentric_subdivision_iterative(a, b, c, subdivisions): # 用队列存储待处理的三角形(顶点a, 顶点b, 顶点c, 剩余细分次数) queue = [(a, b, c, subdivisions)] points = [] while queue: a_tri, b_tri, c_tri, sub = queue.pop() if sub == 0: continue # 获取需要绘制的点 plot_pt = conventional_form(a_tri, b_tri, c_tri) points.append(plot_pt) # 计算细分后的新顶点 ab_mid = (a_tri + b_tri) / 2 bc_mid = (b_tri + c_tri) / 2 ca_mid = (c_tri + a_tri) / 2 centroid = (a_tri + b_tri + c_tri) / 3 # 将新三角形加入队列(保持原递归的深度优先顺序) queue.append((ca_mid, a_tri, centroid, sub-1)) queue.append((c_tri, ca_mid, centroid, sub-1)) queue.append((bc_mid, c_tri, centroid, sub-1)) queue.append((b_tri, bc_mid, centroid, sub-1)) queue.append((ab_mid, b_tri, centroid, sub-1)) queue.append((a_tri, ab_mid, centroid, sub-1)) return np.array(points) def calculate_plotted_points(subdivisions): return 1 + 6 ** (subdivisions - 1) # 输入参数 a = complex(1, 1.5) b = complex(1.5, 2.5) c = complex(0.5, 2.5) subdivisions = 7 plt.figure() # 批量获取所有点并绘制 points = barycentric_subdivision_iterative(a, b, c, subdivisions) plt.scatter(points[:, 0], points[:, 1], color='black', s=2) plt.xlim(0, 1) plt.ylim(0, 1) num_plotted_points = calculate_plotted_points(subdivisions) plt.title(f"Subdivisions: {subdivisions}\nPlotted Points: {num_plotted_points}") plt.show()
关键修改说明
- 迭代替代递归:用队列消除递归函数调用的开销,同时便于监控计算进度
- 批量绘图:将所有绘制操作合并为一次,大幅降低渲染开销
- 简化坐标运算:用numpy实数数组替代复数,合并旋转、平移操作,减少冗余计算
- 减少数据冗余:仅保留需要绘制的顶点,节省内存占用
内容的提问来源于stack exchange,提问作者lambdaserb
相关产品推荐
相关产品推荐

