如何用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
相关产品推荐
相关产品推荐

