如何从PyTorch稀疏张量中提取索引张量
获取PyTorch稀疏张量的索引张量
当你把稠密张量(比如图的邻接矩阵)通过to_sparse()转换成COO格式的稀疏张量后,想要提取其中存储非零元素坐标的索引张量(形状为(2, nnz)),直接用A_sparse[0]这类稠密张量的索引方式是行不通的——因为稀疏张量的存储结构和稠密张量完全不同。
正确的方法是直接访问稀疏张量的indices属性,或者调用indices()方法:
# 假设已得到稀疏张量A_sparse indices_tensor = A_sparse.indices # 或者使用方法调用 indices_tensor = A_sparse.indices()
执行后可以验证形状:
print(indices_tensor.shape) # 输出: torch.Size([2, 10556])
PyTorch的COO稀疏张量会将所有非零元素的行、列坐标,分别存在indices这个二维张量的第一行和第二行,直接访问它就能得到你需要的索引信息。
内容的提问来源于stack exchange,提问作者kiyopi
相关产品推荐
相关产品推荐

