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

如何用C++实现通用版梯度下降算法以通过指定GTest用例

问题:实现兼容多种函数类型的C++通用梯度下降算法

我需要用C++实现一个通用的gradient_descent算法,使其能通过以下GTest测试用例,支持传入函数指针、lambda函数或仿函数作为目标函数。

测试用例代码

#include <cmath>

TEST(HW6Test, TEST1) {
    auto min1 = q1::gradient_descent(0.01, 0.1, cos);
    EXPECT_NEAR(min1, 3.14, 0.1);

    auto min2 = q1::gradient_descent(0.01, 0.01, cos);
    EXPECT_NEAR(min2, 3.14, 0.01);
}

TEST(HW6Test, TEST2) {
    auto min = q1::gradient_descent(0.01, 0.01, [](double a){return sin(a)+cos(a);});
    EXPECT_NEAR(min, -2.36, 0.01);
}

TEST(HW6Test, TEST3) {
    struct Func
    {
        double operator()(double a) {return cos(a);}
    };
    auto min = q1::gradient_descent(0.01, 0.01, Func{});
    EXPECT_NEAR(min, 3.14, 0.01);
}

TEST(HW6Test, TEST4) {
    struct Func
    {
        double operator()(double a) {return sin(a);}
    };
    auto min = q1::gradient_descent<double, Func>(0.0, 0.01);
    EXPECT_NEAR(min, -1.57, 0.01);
}

当前尝试的实现

我编写了如下模板函数,但无法兼容TEST1中传入cos的调用,以及TEST4中显式指定模板参数但不传入函数实例的调用:

namespace q1 
{
    template <typename T, typename Func>
    const T& gradient_descent(const T& init_value, const T& step, Func func) {

    }
};

问题分析与解决方案

1. 修正返回值类型

当前代码返回const T&,但传入的init_value常为临时对象(如0.01),返回其引用会导致悬空引用,触发未定义行为。需将返回值改为值类型T:

template <typename T, typename Func>
T gradient_descent(const T& init_value, const T& step, Func func) {
    // 核心逻辑实现
}

2. 支持无函数实例的调用(针对TEST4)

TEST4中显式指定了Func类型但未传入实例,因此需要提供一个重载版本,利用Func的默认构造函数创建实例:

template <typename T, typename Func>
T gradient_descent(const T& init_value, const T& step) {
    return gradient_descent(init_value, step, Func{});
}

这个重载会调用原有的三参数版本,自动构造Func对象,适配TEST4的调用方式。

3. 完善梯度下降核心逻辑

梯度下降需要通过数值微分计算目标函数的导数(梯度),这里可以用前向差分近似:

template <typename T, typename Func>
T gradient_descent(const T& init_value, const T& step, Func func) {
    T x = init_value;
    const T eps = static_cast<T>(1e-8); // 数值微分的微小增量
    const int max_iter = 100000; // 最大迭代次数,防止死循环
    int iter = 0;

    while (iter < max_iter) {
        // 前向差分计算导数:f'(x) ≈ [f(x+eps) - f(x)] / eps
        T derivative = (func(x + eps) - func(x)) / eps;
        T next_x = x - step * derivative;

        // 收敛条件:两次迭代的x变化小于阈值
        if (std::abs(next_x - x) < static_cast<T>(1e-6)) {
            break;
        }

        x = next_x;
        iter++;
    }

    return x;
}

注:如果需要更精确的导数,可以改用中心差分(func(x+eps) - func(x-eps))/(2*eps),但计算量会翻倍。

4. 适配函数指针的调用(针对TEST1)

cos是标准库中的函数指针(类型为double(*)(double)),上述模板可以自动推导Func为该函数指针类型,因此修正返回值后即可兼容TEST1的调用。

完整实现代码

#include <cmath>
#include <algorithm>

namespace q1 
{
    template <typename T, typename Func>
    T gradient_descent(const T& init_value, const T& step, Func func) {
        T x = init_value;
        const T eps = static_cast<T>(1e-8);
        const int max_iter = 100000;
        int iter = 0;

        while (iter < max_iter) {
            T derivative = (func(x + eps) - func(x)) / eps;
            T next_x = x - step * derivative;

            if (std::abs(next_x - x) < static_cast<T>(1e-6)) {
                break;
            }

            x = next_x;
            iter++;
        }

        return x;
    }

    template <typename T, typename Func>
    T gradient_descent(const T& init_value, const T& step) {
        return gradient_descent(init_value, step, Func{});
    }
};

内容的提问来源于stack exchange,提问作者Kato Megumi

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.04 19:44:56