能否利用Nested parameter pack expansion将矩阵乘法代码简化为嵌套折叠形式?
Answer
Absolutely! You can simplify the matrix multiplication to use nested fold expressions in the form you described, leveraging C++17's fold expressions and compile-time index sequences to generate the required rows, columns, and inner sum indices. Here's how to implement it:
First, let's complete your existing code with the necessary template helpers and the main multiplication function:
#include <array> #include <utility> template<typename T, size_t R, size_t C> using Matrix = std::array<std::array<T, C>, R>; template<typename A, typename B> using mul_el_t = decltype(std::declval<A>()[0][0] * std::declval<B>()[0][0]); // Helper function that performs the nested fold expansion template<typename T, size_t R, size_t K, size_t C, size_t... RowIndices, size_t... ColIndices, size_t... InnerIndices> constexpr auto multiply_impl(const Matrix<T, R, K>& a, const Matrix<T, K, C>& b, std::index_sequence<RowIndices...>, std::index_sequence<ColIndices...>, std::index_sequence<InnerIndices...>) { using ResultElement = mul_el_t<decltype(a), decltype(b)>; return Matrix<ResultElement, R, C>{ // For each row in the result, create a row array std::array<ResultElement, C>{ // For each column in the row, compute the dot product sum ((a[RowIndices][InnerIndices] * b[InnerIndices][ColIndices]) + ...) ... // Expand over all column indices } ... // Expand over all row indices }; } // Main matrix multiplication function template<typename T, size_t R, size_t K, size_t C> constexpr auto multiply(const Matrix<T, R, K>& a, const Matrix<T, K, C>& b) { static_assert(K > 0, "Matrix multiplication requires non-zero inner dimension"); return multiply_impl(a, b, std::make_index_sequence<R>{}, std::make_index_sequence<C>{}, std::make_index_sequence<K>{}); } // Example usage int main() { constexpr Matrix<int, 2, 2> mat1 = {{ {1, 2}, {3, 4} }}; constexpr Matrix<int, 2, 2> mat2 = {{ {5, 6}, {7, 8} }}; constexpr auto result = multiply(mat1, mat2); // Result should be: // Row 0: 1*5 + 2*7 = 19, 1*6 +2*8=22 // Row1:3*5+4*7=43,3*6+4*8=50 static_assert(result[0][0] == 19); static_assert(result[0][1] == 22); static_assert(result[1][0] == 43); static_assert(result[1][1] == 50); return 0; }
How It Works:
- Inner Sum Fold:
((a[RowIndices][InnerIndices] * b[InnerIndices][ColIndices]) + ...)computes the dot product for a single row-column pair by folding over all inner dimension indices (InnerIndices). This is exactly the(a[r][k] * b[k][c]) + ...part you wanted. - Column Expansion: The outer
...after the sum expands this dot product calculation over all column indices (ColIndices), building a complete row as astd::array. - Row Expansion: The final
...expands the row-building logic over all row indices (RowIndices), constructing the full result matrix.
Key Notes:
- This requires C++17 or later since it uses fold expressions and constexpr
std::arrayoperations. - The code is fully constexpr, meaning matrix multiplication can be computed at compile time for constant matrices.
- The
static_assertensures we don't attempt multiplication with an invalid inner dimension (K=0).
This implementation directly maps to the nested fold structure you described, making the code concise and expressive while leveraging compile-time computation for efficiency.
内容的提问来源于stack exchange,提问作者Dmytro Starosud
相关产品推荐
相关产品推荐

