CRTP接口类下适配任意返回类型的方法记忆化通用实现问询
解决方案
完全可以通过单个通用模板函数实现所有返回类型、所有输入类型的记忆化逻辑,你只需要将重复的缓存判断、更新逻辑抽成泛型模板即可,具体实现如下:
第一步:优化相等判断函数,兼容Eigen稠密/稀疏输入
首先将is_memoized重载为支持任意Eigen矩阵/向量类型的模板函数,避免类型转换开销:
#include <type_traits> #include <limits> // 稠密矩阵/向量相等判断 template <typename T, typename U> std::enable_if_t< std::is_base_of_v<Eigen::MatrixBase<T>, T> && std::is_base_of_v<Eigen::MatrixBase<U>, U>, bool> is_memoized(const Eigen::MatrixBase<T>& x, const Eigen::MatrixBase<U>& x_mem) { return ((x - x_mem).array().abs() <= std::numeric_limits<double>::epsilon()).all(); } // 稀疏矩阵/向量相等判断 template <typename T, typename U> std::enable_if_t< std::is_base_of_v<Eigen::SparseMatrixBase<T>, T> && std::is_base_of_v<Eigen::SparseMatrixBase<U>, U>, bool> is_memoized(const Eigen::SparseMatrixBase<T>& x, const Eigen::SparseMatrixBase<U>& x_mem) { if (x.nonZeros() != x_mem.nonZeros()) return false; return ((x - x_mem).cwiseAbs() <= std::numeric_limits<double>::epsilon()).all(); }
第二步:实现通用记忆化模板函数
在objective_memoized类内添加通用的memoize成员模板,收拢所有缓存逻辑:
template <typename X, typename Data, typename Func> decltype(auto) memoize(const X& x, memData<Data>& cache, Func&& eval_func) { if (!is_memoized(x, cache.x)) { cache.data = std::forward<Func>(eval_func)(); cache.x = x; } return cache.data; }
该函数会自动适配任意返回值类型(scalar、Vec、Mat均可),无需针对不同返回类型做特化。
第三步:简化原有业务方法
原来的6个重复方法可以全部简化为统一的调用形式,完全消除冗余:
// 稠密输入重载 template <typename T> inline scalar operator()(const MatrixBase<T>& x) { return memoize(x, mem_operator, [&](){ return static_cast<Derived*>(this)->operator()(x); }); } template <typename T> inline Vec Eval_grad(const MatrixBase<T>& x) { return memoize(x, mem_Eval_grad, [&](){ return static_cast<Derived*>(this)->Eval_grad(x); }); } template <typename T> inline Mat Eval_hessian(const MatrixBase<T>& x) { return memoize(x, mem_Eval_hessian, [&](){ return static_cast<Derived*>(this)->Eval_hessian(x); }); } // 稀疏输入重载 template <typename T> inline scalar operator()(const SparseMatrixBase<T>& x) { return memoize(x, mem_operator, [&](){ return static_cast<Derived*>(this)->operator()(x); }); } template <typename T> inline Vec Eval_grad(const SparseMatrixBase<T>& x) { return memoize(x, mem_Eval_grad, [&](){ return static_cast<Derived*>(this)->Eval_grad(x); }); } template <typename T> inline Mat Eval_hessian(const SparseMatrixBase<T>& x) { return memoize(x, mem_Eval_hessian, [&](){ return static_cast<Derived*>(this)->Eval_hessian(x); }); }
扩展说明
如果你的场景需要支持稀疏类型的输入缓存,只需要将memData改为双模板参数即可,适配不同的输入存储类型:
template <class XType, class DataType> struct memData { XType x; DataType data; template <typename... Args> memData(Args&&... args) : x(std::forward<Args>(args)...) {} };
内容的提问来源于stack exchange,提问作者Jonas Hjulstad
相关产品推荐
相关产品推荐

