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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.02 07:51:24