PyTorch C++中是否支持两个稀疏矩阵相乘的函数?
PyTorch C++ 中稀疏矩阵相乘的实现方法
在PyTorch C++ API中,存在对应Python版torch.sparse.mm的函数——torch::sparse::mm,它支持两个torch::SparseTensor类型的稀疏矩阵直接相乘。
你之前使用的mat1.mm(mat2)是稠密矩阵的乘法方法,仅支持第二个矩阵为稠密张量的场景,稀疏矩阵间的乘法需要调用专门的稀疏模块函数。
示例代码
#include <torch/torch.h> int main() { // 创建第一个稀疏矩阵(3x3) auto indices1 = torch::tensor({{0, 0, 1, 2}, {0, 2, 1, 0}}, torch::kInt64); auto values1 = torch::tensor({1.0, 2.0, 3.0, 4.0}); auto sparse_mat1 = torch::sparse_coo_tensor(indices1, values1, {3, 3}); // 创建第二个稀疏矩阵(3x2) auto indices2 = torch::tensor({{0, 1, 2, 2}, {1, 2, 0, 1}}, torch::kInt64); auto values2 = torch::tensor({5.0, 6.0, 7.0, 8.0}); auto sparse_mat2 = torch::sparse_coo_tensor(indices2, values2, {3, 2}); // 执行稀疏矩阵相乘 auto result = torch::sparse::mm(sparse_mat1, sparse_mat2); // 转为稠密张量查看结果 std::cout << result.to_dense() << std::endl; return 0; }
注意:参与运算的两个稀疏矩阵必须满足矩阵乘法的维度规则(第一个矩阵的列数等于第二个矩阵的行数),否则会触发维度不匹配的错误。
内容的提问来源于stack exchange,提问作者patrick
相关产品推荐
相关产品推荐

