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

C++静态/动态内存张量类如何共享方法以避免代码重复?

复用张量类逻辑:基于容器抽象的模板基类方案

你遇到的问题本质是如何在不同底层容器(std::array/std::vector)之上复用业务逻辑,不用重复实现相同的方法。这里推荐用模板基类+容器策略的方案,完全利用C++编译时多态的特性,既避免重复代码,又不损失性能。

核心思路

把「容器类型」作为模板参数抽象出来,将所有通用方法(比如元素遍历、赋值、数学运算等)放到一个通用基类中,而StaticTensor和DynamicTensor只负责各自特有的逻辑:

  • StaticTensor:处理编译期固定的维度,用std::array作为底层容器
  • DynamicTensor:处理运行期可变的维度,用std::vector作为底层容器

因为std::array和std::vector的接口高度兼容(都支持begin()/end()、operator[]、size()),基类里的通用方法可以直接复用。

具体实现代码

第一步:公共基类 TensorBase

这个基类封装所有和容器无关的通用逻辑,依赖容器的通用接口来实现:

#include <array>
#include <vector>
#include <algorithm>
#include <iostream>
#include <tuple>

// 保留你的Product constexpr函数
template <typename... data_type>
constexpr auto Product(data_type... _values) {
    return (_values * ...);
}

// 通用张量基类:容器类型作为模板参数
template <typename Container>
class TensorBase {
protected:
    Container Entries; // 底层容器,子类直接访问

public:
    // 通用方法:填充所有元素
    void Fill(const typename Container::value_type& val) {
        std::fill(Entries.begin(), Entries.end(), val);
    }

    // 通用方法:遍历打印所有元素
    void Print() const {
        for (const auto& elem : Entries) {
            std::cout << elem << " ";
        }
        std::cout << "\n";
    }

    // 通用方法:访问元素(返回引用)
    typename Container::reference operator[](size_t idx) {
        return Entries[idx];
    }

    typename Container::const_reference operator[](size_t idx) const {
        return Entries[idx];
    }

    // 获取总元素数
    size_t Size() const {
        return Entries.size();
    }
};

第二步:静态张量类 StaticTensor

继承自TensorBase,指定std::array作为容器,处理编译期固定维度:

template <class t_data_type, unsigned... t_dimensions>
class StaticTensor : public TensorBase<std::array<t_data_type, Product(t_dimensions...)>> {
public:
    using Base = TensorBase<std::array<t_data_type, Product(t_dimensions...)>>;
    using Base::Base; // 继承基类构造函数

    // 静态张量特有方法:获取编译期维度
    constexpr auto GetDimensions() const {
        return std::make_tuple(t_dimensions...);
    }

    // 静态维度相关的特有逻辑
    void StaticMethod() {
        std::cout << "Static tensor with dimensions: ";
        auto print_dim = [](auto dim) { std::cout << dim << " "; };
        std::apply([&](auto... dims) { (print_dim(dims), ...); }, GetDimensions());
        std::cout << "\n";
    }
};

第三步:动态张量类 DynamicTensor

同样继承自TensorBase,指定std::vector作为容器,处理运行期可变维度:

template <class t_data_type>
class DynamicTensor : public TensorBase<std::vector<t_data_type>> {
private:
    std::vector<size_t> dimensions; // 存储运行期维度

public:
    using Base = TensorBase<std::vector<t_data_type>>;
    using Base::Base;

    // 动态张量特有方法:调整维度
    template <typename... t_dimensions>
    void Resize(t_dimensions... dims) {
        dimensions = {static_cast<size_t>(dims)...};
        this->Entries.resize(Product(dims...));
    }

    // 获取运行期维度
    const std::vector<size_t>& GetDimensions() const {
        return dimensions;
    }

    // 动态维度相关的特有逻辑
    void DynamicMethod() {
        std::cout << "Dynamic tensor with dimensions: ";
        for (auto dim : dimensions) {
            std::cout << dim << " ";
        }
        std::cout << "\n";
    }
};

测试代码

int main() {
    // 静态张量:2x3的int张量
    StaticTensor<int, 2, 3> static_tensor;
    static_tensor.Fill(5); // 调用基类通用方法
    static_tensor.Print(); // 调用基类通用方法
    static_tensor.StaticMethod(); // 调用子类特有方法

    // 动态张量:先初始化,再调整为3x4的float张量
    DynamicTensor<float> dynamic_tensor;
    dynamic_tensor.Resize(3,4);
    dynamic_tensor.Fill(3.14f);
    dynamic_tensor.Print();
    dynamic_tensor.DynamicMethod();

    return 0;
}

方案优势

  1. 完全避免重复代码:所有通用逻辑都在TensorBase里实现一次,两个子类直接复用
  2. 编译期安全:静态张量的维度在编译期确定,没有运行期开销;动态张量的维度运行期调整,灵活度高
  3. 扩展性强:如果以后要支持其他容器(比如自定义的连续内存容器),只需要继承TensorBase并传入对应的容器类型即可
  4. 性能无损失:没有使用虚函数,所有调用都是编译期解析的,和手写的单容器版本性能一致

补充说明

如果你的通用方法需要维度信息,可以让子类各自实现GetDimensions接口——静态张量返回编译期tuple,动态张量返回运行期vector,基类不需要关心细节,只在需要维度的子类逻辑中使用即可。

内容的提问来源于stack exchange,提问作者niran90

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.30 17:42:50