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

基于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,要么逻辑悄悄出错,要么只能在运行时抛异常,没法在编译阶段就把问题揪出来。

替代方案:基于类型擦除的设计

要实现类似模板化虚函数的效果,同时支持将子类放入同一个容器,可以用类型擦除结合模板中间基类的方式,核心思路是把模板参数的差异封装到内部,对外提供统一的非模板接口。

方案实现步骤

  1. 定义非模板的抽象基类:这个基类只提供统一的接口,不涉及模板参数,专门用来放进vector这类容器。
  2. 定义模板中间基类:继承自抽象基类,实现模板化的Forward/Backward逻辑,并通过类型擦除处理不同秩的Tensor,适配到抽象基类的统一接口。
  3. 子类继承模板中间基类:指定具体的输入/输出秩,实现对应秩的业务逻辑即可。

具体代码示例

第一步:抽象基类(用于容器存储)

#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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.03 09:46:05