C++递归参数包使用中如何避免代码重复?
优化可变参数模板实现张量的代码冗余问题
核心优化思路
- 移除冗余特化:原代码中2维张量的特化版本完全可以被可变参数版本覆盖,直接删除即可消除重复代码。
- 统一递归逻辑:将1维张量作为递归终止节点,多维张量复用递归逻辑,避免为不同维度编写重复的方法实现。
- 提取公共接口:所有维度的张量共享构造、
zero()、迭代器等公共方法,无需重复定义。
优化后的代码实现
#include <cstddef> #include <array> #include <iostream> #include <algorithm> template<typename T, std::size_t... Dims> class Tensor; // 1维张量:递归终止版本 template<typename T, std::size_t N> class Tensor<T, N> { protected: using Type = std::array<T, N>; Type data; public: Tensor() { zero(); } void zero() { fill(T{}); } typename Type::iterator begin() { return data.begin(); } typename Type::iterator end() { return data.end(); } typename Type::const_iterator begin() const { return data.begin(); } typename Type::const_iterator end() const { return data.end(); } void fill(T value) { std::fill(data.begin(), data.end(), value); } void print() const { std::cout << "["; bool first = true; for (const auto& v : data) { if (!first) std::cout << ","; first = false; std::cout << v; } std::cout << "]"; } }; // 多维张量:递归版本(支持2维及以上) template<typename T, std::size_t N, std::size_t... RestDims> class Tensor<T, N, RestDims...> { private: using SubTensor = Tensor<T, RestDims...>; using Type = std::array<SubTensor, N>; Type data; public: Tensor() { zero(); } void zero() { fill(T{}); } typename Type::iterator begin() { return data.begin(); } typename Type::iterator end() { return data.end(); } typename Type::const_iterator begin() const { return data.begin(); } typename Type::const_iterator end() const { return data.end(); } void fill(T value) { for (auto& sub_tensor : data) { sub_tensor.fill(value); } } void print() const { std::cout << "["; bool first = true; for (const auto& sub_tensor : data) { if (!first) std::cout << ","; first = false; sub_tensor.print(); } std::cout << "]"; } }; // 测试示例 int main() { Tensor<int, 2, 3, 4> multi_dim_tensor; multi_dim_tensor.fill(5); multi_dim_tensor.print(); std::cout << "\n"; Tensor<double, 3> one_dim_tensor; one_dim_tensor.fill(3.14); one_dim_tensor.print(); std::cout << "\n"; return 0; }
优化细节说明
- 删除冗余特化:原代码中
Tensor<T, N, M>特化完全可以被Tensor<T, N, RestDims...>(此时RestDims...为单个维度M)处理,删除后直接减少约三分之一的重复代码。 - 递归逻辑复用:多维张量的
fill和print方法通过遍历子张量并调用其对应方法实现递归,自然适配任意维度,无需为每个维度单独编写逻辑。 - 输出格式优化:添加
first标志解决了原代码中末尾多余逗号的问题,输出更规范。 - 接口一致性:所有维度的张量提供完全一致的公共接口,使用时无需区分维度差异。
内容的提问来源于stack exchange,提问作者Cedric Martens
相关产品推荐
相关产品推荐

