mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-26 19:52:38 +00:00
* Convert ABI to a versioned interface. * Convert ORT_THROW_ON_ERROR to inline function to fix link errors.
70 lines
2.4 KiB
C++
70 lines
2.4 KiB
C++
// Copyright (c) Microsoft Corporation. All rights reserved.
|
|
// Licensed under the MIT License.
|
|
|
|
#include "core/framework/allocator.h"
|
|
#include "core/framework/allocatormgr.h"
|
|
#include "core/framework/utils.h"
|
|
#include "core/session/ort_apis.h"
|
|
#include <cstdlib>
|
|
#include <sstream>
|
|
|
|
namespace onnxruntime {
|
|
|
|
void* CPUAllocator::Alloc(size_t size) {
|
|
return utils::DefaultAlloc(size);
|
|
}
|
|
|
|
void CPUAllocator::Free(void* p) {
|
|
utils::DefaultFree(p);
|
|
}
|
|
|
|
const OrtMemoryInfo& CPUAllocator::Info() const { return *memory_info_; }
|
|
} // namespace onnxruntime
|
|
|
|
std::ostream& operator<<(std::ostream& out, const OrtMemoryInfo& info) { return (out << info.ToString()); }
|
|
|
|
ORT_API_STATUS_IMPL(OrtApis::CreateMemoryInfo, _In_ const char* name1, OrtAllocatorType type, int id1,
|
|
OrtMemType mem_type1, _Out_ OrtMemoryInfo** out) {
|
|
if (strcmp(name1, onnxruntime::CPU) == 0) {
|
|
*out = new OrtMemoryInfo(name1, type, OrtDevice(), id1, mem_type1);
|
|
} else if (strcmp(name1, onnxruntime::CUDA) == 0) {
|
|
*out = new OrtMemoryInfo(
|
|
name1, type, OrtDevice(OrtDevice::GPU, OrtDevice::MemType::DEFAULT, static_cast<OrtDevice::DeviceId>(id1)), id1,
|
|
mem_type1);
|
|
} else if (strcmp(name1, onnxruntime::CUDA_PINNED) == 0) {
|
|
*out = new OrtMemoryInfo(
|
|
name1, type, OrtDevice(OrtDevice::CPU, OrtDevice::MemType::CUDA_PINNED, static_cast<OrtDevice::DeviceId>(id1)),
|
|
id1, mem_type1);
|
|
} else {
|
|
return OrtApis::CreateStatus(ORT_INVALID_ARGUMENT, "Specified device is not supported.");
|
|
}
|
|
return nullptr;
|
|
}
|
|
|
|
ORT_API(void, OrtApis::ReleaseMemoryInfo, _Frees_ptr_opt_ OrtMemoryInfo* p) { delete p; }
|
|
|
|
ORT_API_STATUS_IMPL(OrtApis::MemoryInfoGetName, _In_ const OrtMemoryInfo* ptr, _Out_ const char** out) {
|
|
*out = ptr->name;
|
|
return nullptr;
|
|
}
|
|
|
|
ORT_API_STATUS_IMPL(OrtApis::MemoryInfoGetId, _In_ const OrtMemoryInfo* ptr, _Out_ int* out) {
|
|
*out = ptr->id;
|
|
return nullptr;
|
|
}
|
|
|
|
ORT_API_STATUS_IMPL(OrtApis::MemoryInfoGetMemType, _In_ const OrtMemoryInfo* ptr, _Out_ OrtMemType* out) {
|
|
*out = ptr->mem_type;
|
|
return nullptr;
|
|
}
|
|
|
|
ORT_API_STATUS_IMPL(OrtApis::MemoryInfoGetType, _In_ const OrtMemoryInfo* ptr, _Out_ OrtAllocatorType* out) {
|
|
*out = ptr->type;
|
|
return nullptr;
|
|
}
|
|
|
|
ORT_API_STATUS_IMPL(OrtApis::CompareMemoryInfo, _In_ const OrtMemoryInfo* info1, _In_ const OrtMemoryInfo* info2,
|
|
_Out_ int* out) {
|
|
*out = (*info1 == *info2) ? 0 : -1;
|
|
return nullptr;
|
|
}
|