使用std::max实现AutoDiff版ReLU时导数全为0的问题排查
用std::max实现ReLU时AutoDiff导数全为0的原因与解决办法
核心原因
标准库的std::max并没有为自定义的自动微分类型(比如你这里的dual)预设微分规则。大多数基于运算符重载的AutoDiff框架只会跟踪自定义类型的重载运算符/函数,而std::max默认是针对原生数值类型(int、double等)实现的:当你传入dual类型时,要么框架无法识别这个操作的微分逻辑,直接将导数置为0;要么因为类型推导问题,实际调用的是原生类型的max,导致dual携带的微分信息丢失,最终导数全为0。
而用条件语句实现的ReLU_ol和ReLU_if能正常工作,是因为你直接操作dual类型的数值部分,并且返回的是dual对象——框架能跟踪到分支逻辑:当x.value() > 0时返回x(导数为1),否则返回0(导数为0),完全符合ReLU的微分规则。
解决办法
1. 实现适配dual类型的自定义max函数
不要直接用std::max,自己写一个支持dual的max函数,内置微分逻辑:
// 适配dual与原生数值类型的max dual max_dual(const dual& a, double b) { if (a.value() > b) { return a; // 返回a时,对输入的导数就是a自身的导数(此处输入为y,导数为1) } else { return dual(b, 0.0); // 返回常量时,对输入y的导数为0 } } // 用该函数实现ReLU_max dual ReLU_max(const dual& y) { return max_dual(y, 0.0); }
如果需要支持两个dual类型的比较,也可以重载:
dual max_dual(const dual& a, const dual& b) { if (a.value() > b.value()) { return a; } else { return b; } }
注意:当输入x=0时,ReLU的导数是次梯度,通常取0或1都可,你可以根据需求调整分支逻辑。
2. 使用AutoDiff框架内置的max函数
如果你的AutoDiff框架提供了内置的max/min函数(比如多数深度学习框架都有),优先用这些函数——它们已经内置了正确的微分规则,不需要自己实现。
关于结果图的补充
紫蓝线条重叠是因为ReLU_ol和ReLU_if的导数结果完全一致,绘图时自然重叠;ReLU_ol未绘制大概率是绘图代码的逻辑问题(比如未启用该系列的绘制),和微分实现本身无关。
内容的提问来源于stack exchange,提问作者Arek
相关产品推荐
相关产品推荐

