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

C++可变参数模板如何传递修改参数实现通用梯度下降求解器

可constexpr的通用可变参数梯度下降实现

完全可以实现,在C++17及以上标准中仅需借助标准库的tuple和索引序列工具就能完成,不需要复杂的递归模板逻辑;如果接受所有传入参数为double的前提,实现还可以更简洁,且全程保留constexpr特性。

实现注意要点

  • 不要使用原有代码里的std::function做类型擦除:std::function的调用运算符不支持constexpr,会直接破坏编译期运行能力,直接通过模板参数接收任意可调用对象(普通函数、constexpr lambda、仿函数均可)即可,无额外运行时开销。
  • 可变参数的遍历可以通过std::index_sequence配合折叠表达式完成,不需要手写递归展开:将所有传入的参数引用打包为tuple,按索引逐个修改对应参数加epsilon计算偏导,再更新参数值即可,逻辑和原有的固定3参数版本完全一致。
  • 低版本C标准中std::copysign没有constexpr修饰,自己实现一个极简的constexpr版本即可满足需求,C26标准中std::copysign已支持constexpr,可以直接使用标准库实现。

完整实现代码

#include <cmath>
#include <tuple>
#include <utility>
#include <type_traits>

// 兼容C++20及更早版本的constexpr符号复制函数
constexpr double constexpr_copysign(double magnitude, double sign) noexcept {
    return sign >= 0 ? std::abs(magnitude) : -std::abs(magnitude);
}

template<typename ResidualFunc, std::size_t... Indices, typename... Args>
constexpr double gd_impl(const ResidualFunc& res_fun, double step_size, double epsilon,
                         std::index_sequence<Indices...>, Args&... args) {
    auto args_ref = std::tie(args...);
    const double residual = res_fun(*std::get<Indices>(args_ref)...);

    // 折叠表达式遍历每个参数,计算偏导并更新
    auto update_single_arg = [&]<std::size_t I>(std::integral_constant<std::size_t, I>) {
        auto modified_args = args_ref;
        std::get<I>(modified_args) += epsilon;
        const double partial_deriv = (std::apply(res_fun, modified_args) - residual) / epsilon;
        std::get<I>(args_ref) -= constexpr_copysign(step_size, partial_deriv);
    };
    (update_single_arg(std::integral_constant<std::size_t, Indices>{}), ...);

    return residual;
}

template<typename ResidualFunc, typename... Args>
constexpr double gradient_descent(const ResidualFunc& res_fun, double step_size, Args&... args, double epsilon = 0.01) {
    // 校验所有参数均为double,符合提出的可接受前提
    static_assert((std::is_same_v<Args, double> && ...), "All optimization arguments must be double type");
    return gd_impl(res_fun, step_size, epsilon, std::index_sequence_for<Args...>{}, args...);
}

使用示例

// 测试目标函数:f(x,y,z) = (x-1)^2 + (y-2)^2 + (z-3)^2,全局最小值0在(1,2,3)处取得
constexpr auto test_residual = [](double x, double y, double z) constexpr {
    const double dx = x - 1.0;
    const double dy = y - 2.0;
    const double dz = z - 3.0;
    return dx*dx + dy*dy + dz*dz;
};

// 编译期运行验证
constexpr bool compile_time_check() {
    double x = 0.0, y = 0.0, z = 0.0;
    double last_res = 0.0;
    // 迭代100次
    for (int i = 0; i < 100; ++i) {
        last_res = gradient_descent(test_residual, 0.01, x, y, z);
    }
    // 验证结果收敛到足够接近最优值
    return std::abs(x - 1.0) < 0.1 && std::abs(y - 2.0) < 0.1 && std::abs(z - 3.0) < 0.1 && last_res < 0.1;
}
// 静态断言通过即证明constexpr逻辑正常
static_assert(compile_time_check(), "Constexpr gradient descent runs correctly at compile time");

扩展说明

  • 该实现完全保留了原有代码的算法逻辑:固定步长、数值法计算偏导、按导数符号更新参数,没有修改算法本身的行为。
  • 如果后续需要支持非double类型的优化参数,只需要删除函数里的static_assert校验,将相关的double类型替换为推导的泛型类型即可,核心遍历逻辑不需要改动。
  • 如果不需要constexpr支持,直接移除函数前的constexpr修饰,把自定义的constexpr_copysign替换为标准库std::copysign即可正常使用。

内容的提问来源于stack exchange,提问作者Roel Baardman

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.27 21:57:22