C++模板Expression类derivative方法如何根据op模板参数指定返回类型
解决方案
首先修正两处会阻塞编译的前置笔误:
- 运算符重载逻辑写反:当前
operator+返回Multiply类型表达式、operator*返回Add类型表达式,和语义完全不符,需要先调整 Variable<T>::derivative中调用了未定义的constant(1),需要改为Constant<T>(1)
方法1:C++17及以上版本(最简实现)
因为Expression的op是编译期确定的模板参数,用if constexpr替换运行时switch,即可让编译器自动推导符合规则的返回类型:
// 在Expression类内实现derivative方法 auto derivative() const { if constexpr (op == Add) { // 加法求导法则:f’+g’,对应你给出的第一类返回类型规则 return l_.derivative() + r_.derivative(); } else if constexpr (op == Multiply) { // 乘法求导法则:f’g + fg’,对应你给出的第二类返回类型规则 return l_.derivative() * r_ + l_ * r_.derivative(); } }
该实现无需手动指定返回类型,编译器会按照运算逻辑自动生成完全匹配规则的返回类型。
方法2:C++14兼容实现
如果需要兼容C++14,没有if constexpr语法,可以通过类型萃取提前推导返回类型:
首先定义返回类型萃取工具:
#include <type_traits> template<typename L, typename R, OP_enum op> struct DerivativeReturnType { using LDeriv = decltype(std::declval<const L>().derivative()); using RDeriv = decltype(std::declval<const R>().derivative()); // 加法求导返回类型 using AddResult = Expression<LDeriv, RDeriv, Add>; // 乘法求导返回类型 using MultiplyResult = Expression< Expression<LDeriv, R, Multiply>, Expression<L, RDeriv, Multiply>, Add >; using type = typename std::conditional<op == Add, AddResult, MultiplyResult>::type; };
然后在Expression类内实现derivative方法:
typename DerivativeReturnType<L, R, op>::type derivative() const { if (op == Add) { return l_.derivative() + r_.derivative(); } else { return l_.derivative() * r_ + l_ * r_.derivative(); } }
因为op是编译期常量,编译器会自动消除无效分支,不会出现类型不匹配的问题。
修正后的运算符重载代码
template<typename L, typename R> Expression<L, R, Add> operator+(const L & l, const R & r) { return Expression<L, R, Add>(l, r); } template<typename L, typename R> Expression<L, R, Multiply> operator*(const L & l, const R & r) { return Expression<L, R, Multiply>(l, r); }
内容的提问来源于stack exchange,提问作者user14999310
相关产品推荐
相关产品推荐

