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; }
方案优势
- 完全避免重复代码:所有通用逻辑都在
TensorBase里实现一次,两个子类直接复用 - 编译期安全:静态张量的维度在编译期确定,没有运行期开销;动态张量的维度运行期调整,灵活度高
- 扩展性强:如果以后要支持其他容器(比如自定义的连续内存容器),只需要继承
TensorBase并传入对应的容器类型即可 - 性能无损失:没有使用虚函数,所有调用都是编译期解析的,和手写的单容器版本性能一致
补充说明
如果你的通用方法需要维度信息,可以让子类各自实现GetDimensions接口——静态张量返回编译期tuple,动态张量返回运行期vector,基类不需要关心细节,只在需要维度的子类逻辑中使用即可。
内容的提问来源于stack exchange,提问作者niran90
相关产品推荐
相关产品推荐

