Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions onnxruntime/core/framework/plugin_data_transfer.cc
Original file line number Diff line number Diff line change
Expand Up @@ -51,14 +51,14 @@ Status DataTransfer::CopyTensors(const std::vector<SrcDstPair>& src_dst_pairs) c
}

// optimized version for a single copy. see comments above in CopyTensors regarding the OrtValue usage and const_cast
Status DataTransfer::CopyTensorImpl(const Tensor& src_tensor, Tensor& dst_tensor, onnxruntime::Stream* /*stream*/) const {
Status DataTransfer::CopyTensorImpl(const Tensor& src_tensor, Tensor& dst_tensor, onnxruntime::Stream* stream) const {
OrtValue src, dst;
Tensor* src_tensor_ptr = const_cast<Tensor*>(&src_tensor);
src.Init(static_cast<void*>(src_tensor_ptr), ml_tensor_type, no_op_deleter);
dst.Init(static_cast<void*>(&dst_tensor), ml_tensor_type, no_op_deleter);
const OrtValue* src_ptr = &src;
OrtValue* dst_ptr = &dst;
OrtSyncStream* stream_ptr = nullptr; // static_cast<OrtSyncStream*>(stream);
OrtSyncStream* stream_ptr = reinterpret_cast<OrtSyncStream*>(stream);
auto* status = impl_.CopyTensors(&impl_, &src_ptr, &dst_ptr, &stream_ptr, 1);

return ToStatusAndRelease(status);
Expand Down
42 changes: 42 additions & 0 deletions onnxruntime/test/framework/data_transfer_manager_test.cc
Original file line number Diff line number Diff line change
Expand Up @@ -7,12 +7,54 @@
#include "core/common/inlined_containers.h"
#include "core/framework/data_transfer_manager.h"
#include "core/framework/ort_value.h"
#include "core/framework/plugin_data_transfer.h"
#include "core/framework/stream_handles.h"
#include "test/unittest_util/framework_test_utils.h"
#include "test/util/include/asserts.h"

namespace onnxruntime {
namespace test {

TEST(DataTransferManagerTest, PluginCopiesForwardStreams) {
struct TestDataTransfer final : OrtDataTransferImpl {
TestDataTransfer() : OrtDataTransferImpl{} {
ort_version_supported = ORT_API_VERSION;
Release = [](OrtDataTransferImpl*) noexcept {};
CopyTensors = [](OrtDataTransferImpl* impl, const OrtValue**, OrtValue**,
OrtSyncStream** streams, size_t num_tensors) noexcept -> OrtStatus* {
auto& self = *static_cast<TestDataTransfer*>(impl);
self.copied_tensors += num_tensors;
self.last_stream = streams != nullptr && num_tensors > 0 ? streams[0] : nullptr;
return nullptr;
};
}

ORT_DISALLOW_COPY_ASSIGNMENT_AND_MOVE(TestDataTransfer);

size_t copied_tensors = 0;
OrtSyncStream* last_stream = nullptr;
} impl;

plugin_ep::DataTransfer data_transfer{impl};
auto allocator = TestCPUExecutionProvider()->CreatePreferredAllocators()[0];
Tensor source{DataTypeImpl::GetType<float>(), TensorShape{4}, allocator};
Tensor destination{DataTypeImpl::GetType<float>(), TensorShape{4}, allocator};
OrtDevice device;
Stream stream{nullptr, device};

ASSERT_STATUS_OK(data_transfer.CopyTensorAsync(source, destination, stream));
EXPECT_EQ(impl.copied_tensors, 1U);
EXPECT_EQ(impl.last_stream, reinterpret_cast<OrtSyncStream*>(&stream));

ASSERT_STATUS_OK(data_transfer.CopyTensor(source, destination));
EXPECT_EQ(impl.copied_tensors, 2U);
EXPECT_EQ(impl.last_stream, nullptr);

ASSERT_STATUS_OK(data_transfer.CopyTensors({{source, destination, &stream}}));
EXPECT_EQ(impl.copied_tensors, 3U);
EXPECT_EQ(impl.last_stream, reinterpret_cast<OrtSyncStream*>(&stream));
}

// DataTransferManager::CopyTensors should validate sizes match before calling the IDataTransfer implementation
TEST(DataTransferManagerTest, BatchedTensorCopyBadSize) {
auto allocator = TestCPUExecutionProvider()->CreatePreferredAllocators()[0];
Expand Down
Loading