Make IDataTransfer be directly shared with shared providers (#7215)

This commit is contained in:
Ryan Hill 2021-04-01 20:39:16 -07:00 committed by GitHub
parent 0ebeaf529d
commit 5a6d477625
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
5 changed files with 15 additions and 11 deletions

View file

@ -4,7 +4,9 @@
#pragma once
#include "core/common/status.h"
#ifndef SHARED_PROVIDER
#include "core/framework/tensor.h"
#endif
namespace onnxruntime {

View file

@ -353,7 +353,8 @@ struct ProviderHostImpl : ProviderHost {
Status DataTransferManager__CopyTensor(const DataTransferManager* p, const Tensor& src, Tensor& dst, int exec_queue_id) override { return p->CopyTensor(src, dst, exec_queue_id); }
// IDataTransfer
void IDataTransfer__operator_delete(IDataTransfer* p) override { delete p; }
Status IDataTransfer__CopyTensor(const IDataTransfer* p, const Tensor& src, Tensor& dst) override { return p->IDataTransfer::CopyTensor(src, dst); }
Status IDataTransfer__CopyTensors(const IDataTransfer* p, const std::vector<IDataTransfer::SrcDstPair>& src_dst_pairs) override { return p->IDataTransfer::CopyTensors(src_dst_pairs); }
// IndexedSubGraph_MetaDef
std::unique_ptr<IndexedSubGraph_MetaDef> IndexedSubGraph_MetaDef__construct() override { return onnxruntime::make_unique<IndexedSubGraph::MetaDef>(); }

View file

@ -127,7 +127,6 @@ struct Capture;
} // namespace logging
struct ComputeCapability;
struct DataTransferManager;
struct IDataTransfer;
struct IndexedSubGraph;
struct IndexedSubGraph_MetaDef;
struct KernelCreateInfo;
@ -152,6 +151,7 @@ using MLDataType = const DataTypeImpl*;
using NodeArgInfo = ONNX_NAMESPACE::ValueInfoProto;
} // namespace onnxruntime
#include "core/framework/data_transfer.h"
#include "core/framework/execution_provider.h"
#include "provider_interfaces.h"
#include "core/framework/op_kernel.h"

View file

@ -63,6 +63,14 @@ MLDataType DataTypeImpl::GetTensorType<float>() {
return g_host->DataTypeImpl_GetTensorType_float();
}
Status IDataTransfer::CopyTensor(const Tensor& src, Tensor& dst) const {
return g_host->IDataTransfer__CopyTensor(this, src, dst);
}
Status IDataTransfer::CopyTensors(const std::vector<SrcDstPair>& src_dst_pairs) const {
return g_host->IDataTransfer__CopyTensors(this, src_dst_pairs);
}
TensorShape::TensorShape(const int64_t* dimension_sizes, size_t dimension_count)
: std::vector<int64_t>(dimension_count) {
for (size_t i = 0; i < dimension_count; ++i) {

View file

@ -302,7 +302,8 @@ struct ProviderHost {
virtual Status DataTransferManager__CopyTensor(const DataTransferManager* p, const Tensor& src, Tensor& dst, int exec_queue_id) = 0;
// IDataTransfer
virtual void IDataTransfer__operator_delete(IDataTransfer* p) = 0;
virtual Status IDataTransfer__CopyTensor(const IDataTransfer* p, const Tensor& src, Tensor& dst) = 0;
virtual Status IDataTransfer__CopyTensors(const IDataTransfer* p, const std::vector<IDataTransfer::SrcDstPair>& src_dst_pairs) = 0;
// IndexedSubGraph_MetaDef
virtual std::unique_ptr<IndexedSubGraph_MetaDef> IndexedSubGraph_MetaDef__construct() = 0;
@ -716,14 +717,6 @@ struct DataTransferManager {
PROVIDER_DISALLOW_ALL(DataTransferManager)
};
struct IDataTransfer {
static void operator delete(void* p) { g_host->IDataTransfer__operator_delete(reinterpret_cast<IDataTransfer*>(p)); }
IDataTransfer() = delete;
IDataTransfer(const IDataTransfer&) = delete;
void operator=(const IDataTransfer&) = delete;
};
struct IndexedSubGraph_MetaDef {
static std::unique_ptr<IndexedSubGraph_MetaDef> Create() { return g_host->IndexedSubGraph_MetaDef__construct(); }
static void operator delete(void* p) { g_host->IndexedSubGraph_MetaDef__operator_delete(reinterpret_cast<IndexedSubGraph_MetaDef*>(p)); }