Skip to content
Open
Show file tree
Hide file tree
Changes from 1 commit
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
43 changes: 30 additions & 13 deletions src/amdgpu/gpu_data_transfer.cc
Original file line number Diff line number Diff line change
Expand Up @@ -6,8 +6,8 @@

namespace gpu_ep {

DataTransfer::DataTransfer(const ProviderFactory& factory, OrtEpFactory* backend_factory)
: OrtDataTransferImpl{ORT_API_VERSION}, factory_{factory}, backend_factory_{backend_factory}
DataTransfer::DataTransfer(const ProviderFactory& factory)
: OrtDataTransferImpl{ORT_API_VERSION}, factory_{factory}
{
OrtDataTransferImpl::Release = [](OrtDataTransferImpl* this_) noexcept {
// ORT creates one DataTransfer per session and owns it, delete on Release.
Expand All @@ -28,36 +28,53 @@ DataTransfer::DataTransfer(const ProviderFactory& factory, OrtEpFactory* backend
API_CALL_S(DataTransfer, this_, CopyTensors, src_tensors, dst_tensors, streams, num_tensors);
};

// Snapshot the backend transfer once from the backend this session selected.
// A null backend_factory_ (library-registration-time creation) leaves this
// instance inert: CanCopy returns false, CopyTensors returns an error.
if (backend_factory_ != nullptr) {
if (backend_factory_->CreateDataTransfer(backend_factory_, &backend_data_transfer_) != nullptr) {
// The backend is resolved on first use, not here: this constructor also runs at
// library-registration time, before any backend has been selected.
}

OrtDataTransferImpl* DataTransfer::GetBackendDataTransfer() const noexcept {
const auto backend_factory{factory_.GetBackendFactory()};
if (backend_factory == nullptr) {
// No backend selected yet (e.g. called before any CreateEp). Not an error —
// this instance resolves once a backend exists.
return nullptr;
}
if (backend_factory != backend_factory_) {
// First use, or the backend changed (e.g. profile switch) — (re-)query the
// transfer from the currently selected backend factory. The backend owns it
// (Release is a no-op there), so there is nothing to release for the old one.
backend_data_transfer_ = nullptr;
backend_factory_ = nullptr;
if (backend_factory->CreateDataTransfer(backend_factory, &backend_data_transfer_) != nullptr) {
backend_data_transfer_ = nullptr;
backend_factory_ = nullptr;
return nullptr;
}
backend_factory_ = backend_factory;
}
return backend_data_transfer_;
}

bool DataTransfer::CanCopy(const OrtMemoryDevice* src_memory_device,
const OrtMemoryDevice* dst_memory_device) const noexcept
{
if (backend_data_transfer_ == nullptr) {
const auto backend_data_transfer{GetBackendDataTransfer()};
if (backend_data_transfer == nullptr) {
return false;
}
return backend_data_transfer_->CanCopy(backend_data_transfer_,
return backend_data_transfer->CanCopy(backend_data_transfer,
src_memory_device, dst_memory_device);
}

Ort::Status DataTransfer::CopyTensors(const OrtValue** src_tensors,
OrtValue** dst_tensors, OrtSyncStream** streams, size_t num_tensors) const noexcept
{
if (backend_data_transfer_ == nullptr) {
const auto backend_data_transfer{GetBackendDataTransfer()};
if (backend_data_transfer == nullptr) {
return MAKE_STATUS(ORT_EP_FAIL, "invalid backend factory");
}
RETURN_IF_ERROR(backend_data_transfer_->CopyTensors(backend_data_transfer_,
RETURN_IF_ERROR(backend_data_transfer->CopyTensors(backend_data_transfer,
src_tensors, dst_tensors, streams, num_tensors));
return STATUS_OK;
}

} // namespace gpu_ep
} // namespace gpu_ep
26 changes: 17 additions & 9 deletions src/amdgpu/gpu_data_transfer.h
Original file line number Diff line number Diff line change
Expand Up @@ -11,13 +11,16 @@ struct ProviderFactory;

struct DataTransfer : OrtDataTransferImpl {
DataTransfer() = delete;
// Per-session: the backend is snapshotted at construction from the factory that
// the session's CreateEp just selected, and never re-queried. ORT creates one
// DataTransfer per session (CreateDataTransfer -> GetDataTransfer) and owns it,
// Releasing (deleting) it at session teardown. backend_factory may be null when
// ORT creates a transfer at library-registration time (before any backend is
// selected); such an instance is inert.
DataTransfer(const ProviderFactory& factory, OrtEpFactory* backend_factory);
// The backend is resolved lazily, on first use, and re-resolved if the selected
// backend changes (e.g. profile switch) — see GetBackendDataTransfer.
//
// Lazy is required, not merely convenient: ORT creates a DataTransfer per session
// (CreateDataTransfer -> GetDataTransfer) *and* one per factory at library
// registration time, before any CreateEp has run. That registration-time instance
// is the only one reachable from the env-level OrtApi::CopyTensors, so resolving
// the backend in the constructor leaves that path permanently broken.
// ORT owns each instance and Releases (deletes) it at teardown.
explicit DataTransfer(const ProviderFactory& factory);

private:
bool CanCopy(const OrtMemoryDevice* src_memory_device,
Expand All @@ -26,9 +29,14 @@ struct DataTransfer : OrtDataTransferImpl {
[[nodiscard]] Ort::Status CopyTensors(const OrtValue** src_tensors,
OrtValue** dst_tensors, OrtSyncStream** streams, size_t num_tensors) const noexcept;

// Returns the backend's data transfer, re-querying it if the selected backend
// changed (e.g. profile switch). Null until a backend has been selected.
OrtDataTransferImpl* GetBackendDataTransfer() const noexcept;

mutable OrtEpFactory* backend_factory_{};
mutable OrtDataTransferImpl* backend_data_transfer_{};

const ProviderFactory& factory_;
OrtEpFactory* backend_factory_{};
OrtDataTransferImpl* backend_data_transfer_{};
};

} // namespace gpu_ep
11 changes: 6 additions & 5 deletions src/amdgpu/gpu_factory.cc
Original file line number Diff line number Diff line change
Expand Up @@ -326,11 +326,12 @@ void ProviderFactory::ReleaseAllocator(OrtAllocator*) const {
}

Ort::Status ProviderFactory::CreateDataTransfer(OrtDataTransferImpl** data_transfer) {
// Per-session: hand ORT a fresh DataTransfer bound to the backend this session
// selected (GetBackendFactory() reflects the CreateEp that just ran; it is null at
// library-registration time, yielding an inert instance). ORT owns it and calls
// Release (which deletes it) at session teardown.
*data_transfer = std::make_unique<DataTransfer>(*this, GetBackendFactory()).release();
// ORT calls this once per session (after that session's CreateEp) and once per
// factory at library-registration time, before any backend has been selected. The
// registration-time instance is the one behind the env-level OrtApi::CopyTensors,
// so DataTransfer resolves its backend lazily rather than at construction.
// ORT owns each instance and calls Release (which deletes it) at teardown.
*data_transfer = std::make_unique<DataTransfer>(*this).release();
return STATUS_OK;
}

Expand Down