如何实现std::mdspan的矩阵乘法?Eigen仅支持拥有式MatrixXd
用Eigen实现std::mdspan的矩阵乘法
Eigen并非只支持拥有式的矩阵类型,通过Eigen::Map可以直接绑定std::mdspan管理的外部内存,实现零拷贝的矩阵乘法运算,完全适配你想要的mat_mult函数形式。
核心思路
- 用
Eigen::Map将std::mdspan的底层数据映射为Eigen的矩阵视图,无需复制数据。 - 执行矩阵乘法后,将结果写入拥有式的
mdvector容器返回。
静态维度版本(适配编译期确定的矩阵大小)
假设mdvector是C++23的std::mdarray(拥有式多维数组):
#include <mdspan> #include <Eigen/Dense> #include <stdexcept> // 定义拥有式多维向量类型(静态维度) template <typename T, size_t... Extents> using mdvector = std::mdarray<T, std::extents<Extents...>>; template <typename T, typename ExtentsA, typename ExtentsB> auto mat_mult(std::mdspan<T, ExtentsA> a, std::mdspan<T, ExtentsB> b) -> mdvector<T, ExtentsA::extent(0), ExtentsB::extent(1)> { // 检查矩阵乘法维度合法性 if (a.extent(1) != b.extent(0)) { throw std::invalid_argument("矩阵维度不匹配,无法执行乘法"); } const auto result_rows = a.extent(0); const auto result_cols = b.extent(1); // 预分配结果容器 mdvector<T, result_rows, result_cols> md_result; // 用Eigen::Map包装输入输出内存,避免拷贝 Eigen::Map<const Eigen::Matrix<T, Eigen::Dynamic, Eigen::Dynamic>> eigen_a(a.data(), result_rows, a.extent(1)); Eigen::Map<const Eigen::Matrix<T, Eigen::Dynamic, Eigen::Dynamic>> eigen_b(b.data(), b.extent(0), result_cols); Eigen::Map<Eigen::Matrix<T, Eigen::Dynamic, Eigen::Dynamic>> eigen_result(md_result.data(), result_rows, result_cols); // 直接将乘法结果写入md_result内存 eigen_result = eigen_a * eigen_b; return md_result; }
动态维度版本(适配运行时确定的矩阵大小)
如果需要支持运行时动态变化的矩阵维度,可将mdvector定义为基于std::vector的自定义类型:
#include <mdspan> #include <Eigen/Dense> #include <stdexcept> #include <vector> #include <array> // 自定义动态维度拥有式向量类型 template <typename T> struct mdvector { std::vector<T> data; std::array<size_t, 2> dims; mdvector(size_t rows, size_t cols) : data(rows * cols), dims{rows, cols} {} // 可选:提供mdspan视图接口 std::mdspan<T, std::extents<std::dynamic_extent, std::dynamic_extent>> as_mdspan() { return std::mdspan<T, std::extents<std::dynamic_extent, std::dynamic_extent>>(data.data(), dims[0], dims[1]); } }; template <typename T, typename ExtentsA, typename ExtentsB> auto mat_mult(std::mdspan<T, ExtentsA> a, std::mdspan<T, ExtentsB> b) -> mdvector<T> { if (a.extent(1) != b.extent(0)) { throw std::invalid_argument("矩阵维度不匹配,无法执行乘法"); } const size_t result_rows = a.extent(0); const size_t result_cols = b.extent(1); mdvector<T> md_result(result_rows, result_cols); Eigen::Map<const Eigen::Matrix<T, Eigen::Dynamic, Eigen::Dynamic>> eigen_a(a.data(), result_rows, a.extent(1)); Eigen::Map<const Eigen::Matrix<T, Eigen::Dynamic, Eigen::Dynamic>> eigen_b(b.data(), b.extent(0), result_cols); Eigen::Map<Eigen::Matrix<T, Eigen::Dynamic, Eigen::Dynamic>> eigen_result(md_result.data.data(), result_rows, result_cols); eigen_result = eigen_a * eigen_b; return md_result; }
关键说明
Eigen::Map是Eigen的视图类,直接绑定外部内存地址,完全兼容std::mdspan的非拥有式特性,无需复制输入数据。- 上述实现避免了中间矩阵拷贝,直接将乘法结果写入目标
mdvector的内存,效率与原生Eigen矩阵乘法一致。 - 可根据实际需求调整
mdvector的定义,只要它是拥有式的多维容器、能提供连续内存的访问接口即可。
内容的提问来源于stack exchange,提问作者Tom Huntington
相关产品推荐
相关产品推荐

