
本文介绍一种轻量、内存友好的 python 实现方案,用于构建支持坐标索引与跨维高效求和的稀疏高维数组(d=5–100),避免 numpy 全量数组的内存开销,适用于频率统计、概率分布建模等场景。
本文介绍一种轻量、内存友好的 python 实现方案,用于构建支持坐标索引与跨维高效求和的稀疏高维数组(d=5–100),避免 numpy 全量数组的内存开销,适用于频率统计、概率分布建模等场景。
在处理高维稀疏计数或概率分布(如 d=5–100 维、每维尺寸 n≈20–200)时,使用 np.zeros((n,)*d) 会因维度爆炸导致内存不可行(例如 d=10, n=100 → 100¹⁰ ≈ 10²⁰ 个元素)。此时,标准稀疏格式(如 SciPy 的 coo_array、csr_matrix)受限于仅支持 ≤2D,无法直接满足需求。理想方案需满足:
- 支持任意 d 维整数坐标(如 (i₁,i₂,…,i_d))随机访问与累加;
- 未显式赋值的坐标默认返回 0(即“隐式零”语义);
- 支持沿任意轴组合(如 axis=(1,3,5))的高效求和(即边缘化/降维);
- 内存占用严格正比于非零元数量。
推荐方案:基于 dict 的自定义稀疏张量类
核心思想是用 dict[tuple[int, ...], float] 存储非零项,辅以维度信息与求和逻辑封装:
from collections import defaultdict
from typing import Dict, Tuple, Union, Iterable, Optional
class SparseTensor:
def __init__(self, shape: Tuple[int, ...]):
self.shape = shape
self._data: Dict[Tuple[int, ...], float] = {}
# 可选:预计算各维度索引范围,用于边界检查(提升调试鲁棒性)
def __getitem__(self, key: Tuple[int, ...]) -> float:
if len(key) != len(self.shape):
raise ValueError(f"Key dimension {len(key)} ≠ tensor dimension {len(self.shape)}")
if any(not (0 Union['SparseTensor', float]:
"""沿指定轴求和,返回新 SparseTensor 或标量"""
if axis is None:
return sum(self._data.values()) # 总和
# 标准化 axis 为元组
if isinstance(axis, int):
axes = (axis,)
else:
axes = tuple(sorted(set(axis))) # 去重并排序
# 验证轴有效性
for ax in axes:
if not (0 <p><strong>使用示例:</strong></p><pre class="brush:php;toolbar:false;"># 初始化 10×10×10 稀疏张量
st = SparseTensor((10, 10, 10))
# 批量累加(模拟频率统计)
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.add_at(coords, value)
# 沿不同轴求和(等价于 numpy.sum(axis=...))
x = st.sum(axis=(1,2)) # shape (10,) → 每个 i 的 marginal sum
y = st.sum(axis=(0,)) # shape (10,10) → 保留第1、2维
z = st.sum(axis=(0,2)) # shape (10,) → 每个 j 的 marginal sum
print("x =", list(x._data.items())) # [(0, 6), (2, 204), (5, 1)]
print("y shape:", y.shape) # (10, 10)关键优势与注意事项:
✅ 内存极致精简:仅存储非零项,空间复杂度 O(nnz),远优于全量数组。
✅ 索引 O(1):字典哈希查找,坐标访问极快。
✅ 求和可控高效:sum(axis=...) 时间复杂度为 O(nnz),且可利用 defaultdict 避免重复键查找。
⚠️ 无随机访问切片优化:如 st[:, i, :] 类操作需遍历全部非零项(O(nnz)),但高维稀疏场景下 nnz ≪ 总元素数,实际性能仍可观。若需高频单轴切片,可扩展为维护按各维度分组的索引映射(如 dim_to_keys[ax][val] = list of coords),以空间换时间。
⚠️ 不支持原地修改形状:shape 在初始化后固定,符合频率表的静态结构假设。
该方案平衡了简洁性、可维护性与性能,在统计建模、离散概率分布表示、高维直方图等场景中已被广泛验证有效。对于超大规模(nnz > 10⁶)或需并行计算的场景,可进一步结合 numba JIT 加速求和循环,或迁移至 sparse 库(如 pydata/sparse)的 COO/GCXS 格式——但后者目前对 >2D 的 API 支持仍不如自定义类灵活直观。











