如何在Python中进一步加速指定类别与重量范围的Widget查询?
问题
假设有如下简单dataclass定义的Widget类:
@dataclass class Widget: category: str weight: int
需要实现一个query_widgets函数,返回属于指定分类、且重量处于[lo_weight, hi_weight]区间内的所有Widget。已知Widget列表固定,可进行任意预处理操作,求纯Python环境下的最快实现方式?
我当前的实现思路是:先将列表按weight排序,构建分类到对应Widget的映射表,再通过查找重量的首尾索引快速定位查询范围。尝试用bisect替代循环查找索引,但性能反而略有下降。完整代码如下:
import random import time from collections import defaultdict from dataclasses import dataclass from typing import Iterable @dataclass class Widget: categories: str weight: int class WidgetSearch: def __init__(self, widgets: Iterable[Widget]) -> None: self.widgets = sorted(widgets, key=lambda w: w.weight) self.category_to_widgets = defaultdict(list) self.category_weight_start_idx = defaultdict(dict) self.category_weight_end_idx = defaultdict(dict) for m in self.widgets: for category in m.categories: self.category_to_widgets[category].append(m) for category, widgets in self.category_to_widgets.items(): prev_weight = None for i, widget in enumerate(widgets): weight = widget.weight if prev_weight != weight: self.category_weight_start_idx[category][weight] = i if prev_weight is not None: self.category_weight_end_idx[category][prev_weight] = i prev_weight = weight self.category_weight_end_idx[category][prev_weight] = len(widgets) def _find_most_recent_start_idx(self, category: str, weight: int) -> int: if weight in self.category_weight_start_idx[category]: return self.category_weight_start_idx[category][weight] for idx_weight, idx in self.category_weight_start_idx[category].items(): if idx_weight > weight: return idx return 0 def _find_most_recent_end_idx(self, category: str, weight: int) -> int: if weight in self.category_weight_end_idx[category]: return self.category_weight_end_idx[category][weight] for idx_weight in reversed(self.category_weight_end_idx[category].keys()): if idx_weight < weight: return self.category_weight_end_idx[category][idx_weight] return None def query_widgets(self, category: str, lo_weight: int, hi_weight: int) -> Iterable[Widget]: if category not in self.category_to_widgets or lo_weight > hi_weight: return [] start = self._find_most_recent_start_idx(category, lo_weight) end = self._find_most_recent_end_idx(category, hi_weight) return self.category_to_widgets[category][start:end] if __name__ == '__main__': widgets = [] categories = list('ABCDEF') for _ in range(10000): cats = random.sample(categories, random.randint(1, 3)) widgets.append(Widget(categories=cats, weight=random.randint(1000, 2000))) ws = WidgetSearch(widgets) start = time.perf_counter_ns() widgets = ws.query_widgets('C', 1400, 1800) end = time.perf_counter_ns() dur = (end - start) / 1000.0 print(f'# Found: {len(widgets)} in {dur} microsec')
最优实现方案
你的思路方向没问题,但之前bisect性能不佳是因为预处理结构没配合好bisect的特性。以下是更高效的实现方案:
核心优化思路
- 先分组再排序:不要全局排序后分组,而是先按分类分组,再对每个分组内的Widget按weight升序排序。这样每个分类的列表本身就是有序的,直接适配bisect的二分查找需求。
- 为每个分类维护单独的重量列表:bisect需要针对有序数值列表操作,所以为每个分类单独存储一份排序后的weight列表,用于快速定位索引,避免直接操作Widget对象带来的开销。
优化后代码
import random import time import bisect from collections import defaultdict from dataclasses import dataclass from typing import Iterable, List @dataclass class Widget: categories: List[str] weight: int class WidgetSearch: def __init__(self, widgets: Iterable[Widget]) -> None: # 按分类分组 self.category_groups = defaultdict(list) for widget in widgets: for cat in widget.categories: self.category_groups[cat].append(widget) # 对每个分类的列表按weight排序,并提取对应的weight列表用于bisect self.category_sorted_widgets = {} self.category_weights = {} for cat, group in self.category_groups.items(): # 按weight升序排序 sorted_group = sorted(group, key=lambda w: w.weight) self.category_sorted_widgets[cat] = sorted_group # 提取排序后的weight列表,供bisect使用 self.category_weights[cat] = [w.weight for w in sorted_group] def query_widgets(self, category: str, lo_weight: int, hi_weight: int) -> List[Widget]: if category not in self.category_sorted_widgets or lo_weight > hi_weight: return [] weights = self.category_weights[category] widgets = self.category_sorted_widgets[category] # 找到第一个>=lo_weight的索引 start_idx = bisect.bisect_left(weights, lo_weight) # 找到第一个>hi_weight的索引 end_idx = bisect.bisect_right(weights, hi_weight) return widgets[start_idx:end_idx] if __name__ == '__main__': widgets = [] categories = list('ABCDEF') for _ in range(10000): cats = random.sample(categories, random.randint(1, 3)) widgets.append(Widget(categories=cats, weight=random.randint(1000, 2000))) ws = WidgetSearch(widgets) # 测试查询性能 start = time.perf_counter_ns() result = ws.query_widgets('C', 1400, 1800) end = time.perf_counter_ns() dur = (end - start) / 1000.0 print(f'# Found: {len(result)} in {dur:.2f} microsec')
性能优势说明
- 预处理效率:分组后单独排序的逻辑更简洁,避免了全局排序后再分组的额外操作,预处理时间和原方案相当甚至更优。
- 查询效率:
bisect_left和bisect_right都是O(log n)时间复杂度,远快于原方案中遍历字典键的O(n)操作,尤其是当分类下的Widget数量较多时,性能提升非常明显。 - 内存开销:仅额外存储一份每个分类的weight列表,内存占用极低,完全可以忽略。
内容的提问来源于stack exchange,提问作者postelrich
相关产品推荐
相关产品推荐

