You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

寻找包围点的矩形:超大规模数据下的高效并行化方案需求

大规模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)实现,无显式循环。

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.03 21:54:59