如何为bbox_sort函数传入可选参数thresh?替代全局变量方案
解决方案:重构边界框排序函数,避免全局变量
首先,不要用文件存储thresh,这种做法完全没必要,会增加复杂度还降低灵活性,直接通过参数传递是更合理的选择。下面提供几种重构方案,都能避免全局变量,同时让thresh成为可选参数:
方案1:使用闭包(直观易懂)
把bbox_sort定义在一个外层函数里,让它能直接访问外层的thresh参数,调用时可灵活指定阈值:
from functools import cmp_to_key def get_bbox_sorter(thresh=10): def bbox_sort(a, b): # 先判断Y轴高度差是否小于等于阈值 if abs(a[1] - b[1]) <= thresh: # 同一行内按X轴从左到右排序 return a[0] - b[0] # 不同行按Y轴从上到下排序 return a[1] - b[1] return bbox_sort def get_prediction(result, thresh=10): coord_list = [] res = result.to_coco_annotations() for ann in res: x, y, w, h = ann['bbox'] coord_list.append((x, y, w, h)) # 获取绑定了指定thresh的排序函数 cnts = sorted(coord_list, key=cmp_to_key(get_bbox_sorter(thresh))) # 优化索引查找:用字典替代index方法,避免重复bbox匹配错误且提升性能 bbox_to_index = {bbox: idx for idx, bbox in enumerate(cnts)} for ann in res: ann['image_id'] = bbox_to_index[tuple(ann['bbox'])] return res
方案2:使用functools.partial绑定参数
通过partial把thresh参数绑定到bbox_sort上,让它符合cmp_to_key要求的双参数函数格式:
from functools import cmp_to_key, partial def bbox_sort(a, b, thresh=10): if abs(a[1] - b[1]) <= thresh: return a[0] - b[0] return a[1] - b[1] def get_prediction(result, thresh=10): coord_list = [] res = result.to_coco_annotations() for ann in res: x, y, w, h = ann['bbox'] coord_list.append((x, y, w, h)) # 绑定thresh参数,生成符合要求的比较函数 sorted_func = partial(bbox_sort, thresh=thresh) cnts = sorted(coord_list, key=cmp_to_key(sorted_func)) # 优化索引查找 bbox_to_index = {bbox: idx for idx, bbox in enumerate(cnts)} for ann in res: ann['image_id'] = bbox_to_index[tuple(ann['bbox'])] return res
方案3:改用key函数(更高效,推荐)
Python的sorted用key函数比cmp_to_key效率更高,我们可以把排序逻辑转换成生成key的规则,完全不用比较函数:
def get_bbox_key(thresh=10): def key_func(bbox): x, y, _, _ = bbox # 把Y轴坐标按阈值分组,同一组按X轴排序,不同组按Y轴排序 return (y // thresh, x) return key_func def get_prediction(result, thresh=10): coord_list = [] res = result.to_coco_annotations() for ann in res: x, y, w, h = ann['bbox'] coord_list.append((x, y, w, h)) # 使用key函数排序,无需cmp_to_key cnts = sorted(coord_list, key=get_bbox_key(thresh)) # 优化索引查找 bbox_to_index = {bbox: idx for idx, bbox in enumerate(cnts)} for ann in res: ann['image_id'] = bbox_to_index[tuple(ann['bbox'])] return res
补充说明:
- 原代码中
cnts.index(tuple(res[pred]['bbox']))存在性能问题(每次查找是O(n)),且如果有重复bbox会返回第一个匹配的索引,改用字典映射bbox_to_index可将查找复杂度降到O(1),结果也更准确。 - 把
thresh作为get_prediction的可选参数,调用时可根据需求灵活修改,比如get_prediction(result, thresh=15)。
内容的提问来源于stack exchange,提问作者ИНДУС Геймдев
相关产品推荐
相关产品推荐

