使用NumPy meshgrid提取n维空间立方体与超立方体坐标
n维空间分箱提取超立方体单元格实现方案
问题定位
你卡壳的核心原因有两个:
- 调用
meshgrid时未指定indexing='ij'参数,默认的xy索引模式会调换前两个维度的顺序,3维及以上场景下生成的坐标和维度对应关系完全错乱,自然匹配不到正确的立方体 - 没理清网格顶点数组的索引逻辑:生成的顶点数组形状为
(n_bins+1, n_bins+1, ..., d),前d个维度的索引值,对应该顶点在对应维度上的分箱边界序号;单个超立方体的所有顶点,本质就是每个维度取相邻两个分箱边界,围成的区域内的所有点。
补充:3维场景下每个维度分3箱时总共有3^3=27个立方体单元格,你提到的9个单元格是2维场景下33分箱的结果,按需调整维度层数即可。*
实现代码
第一步:生成分箱边界
兼容你现有的逻辑,支持每个维度自定义数值范围、分箱粒度:
import numpy as np # 每个维度单独定义分箱边界,范围、分箱数可以不一样 dim_boundaries = [ np.linspace(0, 2, 3 + 1), # 维度0:0~2范围分3箱 np.linspace(0, 3, 3 + 1), # 维度1:0~3范围分3箱 np.linspace(0, 1, 3 + 1) # 维度2:0~1范围分3箱 ] d = len(dim_boundaries) n_bins_per_dim = [len(b) - 1 for b in dim_boundaries] # 每个维度的分箱数量
第二步:生成维度顺序正确的网格顶点
补上indexing='ij'参数,保证维度顺序和输入完全一致:
grid = np.stack(np.meshgrid(*dim_boundaries, indexing='ij')).T
第三步:提取所有超立方体单元格
固定3维场景写法
逻辑直观,方便调试:
cubes = [] for i in range(n_bins_per_dim[0]): for j in range(n_bins_per_dim[1]): for k in range(n_bins_per_dim[2]): # 切片取每个维度相邻两个边界围成区域的所有顶点,共2^3=8个 cube_vertices = grid[i:i+2, j:j+2, k:k+2, :] # 直接获取当前立方体在各维度的数值范围 cube_ranges = [ [dim_boundaries[0][i], dim_boundaries[0][i+1]], [dim_boundaries[1][j], dim_boundaries[1][j+1]], [dim_boundaries[2][k], dim_boundaries[2][k+1]] ] cubes.append( (cube_ranges, cube_vertices) )
任意d维通用写法
不需要根据维度写死循环,2维、3维及更高维场景都能用:
cubes = [] # 遍历所有分箱的索引组合 for bin_index in np.ndindex(*n_bins_per_dim): # 构造每个维度的切片规则,取相邻两个边界的区间 slices = tuple(slice(idx, idx + 2) for idx in bin_index) + (slice(None),) cube_vertices = grid[slices] # 提取当前立方体的各维度边界范围 cube_ranges = [] for dim, idx in enumerate(bin_index): b = dim_boundaries[dim] cube_ranges.append( [b[idx], b[idx+1]] ) cubes.append( (cube_ranges, cube_vertices) )
扩展使用提示
- 如果需要判断任意坐标点属于哪个单元格,不用遍历所有立方体,用
np.digitize分别计算点在每个维度所属的bin序号,组合起来就是对应立方体的索引,查询复杂度为O(d),高维场景下性能远高于遍历 - 提取出的
cube_vertices形状固定为(2,2,...,2,d),d维场景下每个立方体正好有2^d个顶点,和几何定义完全一致 - 如果只需要单元格的边界范围不需要顶点坐标,可以跳过meshgrid生成步骤,直接遍历分箱索引组合生成
cube_ranges,内存占用会大幅降低,高维大分箱场景下优化效果非常明显
内容的提问来源于stack exchange,提问作者TheHoly Pebble
相关产品推荐
相关产品推荐

