基于Eigen Tensor的Layer类优化:减少重复虚函数的设计方案
问题描述
我正在用支持模板的Eigen Tensor库编写类库,有一个名为Layer的基类,多个子类继承自该基类。每个子类需要实现特定的虚函数,比如void Forward(const Tensor<2> &input)和void Backward(const Tensor<2> &gradient)。但每个子类仅接受特定秩(Rank)的输入和梯度,比如Tensor<2>、Tensor<3>这类,因此基类最终包含一堆重载的虚函数:
virtual void Forward(const Tensor<2> &input) virtual void Forward(const Tensor<3> &input) virtual void Forward(const Tensor<4> &input) virtual void Backward(const Tensor<2> &gradient) virtual void Backward(const Tensor<3> &gradient) virtual void Backward(const Tensor<4> &gradient)
由于C++无法将虚函数模板化,想问下当前设计是不是有问题?有没有能实现类似模板化虚函数功能的替代方案,同时还能把子类对象放进同一个vector里?理想中的模板化虚函数大概是这样:
template<int InputRank> virtual void Forward(const Tensor<Rank> &input) template<int OutputRank> virtual void Backward(const Tensor<OutputRank> &gradient)
分析与解决方案
当前设计的问题
当前这种写法确实有不少槽点:
- 代码冗余到爆炸:基类得给每一种可能用到的Tensor秩都写一遍
Forward和Backward的重载,以后要是要支持5维、6维Tensor,还得继续加,维护起来简直头疼。 - 接口设计不合理:每个子类明明只需要处理特定秩的Tensor,但基类逼着所有子类都要实现所有重载——大部分重载要么是空函数,要么只能扔个运行时错误,完全违背了“只给子类需要的接口”的原则。
- 类型安全没保障:要是不小心给子类传了它不支持的秩的Tensor,要么逻辑悄悄出错,要么只能在运行时抛异常,没法在编译阶段就把问题揪出来。
替代方案:基于类型擦除的设计
要实现类似模板化虚函数的效果,同时支持将子类放入同一个容器,可以用类型擦除结合模板中间基类的方式,核心思路是把模板参数的差异封装到内部,对外提供统一的非模板接口。
方案实现步骤
- 定义非模板的抽象基类:这个基类只提供统一的接口,不涉及模板参数,专门用来放进
vector这类容器。 - 定义模板中间基类:继承自抽象基类,实现模板化的
Forward/Backward逻辑,并通过类型擦除处理不同秩的Tensor,适配到抽象基类的统一接口。 - 子类继承模板中间基类:指定具体的输入/输出秩,实现对应秩的业务逻辑即可。
具体代码示例
第一步:抽象基类(用于容器存储)
#include <Eigen/Core> #include <Eigen/Tensor> #include <any> #include <memory> #include <vector> class Layer { public: virtual ~Layer() = default; // 统一的Forward/Backward接口,用std::any擦除Tensor的具体类型 virtual void Forward(const std::any& input) = 0; virtual void Backward(const std::any& gradient) = 0; };
第二步:模板中间基类(处理类型擦除)
template<int InputRank, int OutputRank> class TypedLayer : public Layer { public: void Forward(const std::any& input) override { try { // 尝试将any转换为指定秩的Tensor const auto& tensor = std::any_cast<const Eigen::Tensor<double, InputRank>&>(input); // 调用子类实现的具体前向逻辑 ForwardImpl(tensor); } catch (const std::bad_any_cast&) { throw std::runtime_error("Forward输入Tensor秩不匹配"); } } void Backward(const std::any& gradient) override { try { const auto& tensor = std::any_cast<const Eigen::Tensor<double, OutputRank>&>(gradient); BackwardImpl(tensor); } catch (const std::bad_any_cast&) { throw std::runtime_error("Backward梯度Tensor秩不匹配"); } } // 子类需要实现的、绑定了具体秩的纯虚函数 virtual void ForwardImpl(const Eigen::Tensor<double, InputRank>& input) = 0; virtual void BackwardImpl(const Eigen::Tensor<double, OutputRank>& gradient) = 0; };
第三步:子类实现(指定具体秩)
// 示例:只接受2维输入和2维梯度的全连接层 class DenseLayer : public TypedLayer<2, 2> { public: void ForwardImpl(const Eigen::Tensor<double, 2>& input) override { // 这里写全连接层的前向传播逻辑 } void BackwardImpl(const Eigen::Tensor<double, 2>& gradient) override { // 这里写全连接层的反向传播逻辑 } }; // 示例:接受4维输入(比如图像数据)和2维梯度的卷积层 class ConvLayer : public TypedLayer<4, 2> { public: void ForwardImpl(const Eigen::Tensor<double, 4>& input) override { // 这里写卷积层的前向传播逻辑 } void BackwardImpl(const Eigen::Tensor<double, 2>& gradient) override { // 这里写卷积层的反向传播逻辑 } };
使用方式
int main() { std::vector<std::unique_ptr<Layer>> layers; layers.push_back(std::make_unique<DenseLayer>()); layers.push_back(std::make_unique<ConvLayer>()); // 构造对应秩的Tensor输入 Eigen::Tensor<double, 2> dense_input(32, 128); // batch=32, feature=128 Eigen::Tensor<double, 4> conv_input(32, 28, 28, 3); // batch=32, 28x28图像, 3通道 // 调用对应层的Forward layers[0]->Forward(dense_input); layers[1]->Forward(conv_input); return 0; }
方案优势
- 接口简洁:基类只需要两个统一的接口,不用为每个秩写一堆重载。
- 编译期类型检查:子类继承
TypedLayer时指定了输入/输出秩,ForwardImpl/BackwardImpl的参数类型是确定的,编译阶段就能保证类型正确。 - 扩展性强:后续要支持新的秩,只需要创建对应参数的
TypedLayer子类,完全不用修改基类。 - 支持容器存储:所有子类都继承自非模板的
Layer基类,可以安全放进vector等容器统一管理。
其他可选方案
如果觉得std::any太灵活、不够安全,也可以用std::variant限定允许的Tensor类型,比如:
#include <variant> // 预先定义支持的所有Tensor类型 using TensorVariant = std::variant< Eigen::Tensor<double,2>, Eigen::Tensor<double,3>, Eigen::Tensor<double,4> >; class Layer { public: virtual ~Layer() = default; virtual void Forward(const TensorVariant& input) = 0; virtual void Backward(const TensorVariant& gradient) = 0; };
这种方式比std::any更严谨,因为它限定了只能是指定的几种Tensor类型,编译期就能排除非法类型。
内容的提问来源于stack exchange,提问作者rkuang25
相关产品推荐
相关产品推荐

