本文介绍一种基于字典的轻量级稀疏多维容器设计,专为高维(5–100维)、中等尺寸(每维20–200)坐标索引与跨维高效求和场景优化,兼顾内存效率与操作灵活性,适用于频率建模、边际分布计算等任务。
本文介绍一种基于字典的轻量级稀疏多维容器设计,专为高维(5–100维)、中等尺寸(每维20–200)坐标索引与跨维高效求和场景优化,兼顾内存效率与操作灵活性,适用于频率建模、边际分布计算等任务。
在处理高维稀疏计数或概率分布(如多变量离散联合分布)时,传统 numpy.ndarray 因需预分配完整 n^d 空间而极易造成内存爆炸——例如 d=50, n=100 时,即使每个元素仅占 8 字节,理论内存也远超宇宙原子总数。此时,稀疏性成为核心突破口:实际非零项通常不足百万分之一。Scipy 的 coo_array 虽支持稀疏存储,但仅限 1D/2D,无法直接满足 d ≥ 5 的高维需求。
理想的解决方案应满足三点本质要求:
✅ 隐式零初始化:未显式赋值的坐标默认返回 0,无需预先构造全零结构;
✅ 坐标式随机访问:支持 obj[(i1, i2, ..., id)] 语法读写,时间复杂度 O(1);
✅ 维度聚合高效:能快速计算沿任意轴组合的 sum(axis=(a,b,c)),避免全量展开。
推荐方案:SparseTensor —— 基于 dict[tuple, scalar] 的定制类
我们不依赖外部稀疏库,而是构建一个语义清晰、内存极致精简的纯 Python 容器:
from collections import defaultdict
from typing import Dict, Tuple, Any, Union, Iterable
class SparseTensor:
def __init__(self, shape: Tuple[int, ...], dtype=float):
self.shape = shape
self._data: Dict[Tuple[int, ...], Any] = {}
self._dtype = dtype
def __getitem__(self, key: Tuple[int, ...]) -> Any:
# 边界检查(可选)
if not all(0 'SparseTensor':
"""沿指定轴求和,返回新的 SparseTensor"""
if axis is None:
return sum(self._data.values(), self._dtype(0))
axes = (axis,) if isinstance(axis, int) else tuple(axis)
axes = tuple(a % len(self.shape) for a in axes) # 支持负索引
keep_dims = tuple(i for i in range(len(self.shape)) if i not in axes)
result = defaultdict(self._dtype)
for coords, val in self._data.items():
reduced_key = tuple(coords[i] for i in keep_dims)
result[reduced_key] += val
# 构造新 shape:被 sum 的轴长度变为 1(若需保留维度),但此处返回降维结果
new_shape = tuple(self.shape[i] for i in keep_dims)
out = SparseTensor(new_shape, self._dtype)
out._data = dict(result)
return out
def to_dense(self) -> 'np.ndarray':
"""仅用于调试:转为稠密数组(慎用!)"""
import numpy as np
arr = np.zeros(self.shape, dtype=self._dtype)
for coords, val in self._data.items():
arr[coords] = val
return arr
使用示例与性能说明
# 初始化 10×10×10 稀疏张量
st = SparseTensor((10, 10, 10))
# 批量更新(等价于你的 data 预处理)
data = [((1,0,5),6), ((2,6,5),100), ((5,3,1),1), ((2,0,5),4), ((2,6,5),100)]
for coords, value in data:
st[coords] += value
# 高效跨维求和(仅遍历非零项)
x = st.sum(axis=(1, 2)) # shape (10,), 沿后两维压缩 → O(nnz) 时间
y = st.sum(axis=0) # shape (10, 10)
z = st.sum(axis=(0, 2)) # shape (10,)
print("Marginal x:", x._data) # {(0,): 6.0, (2,): 204.0, (5,): 1.0}
- 内存优势:仅存储 (coords, value) 对,nnz × (d × 8 + 8) 字节(典型情况);
- 求和复杂度:O(nnz),与非零元数量线性相关,不受总维度 n^d 影响;
- 扩展性:天然支持任意 d ≥ 1,无维度硬限制;
-
注意事项:
- 若需频繁切片(如 st[:, 3, :]),建议额外维护按各轴排序的索引映射(如 axis_to_keys[ax] = sorted(keys, key=lambda k: k[ax])),配合二分查找加速;
- 对于超大规模 nnz > 10⁷ 场景,可替换底层为 numba 加速的 typed.Dict 或 pyarrow 表提升吞吐;
- 若需广播、张量积等高级操作,建议封装后对接 xarray 或 einops。
该设计回归数据本质:高维稀疏结构的核心不是“矩阵”,而是键值映射。通过放弃稠密布局的假定,换取了对真实问题规模的精准刻画——这正是科学计算中“合适工具”的真正含义。
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!











