Python中基于字典的稀疏矩阵乘法实现问题求助
稀疏矩阵乘法问题修正
原代码的核心错误
你的代码存在两个关键问题:
- 乘法逻辑完全错误:仅计算了结果矩阵的对角线元素,且元素配对方式不符合矩阵乘法的点积规则。矩阵乘法
mat2 × mat1中,结果矩阵的第x行第y列元素,是mat2第x行与mat1第y列的逐元素乘积之和,而非仅针对(i,i)位置的计算。 - 坐标对应关系混乱:混淆了两个矩阵的坐标与行列的对应关系,导致相乘的元素并非矩阵乘法中需要配对的元素。
正确实现思路
根据题目定义:坐标(j,i)代表第i行第j列,矩阵乘法C = mat2 × mat1的元素计算规则为:
- 结果矩阵
C的第x行第y列元素(对应坐标(y, x)),等于sum_{k=1到n} mat2第x行第k列元素 × mat1第k行第y列元素 - 对应到稀疏矩阵字典的键:
- mat2第x行第k列的键是
(k, x) - mat1第k行第y列的键是
(y, k)
- mat2第x行第k列的键是
修正后的代码
def sparse_mult(n, mat2, mat1): result = {} # 预处理mat2:按行分组,记录每行的非零列与对应值 mat2_rows = {} for (k, x), val2 in mat2.items(): if x not in mat2_rows: mat2_rows[x] = {} mat2_rows[x][k] = val2 # 预处理mat1:按列分组,记录每列的非零行与对应值 mat1_cols = {} for (y, k), val1 in mat1.items(): if y not in mat1_cols: mat1_cols[y] = {} mat1_cols[y][k] = val1 # 计算所有非零乘积元素 for x in mat2_rows: for y in mat1_cols: total = 0 # 取mat2行的列与mat1列的行的交集,仅计算有效点积 common_k = set(mat2_rows[x].keys()) & set(mat1_cols[y].keys()) for k in common_k: total += mat2_rows[x][k] * mat1_cols[y][k] if total != 0: result[(y, x)] = total return result
代码优化说明
- 预处理分组避免了全量遍历n次,适配稀疏矩阵的特性,提升计算效率
- 仅在存在共同非零元素的行和列间计算点积,减少无效运算
- 严格遵循题目中的坐标规则,确保结果矩阵的键格式正确
内容的提问来源于stack exchange,提问作者YYY
相关产品推荐
相关产品推荐

