求助优化Python代码运行速度:8.25万行CSV坐标分组任务
坐标分组代码的效率优化建议
原问题概述
需要处理含约82500行的CSV文件,将x/y/z三个坐标相互差值均不超过±5(由margin控制)的行归为一组,为对应行新增第4列标记分组ID;无匹配项的行该列留空。原代码在小测试文件上可行,但处理全量数据耗时约9小时,需要优化性能。
原代码的核心性能问题
- 双重循环导致O(n²)复杂度:82500行的情况下,循环次数约为(82500)²≈6.8e9次,这是耗时的根本原因。
- 重复类型转换:每次循环都将字符串格式的坐标转为float,重复计算浪费资源。
- 冗余逻辑判断:多次重复检查行的长度、坐标边界,且分组逻辑未处理传递性(比如A和B相似、B和C相似,但原代码可能把A和C分到不同组)。
优化方案
1. 预处理数据,避免重复计算
先将所有坐标行的字符串转为数值类型,同时保留原始索引,方便后续对应到原数据行。
2. 使用Union-Find(并查集)算法处理连通分组
这是处理“相似即归组”这类具有传递性分组问题的高效算法,合并和查找操作的时间复杂度接近O(1),整体复杂度主要由排序决定(O(n log n))。
3. 排序后减少比较范围
先按x坐标排序,这样只需要检查x坐标在当前行±margin范围内的行,无需遍历所有行,进一步减少计算量。
优化后的代码
import csv class UnionFind: def __init__(self, size): self.parent = list(range(size)) self.rank = [0]*size def find(self, x): if self.parent[x] != x: self.parent[x] = self.find(self.parent[x]) return self.parent[x] def union(self, x, y): x_root = self.find(x) y_root = self.find(y) if x_root == y_root: return if self.rank[x_root] < self.rank[y_root]: self.parent[x_root] = y_root else: self.parent[y_root] = x_root if self.rank[x_root] == self.rank[y_root]: self.rank[x_root] += 1 def main(): margin = 5 input_path = 'just xyz.csv' output_path = 'new manual clusters.csv' # 读取并预处理数据:保留表头,将坐标转为数值,记录原始索引 with open(input_path, newline='') as f: reader = csv.reader(f) header = next(reader) data = [] for idx, row in enumerate(reader, start=1): # idx从1开始,对应原代码的行号 x = float(row[0]) y = float(row[1]) z = float(row[2]) data.append( (x, y, z, idx) ) # 按x坐标排序,减少后续比较范围 data.sort() n = len(data) uf = UnionFind(n) # 遍历每个点,只检查x在当前点±margin范围内的后续点 for i in range(n): x_i, y_i, z_i, idx_i = data[i] # 因为已排序,x只会递增,超过x_i+margin就可以停止 for j in range(i+1, n): x_j, y_j, z_j, idx_j = data[j] if x_j - x_i > margin: break # 检查y和z的差值 if abs(y_j - y_i) <= margin and abs(z_j - z_i) <= margin: uf.union(i, j) # 构建分组结果:key是原始行号,value是分组ID group_map = {} # 给每个连通分量分配唯一ID root_to_id = {} current_id = 0 for i in range(n): root = uf.find(i) if root not in root_to_id: root_to_id[root] = current_id current_id += 1 original_idx = data[i][3] group_map[original_idx] = root_to_id[root] # 生成输出数据:给原行添加分组ID # 先初始化结果列表,表头加"Group"列 result = [header + ["Group"]] # 原数据行数是n+1(表头+数据行),初始化空行 for _ in range(n): result.append( ['', '', '', ''] ) # 重新读取原数据,填充内容和分组ID with open(input_path, newline='') as f: reader = csv.reader(f) next(reader) # 跳过表头 for row_idx, row in enumerate(reader, start=1): result[row_idx] = row + [group_map.get(row_idx, '')] # 写入输出文件 with open(output_path, 'w', newline='') as f: writer = csv.writer(f) writer.writerows(result) print("Done") if __name__ == "__main__": main()
优化效果说明
- 时间复杂度从O(n²)降到O(n log n),8万行数据的处理时间会从小时级缩短到分钟甚至秒级。
- 处理了分组的传递性:只要两个点通过中间点间接相似,就会被分到同一组,符合实际需求。
- 避免了重复的类型转换和冗余判断,进一步提升效率。
内容的提问来源于stack exchange,提问作者compto2017
相关产品推荐
相关产品推荐

