Python如何实现兼容NumPy接口的混合稀疏稠密轴COO数组
部分轴稠密的分块COO稀疏数组实现
你需要的「前N轴稀疏、剩余轴稠密」的高维稀疏结构,不需要修改现有稀疏库的底层逻辑,封装一个轻量类即可实现,完全兼容numpy的使用习惯。
核心存储逻辑
- 仅记录稀疏轴上非零块的坐标,坐标存储为形状为
(n_sparse_axis, n_nonzero_blocks)的整数数组,和标准COO格式的坐标存储逻辑一致 - 所有非零块统一存储为一个稠密numpy数组,形状为
(n_nonzero_blocks, *dense_axis_shape),避免拆分标量带来的存储开销 - 所有运算、索引逻辑优先复用numpy原生实现,不需要额外编写复杂的稀疏运算逻辑
完整实现代码
import numpy as np class BlockCOO: def __init__(self, coords, blocks, shape, sparse_axes=(0,1)): """ 分块COO稀疏数组,指定轴为稀疏轴,其余轴为稠密轴 :param coords: 稀疏轴坐标,形状(n_sparse_axes, n_nonzero_blocks) :param blocks: 非零稠密块,形状(n_nonzero_blocks, *dense_axes_shape) :param shape: 数组完整形状 :param sparse_axes: 指定为稀疏轴的轴索引,默认前两轴 """ self.sparse_axes = tuple(sparse_axes) self.dense_axes = tuple(i for i in range(len(shape)) if i not in self.sparse_axes) self.shape = tuple(shape) # 坐标合法性校验 self.coords = np.asarray(coords, dtype=np.int64) if self.coords.ndim != 2 or self.coords.shape[0] != len(self.sparse_axes): raise ValueError("坐标形状不匹配稀疏轴数量") self.n_blocks = self.coords.shape[1] # 稠密块合法性校验 self.blocks = np.asarray(blocks) expected_block_shape = tuple(self.shape[i] for i in self.dense_axes) if self.blocks.shape != (self.n_blocks, *expected_block_shape): raise ValueError(f"稠密块形状应为{(self.n_blocks, *expected_block_shape)},输入为{self.blocks.shape}") # 去重重复坐标 unique_coords, unique_idx = np.unique(self.coords, axis=1, return_index=True) self.coords = unique_coords self.blocks = self.blocks[unique_idx] def __getitem__(self, key): # 对齐numpy索引规则,先把索引补全到和维度数一致 if not isinstance(key, tuple): key = (key,) key = list(key) while len(key) < len(self.shape): key.append(slice(None)) # 先处理稀疏轴的索引筛选 mask = np.ones(self.n_blocks, dtype=bool) new_sparse_axes = [] new_sparse_shape = [] coord_transform = [] d_ptr = 0 for ax_idx, s_ax in enumerate(self.sparse_axes): ax_key = key[s_ax] ax_size = self.shape[s_ax] if isinstance(ax_key, (int, np.integer)): # 整数索引,该稀疏轴被消除,筛选对应坐标的块 mask &= (self.coords[ax_idx] == ax_key) coord_transform.append(None) elif isinstance(ax_key, slice): # 切片索引,计算坐标偏移 start, stop, step = ax_key.indices(ax_size) ax_coords = self.coords[ax_idx] block_mask = (ax_coords >= start) & (ax_coords < stop) & ((ax_coords - start) % step == 0) mask &= block_mask new_coords = (ax_coords[mask] - start) // step new_size = len(range(start, stop, step)) new_sparse_axes.append(ax_idx - d_ptr) new_sparse_shape.append(new_size) coord_transform.append(new_coords) else: # 花式索引/布尔索引直接转稠密处理,简化逻辑 return self.todense()[key] # 处理稠密轴的索引 dense_key = tuple(key[d_ax] for d_ax in self.dense_axes) selected_blocks = self.blocks[mask][(slice(None), *dense_key)] # 组装新形状 new_shape = [] s_ptr = 0 d_idx = 0 for i in range(len(self.shape)): if i in self.sparse_axes: if coord_transform[s_ptr] is not None: new_shape.append(new_sparse_shape[d_idx]) d_idx +=1 else: d_ptr +=1 s_ptr +=1 else: new_shape.append(selected_blocks.shape[1 + list(self.dense_axes).index(i)]) # 如果没有剩余稀疏轴,直接返回稠密数组 if len(new_sparse_axes) == 0: return selected_blocks[0] if selected_blocks.shape[0] == 1 else selected_blocks # 否则返回新的BlockCOO实例 new_coords = np.stack([ct for ct in coord_transform if ct is not None]) return BlockCOO(new_coords, selected_blocks, tuple(new_shape), sparse_axes=tuple(new_sparse_axes)) def todense(self): """转成原生numpy稠密数组""" arr = np.zeros(self.shape, dtype=self.blocks.dtype) # 构造稀疏轴的索引元组 sparse_idx = tuple(self.coords[i] for i in range(len(self.sparse_axes))) arr[sparse_idx] = self.blocks return arr def __array_ufunc__(self, ufunc, method, *inputs, **kwargs): """兼容numpy广播运算,自动把操作数转成稠密块参与运算""" proc_inputs = [] for inp in inputs: if isinstance(inp, BlockCOO): proc_inputs.append(inp.blocks) else: proc_inputs.append(inp) # 广播处理 broadcasted = np.broadcast_arrays(*proc_inputs) res_blocks = getattr(ufunc, method)(*broadcasted, **kwargs) # 输出形状处理 out_shape = np.broadcast_shapes(*[inp.shape if isinstance(inp, BlockCOO) else np.asarray(inp).shape for inp in inputs]) return BlockCOO(self.coords, res_blocks, out_shape, sparse_axes=self.sparse_axes)
使用示例
对应你提到的三维数组场景(第一轴稀疏、后两轴稠密):
# 构造示例:第一轴0位置是2x2稠密块,1位置全零 coords = np.array([[0]]) # 稀疏轴(第一轴)的非零坐标 blocks = np.array([[[1,2],[3,4]]]) # 对应坐标的稠密块 shape = (2,2,2) # 完整数组形状 sparse_arr = BlockCOO(coords, blocks, shape, sparse_axes=(0,)) # 索引测试 print(sparse_arr[0]) # 输出 [[1 2] # [3 4]] print(sparse_arr[1]) # 输出 [[0 0] # [0 0]] # 运算测试(兼容numpy广播) print((sparse_arr + 1)[0]) # 输出 [[2 3] # [4 5]] # 转稠密数组 print(sparse_arr.todense()) # 输出 [[[1 2] # [3 4]] # [[0 0] # [0 0]]]
如果是你需要的前两轴稀疏、后两轴稠密的四维数组,初始化时指定sparse_axes=(0,1)即可,存储逻辑和你预期完全一致:仅记录前两轴的非零坐标,每个坐标对应一个二维稠密矩阵块。
特性说明
- 存储效率远高于全轴COO:稠密块不需要拆分单个元素存储坐标,内存占用和预期的分块存储形式完全一致
- 索引接口完全对齐numpy:支持整数、切片、花式索引,索引结果自动根据剩余维度返回稠密数组或新的稀疏实例
- 运算兼容numpy广播:支持直接和标量、numpy数组做逐元素运算,规则和原生numpy保持一致
- 轻量无强依赖:仅依赖numpy,不需要安装额外稀疏计算库即可运行
如果需要和现有稀疏生态对接,可以自行扩展
to_coo方法,把稠密块展开成标量坐标即可生成标准COO格式对象。
内容的提问来源于stack exchange,提问作者Frank Seidl
相关产品推荐
相关产品推荐

