PyTorch创建稀疏张量时维度不匹配RuntimeError问题求助
错误根源分析
报错RuntimeError: number of dimensions must be sparse_dim (59697) + dense_dim (0), but got 3的核心是创建稀疏张量时,索引的稀疏维度数与指定的张量形状维度数不匹配:
索引张量维度异常
创建torch.sparse.FloatTensor时,第一个参数indices的最后一维长度即为sparse_dim。报错显示sparse_dim=59697,说明传入的X_I.t()是形状为(59697, num_nonzero)的张量——完全不符合参数注释中X_I应为(num_nonzero, 2)的COO索引格式(转置后应为(2, num_nonzero),对应行、列两个稀疏维度)。代码逻辑与注释矛盾
函数注释明确X_shape是稀疏矩阵的2维形状,但代码中却执行B, N, _ = X_shape,强行解包3个元素,说明实际传入的X_shape是3维,而代码逻辑未做对应调整,导致后续计算B*N后生成的稀疏张量形状完全错误。稀疏张量构造逻辑错误
代码中直接将批量稀疏矩阵展平为(B*N, Y.shape[1])的2维张量,违背了批量稀疏矩阵乘法的维度逻辑——批量稀疏矩阵应保持(batch_size, rows, cols)的3维结构,而非展平。
具体解决步骤
1. 修正输入参数的维度匹配
- 确保传入的
X_I是COO格式的索引:批量场景下为(num_nonzero, 3)(每行对应(batch_idx, row_idx, col_idx)),单矩阵场景为(num_nonzero, 2)。 - 确保
X_shape与X_I匹配:批量场景下为(B, N, M)(B=批量大小,N=矩阵行数,M=矩阵列数),单矩阵场景为(N, M)。
2. 修复sparse_bmm函数逻辑
替换原函数代码,适配批量/单矩阵两种场景的正确逻辑:
def sparse_bmm(X_I, X_V, X_shape, Y): """ 执行批量/单矩阵稀疏矩阵乘法(X * Y)。 :param X_I: COO格式稀疏矩阵的索引(num_nonzero, 3)→ 批量场景;或(num_nonzero,2)→单矩阵场景 :param X_V: 稀疏矩阵的值(num_nonzero,) :param X_shape: 稀疏矩阵的形状(3,)→批量场景;或(2,)→单矩阵场景 :param Y: 稠密矩阵(B, M, K)→批量场景;或(M,K)→单矩阵场景 :return: 乘法结果(B*N, K)→展平格式;或(N,K)→单矩阵格式 """ if len(X_shape) == 3: B, N, M = X_shape # 转置索引为torch.sparse.FloatTensor要求的格式:(sparse_dim, num_nonzero) X_I = X_I.t() # 创建3维稀疏张量 X_sparse = torch.sparse.FloatTensor(X_I, X_V, torch.Size(X_shape)) # 转为稠密后执行批量矩阵乘法 X_dense = X_sparse.to_dense() result = torch.bmm(X_dense, Y) # 展平结果返回 return result.view(B * N, -1) else: N, M = X_shape X_I = X_I.t() X_sparse = torch.sparse.FloatTensor(X_I, X_V, torch.Size(X_shape)) X_dense = X_sparse.to_dense() result = torch.matmul(X_dense, Y) return result
3. 调整输入Y的维度
- 批量场景下,确保
Y的形状为(B, M, K)(与X的列数M匹配),而非原代码中的(B*N, K)。如果Y是展平格式,先通过Y.view(B, M, K)调整维度。
4. 验证模型加载后的参数维度
加载model/repro/wikimovie/best_model_doc后,检查传入sparse_bmm的X_I、X_shape、Y的实际形状,确认与函数要求一致。
内容的提问来源于stack exchange,提问作者Ihsan Ullah Khan

