如何推导带initializer_list构造的编译期Matrix类的维度与元素类型
C++编译期矩阵类的模板参数推导问题
原实现代码
#include <iostream> #include <array> template <typename T, size_t Rows, size_t Cols> class Matrix { protected: std::array<std::array<T, Cols>, Rows> matrix; public: constexpr Matrix() = default; constexpr explicit Matrix(std::array<std::array<T, Cols>, Rows> matrix) : matrix(std::move(matrix)) { } consteval Matrix(std::initializer_list<std::initializer_list<T>> matrix) : matrix() { if (matrix.size() != Rows) { throw std::invalid_argument("Invalid matrix Rows count"); } auto current_row = matrix.begin(); for (size_t i = 0; i < matrix.size(); i++, current_row++) { if (current_row->size() != Cols) { throw std::invalid_argument("Invalid matrix column count"); } std::copy(current_row->begin(), current_row->end(), this->matrix[i].begin()); } } constexpr auto& operator[](this auto&& self, size_t index) { return self.matrix[index]; } }; template <typename T, size_t N> struct Vector : Matrix<T, 1, N> { using Matrix<T, 1, N>::Matrix; template <typename U, typename... Us> consteval explicit Vector(U u, Us... us) { this->matrix[0] = { u, us... }; } constexpr auto& operator[](this auto&& self, size_t index) { return self.matrix[0][index]; } }; template <typename T, typename... U> Vector(T, U...) -> Vector<T, 1 + sizeof...(U)>; int main() { constexpr Matrix<int, 4, 3> matrix { { 1, 2, 3 }, { 1, 2, 3 }, { 1, 2, 3 }, { 1, 2, 3 } }; static_assert(matrix[2][1] == 2); constexpr Vector v { 1, 2, 3 }; static_assert(v[1] == 2); }
问题背景
上述代码实现了编译期可用的Matrix和Vector类,其中Vector借助推导指南可以省略类型和元素数量的指定,但Matrix的模板参数推导存在障碍——std::initializer_list的size()不是常量表达式,无法用来推导Rows和Cols参数。
待解决的两个问题
问题1
如何编写如下形式的初始化表达式,让编译器自动推导Matrix的元素类型与行列维度?
constexpr Matrix matrix { { 1, 2, 3 }, { 1, 2, 3 }, { 1, 2, 3 }, { 1, 2, 3 } };
问题2
如何编写如下形式的初始化表达式,让编译器自动推导Matrix的行列维度,并将初始元素转换为指定的元素类型?
constexpr Matrix<double> matrix { { 1, 2, 3 }, { 1, 2, 3 }, { 1, 2, 3 }, { 1, 2, 3 } };
解决方案
核心思路是放弃依赖std::initializer_list,改用编译期可推导的std::array作为行载体,结合模板推导指南和辅助构造函数实现参数推导。
完整改进代码
#include <iostream> #include <array> #include <type_traits> template <typename T, size_t Rows, size_t Cols> class Matrix { protected: std::array<std::array<T, Cols>, Rows> matrix; public: constexpr Matrix() = default; constexpr explicit Matrix(std::array<std::array<T, Cols>, Rows> matrix) : matrix(std::move(matrix)) { } consteval Matrix(std::initializer_list<std::initializer_list<T>> matrix) : matrix() { if (matrix.size() != Rows) { throw std::invalid_argument("Invalid matrix Rows count"); } auto current_row = matrix.begin(); for (size_t i = 0; i < matrix.size(); i++, current_row++) { if (current_row->size() != Cols) { throw std::invalid_argument("Invalid matrix column count"); } std::copy(current_row->begin(), current_row->end(), this->matrix[i].begin()); } } // 新增:支持类型转换的行构造函数 template <typename U, size_t... RowSizes> consteval Matrix(std::array<U, RowSizes>... rows) requires ((RowSizes == Cols) && ...) { static_assert(sizeof...(rows) == Rows, "Row count mismatch with template parameter"); static_assert((std::is_convertible_v<U, T> && ...), "Element type cannot be converted to Matrix's value type"); size_t row_idx = 0; ((std::copy(rows.begin(), rows.end(), matrix[row_idx++].begin())), ...); } constexpr auto& operator[](this auto&& self, size_t index) { return self.matrix[index]; } }; // 辅助结构体:推导Matrix的模板参数 template <typename FirstRow, typename... RestRows> struct MatrixDeduction { using ValueType = typename FirstRow::value_type; static constexpr size_t TotalRows = 1 + sizeof...(RestRows); static constexpr size_t TotalCols = FirstRow::size(); static_assert((std::is_same_v<FirstRow, RestRows> && ...), "All rows must have identical type"); }; // 问题1的推导指南:自动推导所有参数 template <typename FirstRow, typename... RestRows> Matrix(FirstRow, RestRows...) -> Matrix< typename MatrixDeduction<FirstRow, RestRows...>::ValueType, MatrixDeduction<FirstRow, RestRows...>::TotalRows, MatrixDeduction<FirstRow, RestRows...>::TotalCols >; // 问题2的辅助构造函数:指定类型推导行列 template <typename T, typename... Rows> consteval auto make_matrix(Rows&&... rows) { using RowType = std::decay_t<decltype(rows)...>; return Matrix<T, 1 + sizeof...(Rows), RowType::size()>{std::forward<Rows>(rows)...}; } template <typename T, size_t N> struct Vector : Matrix<T, 1, N> { using Matrix<T, 1, N>::Matrix; template <typename U, typename... Us> consteval explicit Vector(U u, Us... us) { this->matrix[0] = { u, us... }; } constexpr auto& operator[](this auto&& self, size_t index) { return self.matrix[0][index]; } }; template <typename T, typename... U> Vector(T, U...) -> Vector<T, 1 + sizeof...(U)>; int main() { // 问题1:自动推导类型与行列 constexpr Matrix matrix1 { std::array{1, 2, 3}, std::array{1, 2, 3}, std::array{1, 2, 3}, std::array{1, 2, 3} }; static_assert(std::is_same_v<decltype(matrix1), Matrix<int, 4, 3>>); static_assert(matrix1[2][1] == 2); // 问题2:指定类型,自动推导行列 constexpr Matrix matrix2 = make_matrix<double>( std::array{1, 2, 3}, std::array{1, 2, 3}, std::array{1, 2, 3}, std::array{1, 2, 3} ); static_assert(std::is_same_v<decltype(matrix2), Matrix<double, 4, 3>>); static_assert(matrix2[2][1] == 2.0); constexpr Vector v { 1, 2, 3 }; static_assert(v[1] == 2); }
方案说明
问题1解决方式:
- 借助
std::array作为行的容器,利用其编译期可知的size()属性推导行列数 - 通过辅助结构体
MatrixDeduction统一提取元素类型、行数和列数,结合推导指南实现自动推导
- 借助
问题2解决方式:
- 提供
make_matrix辅助函数,接收指定的元素类型和std::array行数据 - 函数内部自动推导行列数,并构造对应类型的
Matrix对象,同时支持元素类型的隐式转换
- 提供
内容的提问来源于stack exchange,提问作者Osama Ahmad
相关产品推荐
相关产品推荐

