C++矩阵库特征值/特征向量计算API重复代码优化方案问询
代码去重方案:编译期标签调度
你当前的代码重复问题可以通过编译期标签调度+统一内部实现的方案解决,完全保留原有对外API的同时实现逻辑收敛,零运行时额外开销,非常适合你的场景。
核心思路
定义两个无状态的空标签类型,作为模板参数传给内部统一实现函数,通过C++17的if constexpr在编译期决定是否计算特征向量、返回哪种类型,不需要的分支代码会被编译器直接优化掉,和原有拆分实现的性能完全一致。
具体实现步骤
1. 定义调度标签
// 放在公共头或者内部命名空间里都可以 struct compute_eigenval_only {}; struct compute_eigenval_with_vec {};
2. 收敛Hessenberg逻辑到统一内部实现
把两份Hessenberg的重复代码合并到同一个内部实现函数,用标签控制分支:
namespace internal { template <typename Tag, typename Derived, isScalar U, isScalar T = ScalarTypeT<U>> requires ScalarTypeTo<U, T> auto HessenbergImpl(const MatrixBase<Derived, U, 2>& mat, bool both = false) { std::size_t n = mat.dims(0); std::size_t C = mat.dims(1); if (n != C) { throw std::invalid_argument("Not a square matrix, cannot transform to Hessenberg"); } // 编译期决定是否初始化变换矩阵V Mat<T> V = [&]() { if constexpr (std::is_same_v<Tag, compute_eigenval_with_vec>) { return identity<T>(n); } else { return Mat<T>{}; } }(); if (n < 3) { if constexpr (std::is_same_v<Tag, compute_eigenval_with_vec>) { return std::pair{Mat<T>(mat), V}; } else { return Mat<T>(mat); } } auto conjif = [&](const auto& v) { if constexpr (isComplex<U>) { return conj(v); } else { return v; } }; Mat<T> H = mat; if (both) { for (std::size_t k = 0; k < n - 1; ++k) { auto ck1 = H.col(k).submatrix(k + 1); if (norm(ck1) > tolerance_soft) { auto vk1 = Householder(ck1); auto Sub1 = H.submatrix({k + 1, k}); Sub1 -= outer(vk1, dot(2.0f * conjif(vk1), Sub1)); } auto ck2 = H.row(k).submatrix(k + 1); if (norm(ck2) > tolerance_soft) { auto vk2 = Householder(ck2); auto Sub2 = H.submatrix({k, k + 1}); Sub2 -= outer(dot(Sub2, 2.0f * vk2), conjif(vk2)); // 仅计算特征向量时才执行这段代码 if constexpr (std::is_same_v<Tag, compute_eigenval_with_vec>) { auto VSub2 = V.submatrix({k, k + 1}); VSub2 -= outer(dot(VSub2, 2.0f * vk2), conjif(vk2)); } } } } else { for (std::size_t k = 0; k < n - 2; ++k) { auto ck = H.col(k).submatrix(k + 1); if (norm(ck) < tolerance_soft) { continue; } auto vk = Householder(ck); auto Sub1 = H.submatrix({k + 1, k}); Sub1 -= outer(vk, dot(2.0f * conjif(vk), Sub1)); auto Sub2 = H.submatrix({0, k + 1}); Sub2 -= outer(dot(Sub2, 2.0f * vk), conjif(vk)); // 仅计算特征向量时才执行这段代码 if constexpr (std::is_same_v<Tag, compute_eigenval_with_vec>) { auto VSub2 = V.submatrix({0, k + 1}); VSub2 -= outer(dot(VSub2, 2.0f * vk), conjif(vk)); } } } if constexpr (std::is_same_v<Tag, compute_eigenval_with_vec>) { return std::pair{H, V}; } else { return H; } } } // namespace internal
3. 封装对外的Hessenberg接口
原有对外接口完全保留,只是做薄封装:
template <typename Derived, isScalar U, isScalar T = ScalarTypeT < U>> requires ScalarTypeTo<U, T> Mat<T> Hessenberg(const MatrixBase<Derived, U, 2>& mat, bool both = false) { return internal::HessenbergImpl<compute_eigenval_only>(mat, both); } template <typename Derived, isScalar U, isScalar T = ScalarTypeT<U>> requires ScalarTypeTo<U, T> std::pair<Mat<T>, Mat<T>> HessenbergWithVec(const MatrixBase<Derived, U, 2>& mat) { return internal::HessenbergImpl<compute_eigenval_with_vec>(mat, false); }
4. 同理收敛顶层特征计算逻辑
namespace internal { template <typename Tag, typename Derived, isScalar U, isScalar T = CmpTypeT<U>> requires CmpTypeTo<U, T> auto EigenImpl(const MatrixBase<Derived, U, 2>& M) { std::size_t n = M.dims(0); std::size_t C = M.dims(1); if (n != C) { throw std::invalid_argument("Not a square Matrix, cannot compute eigenvalues"); } if (n == 1) { T val = M[{0, 0}]; if constexpr (std::is_same_v<Tag, compute_eigenval_with_vec>) { return std::vector<std::pair<T, Vec<T>>>{{val, Vec<T>{T{1}}}}; } else { return std::vector<T>{val}; } } else if (n == 2) { if constexpr (std::is_same_v<Tag, compute_eigenval_with_vec>) { return eigenVecTwo(M); } else { return eigenTwo(M); } } else if (n == 3) { if constexpr (std::is_same_v<Tag, compute_eigenval_with_vec>) { return eigenVecThree(M); } else { return eigenThree(M); } } else { if constexpr (std::is_same_v<Tag, compute_eigenval_with_vec>) { auto [H, V] = HessenbergWithVec(M); return QRIterationWithVec(H, V); } else { auto H = Hessenberg(M); return QRIteration(H); } } } } // namespace internal
5. 封装对外的特征计算接口
原有API完全不变,用户无感知:
template <typename Derived, isScalar U, isScalar T = CmpTypeT<U>> requires CmpTypeTo<U, T> std::vector<T> eigenval(const MatrixBase<Derived, U, 2>& M) { return internal::EigenImpl<compute_eigenval_only>(M); } template <typename Derived, isScalar U, isScalar T = CmpTypeT<U>> requires CmpTypeTo<U, T> std::vector<std::pair<T, Vec<T>>> eigenvec(const MatrixBase<Derived, U, 2>& M) { return internal::EigenImpl<compute_eigenval_with_vec>(M); }
方案优势
- 对外API完全保持原有形态,用户不需要修改任何调用代码,不会出现optional方案的体验问题
- 所有重复逻辑收敛到内部实现,后续改bug、加特性只需要修改一份代码
- 零运行时开销:所有标签判断都在编译期完成,不需要的分支会被编译器直接优化掉,性能和原有拆分实现完全一致
- 扩展性好,后续如果要新增比如仅计算前K个特征值的接口,只需要新增标签即可,不需要重复复制逻辑
注:原有Hessenberg实现的upper分支里有笔误,把vk写成了vk1,合并实现的时候可以顺便修正。
内容的提问来源于stack exchange,提问作者frozenca
相关产品推荐
相关产品推荐

