You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

能否利用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:

  1. 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.
  2. Column Expansion: The outer ... after the sum expands this dot product calculation over all column indices (ColIndices), building a complete row as a std::array.
  3. 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::array operations.
  • The code is fully constexpr, meaning matrix multiplication can be computed at compile time for constant matrices.
  • The static_assert ensures 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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.26 10:12:23