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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.30 15:27:03