blob: d92f8d19befec35122ae546c86d61c4f7cdc283f [file]
/*
* Copyright (c) Meta Platforms, Inc. and affiliates.
* All rights reserved.
*
* This source code is licensed under the BSD-style license found in the
* LICENSE file in the root directory of this source tree.
*/
#include <executorch/runtime/core/exec_aten/exec_aten.h>
#include <executorch/runtime/core/exec_aten/util/dim_order_util.h>
#include <executorch/runtime/core/exec_aten/util/tensor_util.h>
#include <executorch/runtime/platform/assert.h>
#include <memory>
// NOTE: required by torchchat install_et.sh script.
// @nolint PATTERNLINT Ok to use stdlib for this optional library
#include <vector>
#ifdef USE_ATEN_LIB
#include <torch/torch.h>
#else
#include <executorch/runtime/core/portable_type/tensor.h>
#endif
#pragma once
namespace torch {
namespace executor {
/**
* A tensor wrapper takes ownership of all the memory of the necessary metadata
* for torch::executor::Tensor. Note that it doesn't own the data memory.
*/
class ManagedTensor {
public:
/// The type used for elements of `sizes()`.
using SizesType = exec_aten::SizesType;
/// The type used for elements of `dim_order()`.
using DimOrderType = exec_aten::DimOrderType;
/// The type used for elements of `strides()`.
using StridesType = exec_aten::StridesType;
ManagedTensor() = delete;
explicit ManagedTensor(
void* data,
const std::vector<SizesType>& sizes,
ScalarType dtype)
: dtype_(dtype), sizes_(sizes), data_ptr_(data) {
#ifdef USE_ATEN_LIB
tensor_ = torch::from_blob(data, sizes, dtype_);
#else
ssize_t dim = sizes.size();
dim_order_.resize(dim);
strides_.resize(dim);
for (size_t i = 0; i < dim; ++i) {
dim_order_[i] = i;
}
dim_order_to_stride_nocheck(
sizes.data(), dim_order_.data(), dim, strides_.data());
tensor_impl_ = std::make_unique<TensorImpl>(
dtype_,
dim,
sizes_.data(),
data_ptr_,
dim_order_.data(),
strides_.data(),
TensorShapeDynamism::DYNAMIC_BOUND);
#endif
}
void resize(const std::vector<SizesType>& new_sizes) {
ET_CHECK_MSG(
new_sizes.size() == sizes_.size(),
"Cannot change rank of a managed tensor");
auto err = resize_tensor(
this->get_aliasing_tensor(),
exec_aten::ArrayRef<SizesType>(new_sizes.data(), new_sizes.size()));
ET_CHECK(err == Error::Ok);
}
/**
* Get the underlying Tensor object. This is assuming the copying is cheap.
*/
Tensor get_aliasing_tensor() {
#ifdef USE_ATEN_LIB
return tensor_;
#else
return Tensor(tensor_impl_.get());
#endif
}
private:
ScalarType dtype_;
std::unique_ptr<TensorImpl> tensor_impl_;
std::vector<SizesType> sizes_;
std::vector<StridesType> strides_;
std::vector<DimOrderType> dim_order_;
void* data_ptr_ = nullptr;
#ifdef USE_ATEN_LIB
Tensor tensor_;
#endif
};
} // namespace executor
} // namespace torch