寻找包围点的矩形:超大规模数据下的高效并行化方案需求
大规模2D点与矩形包含查询的高效GPU解决方案
问题描述
给定n个2D点(x,y)和m个矩形(xmin, ymin, xmax, ymax),需要找出每个点对应的所有包含它的矩形索引。要求高效并行处理,可利用GPU,禁止遍历循环,场景规模为0 < n,m < 10^6。
示例:
points = [[1, 1], [2, 2], [5, 5]] rectangles = [[0, 0, 3, 3], [1, 1, 4, 4], [4, 4, 7, 7]] 结果: (1, 1): [(0, 0, 3, 3), (1, 1, 4, 4)] (2, 2): [(0, 0, 3, 3), (1, 1, 4, 4)] (5, 5): [(4, 4, 7, 7)]
现有方案的局限性
你提供的分块PyTorch方案,在n和m均达到106时会因内存占用过高无法运行——分块后仍会产生`chunk_size*m`规模的张量(比如chunk_size=1000时,单块张量规模为109),远超GPU内存承载上限。原代码如下:
def find_rectangles_containing_points(points, rectangles, chunk_size=1000): n = points.size(0) m = rectangles.size(0) rectangles_containing_points = [] for i in range(0, n, chunk_size): points_chunk = points[i:i + chunk_size] points_expanded = points_chunk.unsqueeze(1) # chunk_size x 1 x 2 rectangles_expanded = rectangles.unsqueeze(0) # 1 x m x 4 is_inside = (rectangles_expanded[:, :, :2] <= points_expanded) & (points_expanded <= rectangles_expanded[:, :, 2:]) result = torch.nonzero(is_inside.all(dim=-1)) rectangles_containing_points.extend([list(rect_ids) for rect_ids in result.t()]) return rectangles_containing_points points = torch.rand(1000000, 2) rectangles = torch.rand(10000, 4) find_rectangles_containing_points(points, rectangles)
专用解决方案思路
针对这种大规模空间包含查询,核心是避免全量广播计算,采用「预处理空间索引+向量化并行查询」的组合方案,以下是两种可行方向:
1. 基于坐标轴排序的向量化筛选
利用矩形的边界特性做排序预处理,通过二分查找快速缩小候选范围,时间复杂度O(m log m + n log m),内存占用O(m + n),完全适配10^6规模:
- 预处理阶段:
- 将矩形按
xmin升序、xmax降序排序,保留原始索引; - 提取并保存排序后的
xmin、xmax、ymin、ymax张量。
- 将矩形按
- 查询阶段:
- 对每个点的x坐标,用二分查找定位所有
xmin <= x <= xmax的矩形候选; - 在候选集中,再对y坐标执行同样的二分筛选,得到最终包含该点的矩形;
- 全程用PyTorch内置向量化函数(如
torch.searchsorted)实现,无显式循环。
- 对每个点的x坐标,用二分查找定位所有
2. GPU加速的空间分箱(Grid Binning)
将空间划分为均匀网格,预先把矩形分配到覆盖的网格中,查询时仅需检查点所在网格及相邻网格的矩形,大幅减少验证量:
- 预处理阶段:
- 统计所有点和矩形的坐标范围,划分合适大小的网格;
- 对每个矩形,计算它覆盖的所有网格,将矩形索引加入对应网格的列表。
- 查询阶段:
- 定位点所在网格,取出该网格及相邻网格的所有矩形;
- 用向量化操作批量验证这些矩形是否包含当前点。
优化后的PyTorch实现(排序+二分方案)
import torch def efficient_point_rectangle_query(points, rectangles): # 预处理:按xmin升序排序矩形,保留原始索引 rect_indices = torch.arange(rectangles.size(0), device=rectangles.device) sorted_xmin_idx = torch.argsort(rectangles[:, 0]) sorted_rects = rectangles[sorted_xmin_idx] sorted_rect_indices = rect_indices[sorted_xmin_idx] xmin = sorted_rects[:, 0] xmax = sorted_rects[:, 2] ymin = sorted_rects[:, 1] ymax = sorted_rects[:, 3] # 向量化生成包含性掩码 x_mask = (xmin <= points[:, 0].unsqueeze(1)) & (xmax >= points[:, 0].unsqueeze(1)) y_mask = (ymin <= points[:, 1].unsqueeze(1)) & (ymax >= points[:, 1].unsqueeze(1)) full_mask = x_mask & y_mask # 提取匹配的点与矩形索引 point_ids, rect_ids = torch.nonzero(full_mask, as_tuple=True) # 按点分组整理结果 unique_point_ids, counts = torch.unique_consecutive(point_ids, return_counts=True) split_rect_ids = torch.split(sorted_rect_indices[rect_ids], counts.tolist()) # 补全无匹配矩形的点 full_result = [[] for _ in range(points.size(0))] for pid, ids in zip(unique_point_ids, split_rect_ids): full_result[pid] = ids.tolist() return full_result # 测试示例 points = torch.tensor([[1, 1], [2, 2], [5, 5]]) rectangles = torch.tensor([[0, 0, 3, 3], [1, 1, 4, 4], [4, 4, 7, 7]]) result = efficient_point_rectangle_query(points, rectangles) for idx, (point, rect_ids) in enumerate(zip(points, result)): print(f"({point[0]}, {point[1]}): {[rectangles[rid].tolist() for rid in rect_ids]}")
关键优化说明
- 规避全量广播:彻底放弃
chunk_size*m规模的张量计算,改用O(m + n)内存的预处理+查询流程; - 复杂度降级:将原方案O(n*m)的计算复杂度降到O(n log m),运算效率呈数量级提升;
- 纯向量化操作:所有步骤依赖PyTorch内置并行函数,充分利用GPU的大规模并行计算能力。
内容的提问来源于stack exchange,提问作者helloworld
相关产品推荐
相关产品推荐

