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

如何为自定义Tensor容器创建只读视图?C++20实现疑问

关于Tensor只读视图的设计方案评估

你的基类继承方案合理性

你的方案是合理的,能有效避免代码重复,但需要注意几个关键细节:

  • TensorView作为基类,仅包含const成员函数(如const版本的at()、打印方法)和protected的核心数据(指针、维度、内存持有标记),不负责内存管理;
  • 派生类Tensor添加可变接口(非const的at()、zero()等),并通过owns_allocation_控制析构时的内存释放逻辑;
  • 基类析构函数无需设为虚函数:因为视图本身不持有内存,不会通过基类指针销毁对象,避免不必要的虚函数开销。

该方案的小缺陷是protected成员可能增加误用风险,且继承关系会让类型系统稍显复杂。

更优实现:模板参数控制可变性

通过模板类统一实现核心逻辑,用模板参数区分可变/只读版本,完全复用代码且规避继承的潜在问题:

#include <vector>
#include <initializer_list>

template <typename T, bool IsMutable>
class TensorBase {
public:
    // 统一实现const接口
    void print() const {
        // 打印多维数组逻辑
    }

    size_t dimension(size_t idx) const {
        return dims_.at(idx);
    }

    const T& at(std::initializer_list<size_t> indices) const {
        // 计算索引并返回const引用
        return data_[calculate_index(indices)];
    }

    // 切片返回只读视图,无论当前是否可变
    TensorBase<T, false> slice(size_t start_dim, size_t end_dim) const {
        // 计算子数组指针与维度
        T* slice_data = data_ + calculate_offset(start_dim);
        std::vector<size_t> slice_dims(dims_.begin() + start_dim, dims_.begin() + end_dim);
        return TensorBase<T, false>(slice_data, slice_dims, false);
    }

protected:
    T* data_;
    std::vector<size_t> dims_;
    bool owns_allocation_;

    // 构造函数仅允许内部/派生类调用
    TensorBase(T* data, std::vector<size_t> dims, bool owns_allocation)
        : data_(data), dims_(std::move(dims)), owns_allocation_(owns_allocation) {}

    ~TensorBase() {
        // 仅当可变且持有内存时释放
        if (owns_allocation_ && IsMutable) {
            delete[] data_;
        }
    }

private:
    size_t calculate_index(std::initializer_list<size_t> indices) const {
        // 实现多维索引转一维的逻辑
        size_t idx = 0;
        size_t stride = 1;
        for (auto it = indices.rbegin(); it != indices.rend(); ++it) {
            idx += *it * stride;
            stride *= dims_[std::distance(it, indices.rend()) - 1];
        }
        return idx;
    }

    size_t calculate_offset(size_t start_dim) const {
        // 计算切片起始偏移
        size_t offset = 1;
        for (size_t i = start_dim; i < dims_.size(); ++i) {
            offset *= dims_[i];
        }
        return offset;
    }
};

// 可变Tensor类型
template <typename T>
class Tensor : public TensorBase<T, true> {
public:
    Tensor(std::vector<size_t> dims)
        : TensorBase<T, true>(new T[calculate_total_size(dims)], std::move(dims), true) {}

    // 可变访问接口
    T& at(std::initializer_list<size_t> indices) {
        return this->data_[this->calculate_index(indices)];
    }

    void zero() {
        // 遍历数据置零
        size_t total = calculate_total_size(this->dims_);
        for (size_t i = 0; i < total; ++i) {
            this->data_[i] = T{};
        }
    }

private:
    size_t calculate_total_size(const std::vector<size_t>& dims) const {
        size_t total = 1;
        for (size_t d : dims) total *= d;
        return total;
    }
};

// 只读视图类型
template <typename T>
using TensorView = TensorBase<T, false>;

该方案的优势

  • 严格遵循DRY原则,所有核心逻辑(索引计算、打印、切片)仅实现一次;
  • 类型系统清晰:Tensor是可变容器,TensorView是只读视图,无法通过const_cast将视图转为可变容器(模板参数不同,类型不兼容);
  • 内存管理逻辑集中,避免重复实现;
  • 切片方法直接返回TensorView,从类型层面保证只读,彻底解决调用者修改原对象的问题。

如果需要支持可变切片(允许修改原数组的子区域),可以给Tensor添加mutable_slice方法,返回不持有内存的Tensor实例即可。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.16 20:21:10