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
35 changes: 29 additions & 6 deletions src/bridge/infini/rt.hpp
Original file line number Diff line number Diff line change
@@ -1,6 +1,5 @@
#pragma once

#include "infinicore.h"
#include "infinirt.h"

#include <infini/rt.h>
Expand All @@ -16,34 +15,58 @@ inline infiniStatus_t translate(::infini::rt::runtime::Error error) {
}
}

inline ::infini::rt::Device::Type translate(infiniDevice_t device) {
inline ::infini::rt::Device::Type translate_to(infiniDevice_t device) {
switch (device) {
case INFINI_DEVICE_CPU:
return ::infini::rt::Device::Type::kCpu;
case INFINI_DEVICE_NVIDIA:
return ::infini::rt::Device::Type::kNvidia;
case INFINI_DEVICE_CAMBRICON:
return ::infini::rt::Device::Type::kCambricon;
case INFINI_DEVICE_ASCEND:
return ::infini::rt::Device::Type::kAscend;
case INFINI_DEVICE_METAX:
return ::infini::rt::Device::Type::kMetax;
case INFINI_DEVICE_MOORE:
return ::infini::rt::Device::Type::kMoore;
case INFINI_DEVICE_ILUVATAR:
return ::infini::rt::Device::Type::kIluvatar;
case INFINI_DEVICE_HYGON:
return ::infini::rt::Device::Type::kHygon;
default:
return ::infini::rt::Device::Type::kCount;
}
}

inline infiniDevice_t translate(::infini::rt::Device::Type device) {
inline infiniDevice_t translate_from(::infini::rt::Device::Type device) {
switch (device) {
case ::infini::rt::Device::Type::kCpu:
return INFINI_DEVICE_CPU;
case ::infini::rt::Device::Type::kNvidia:
return INFINI_DEVICE_NVIDIA;
case ::infini::rt::Device::Type::kCambricon:
return INFINI_DEVICE_CAMBRICON;
case ::infini::rt::Device::Type::kAscend:
return INFINI_DEVICE_ASCEND;
case ::infini::rt::Device::Type::kMetax:
return INFINI_DEVICE_METAX;
case ::infini::rt::Device::Type::kMoore:
return INFINI_DEVICE_MOORE;
case ::infini::rt::Device::Type::kIluvatar:
return INFINI_DEVICE_ILUVATAR;
case ::infini::rt::Device::Type::kHygon:
return INFINI_DEVICE_HYGON;
default:
return INFINI_DEVICE_TYPE_COUNT;
}
}

inline ::infini::rt::runtime::Stream translate(infinirtStream_t stream) {
inline ::infini::rt::runtime::Stream translate_to(infinirtStream_t stream) {
return reinterpret_cast<::infini::rt::runtime::Stream>(stream);
}

inline ::infini::rt::runtime::Stream *translate(infinirtStream_t *stream) {
return reinterpret_cast<::infini::rt::runtime::Stream *>(stream);
inline infinirtStream_t translate_from(::infini::rt::runtime::Stream stream) {
return reinterpret_cast<infinirtStream_t>(stream);
}

} // namespace infinicore::bridge::infini::rt
9 changes: 7 additions & 2 deletions src/infinicore/context/allocators/stream_ordered_allocator.cc
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,10 @@ std::byte *StreamOrderedAllocator::allocate(size_t size) {
}
void *ptr = nullptr;
if (device_.getType() != Device::Type::CPU) {
INFINICORE_CHECK_ERROR(bridge::infini::rt::translate(infini::rt::runtime::MallocAsync(&ptr, size, bridge::infini::rt::translate(context::getStream()))));
INFINICORE_CHECK_ERROR(bridge::infini::rt::translate(infini::rt::runtime::MallocAsync(
&ptr,
size,
bridge::infini::rt::translate_to(context::getStream()))));
} else {
INFINICORE_CHECK_ERROR(bridge::infini::rt::translate(infini::rt::runtime::Malloc(&ptr, size)));
}
Expand All @@ -25,7 +28,9 @@ void StreamOrderedAllocator::deallocate(std::byte *ptr) {
return;
}
if (device_.getType() != Device::Type::CPU) {
INFINICORE_CHECK_ERROR(bridge::infini::rt::translate(infini::rt::runtime::FreeAsync(ptr, bridge::infini::rt::translate(context::getStream()))));
INFINICORE_CHECK_ERROR(bridge::infini::rt::translate(infini::rt::runtime::FreeAsync(
ptr,
bridge::infini::rt::translate_to(context::getStream()))));
} else {
INFINICORE_CHECK_ERROR(bridge::infini::rt::translate(infini::rt::runtime::Free(ptr)));
}
Expand Down
14 changes: 12 additions & 2 deletions src/infinicore/context/context_impl.cc
Original file line number Diff line number Diff line change
Expand Up @@ -64,7 +64,7 @@ ContextImpl &ContextImpl::singleton() {
ContextImpl::ContextImpl() {
std::vector<int> device_counter(static_cast<size_t>(Device::Type::COUNT), 0);

infini::rt::set_runtime_device_type(bridge::infini::rt::translate(INFINI_DEVICE_CPU));
infini::rt::set_runtime_device_type(bridge::infini::rt::translate_to(INFINI_DEVICE_CPU));
INFINICORE_CHECK_ERROR(bridge::infini::rt::translate(infini::rt::runtime::GetDeviceCount(&device_counter[static_cast<int>(Device::Type::CPU)])));

runtime_table_[static_cast<int>(Device::Type::CPU)].resize(device_counter[static_cast<int>(Device::Type::CPU)]);
Expand All @@ -73,7 +73,7 @@ ContextImpl::ContextImpl() {
}

if constexpr (infini::rt::DeviceEnabled<infini::rt::Device::Type::kNvidia>::value) {
infini::rt::set_runtime_device_type(bridge::infini::rt::translate(INFINI_DEVICE_NVIDIA));
infini::rt::set_runtime_device_type(bridge::infini::rt::translate_to(INFINI_DEVICE_NVIDIA));
INFINICORE_CHECK_ERROR(bridge::infini::rt::translate(infini::rt::runtime::GetDeviceCount(&device_counter[static_cast<int>(Device::Type::NVIDIA)])));
runtime_table_[static_cast<int>(Device::Type::NVIDIA)].resize(device_counter[static_cast<int>(Device::Type::NVIDIA)]);
if (device_counter[static_cast<int>(Device::Type::NVIDIA)] > 0) {
Expand All @@ -82,6 +82,16 @@ ContextImpl::ContextImpl() {
}
}

if constexpr (infini::rt::DeviceEnabled<infini::rt::Device::Type::kAscend>::value) {
infini::rt::set_runtime_device_type(bridge::infini::rt::translate_to(INFINI_DEVICE_ASCEND));
INFINICORE_CHECK_ERROR(bridge::infini::rt::translate(infini::rt::runtime::GetDeviceCount(&device_counter[static_cast<int>(Device::Type::ASCEND)])));
runtime_table_[static_cast<int>(Device::Type::ASCEND)].resize(device_counter[static_cast<int>(Device::Type::ASCEND)]);
if (device_counter[static_cast<int>(Device::Type::ASCEND)] > 0) {
runtime_table_[static_cast<int>(Device::Type::ASCEND)][0] = std::unique_ptr<Runtime>(new Runtime(Device(Device::Type::ASCEND, 0)));
current_runtime_ = runtime_table_[static_cast<int>(Device::Type::ASCEND)][0].get();
}
}

Comment thread
voltjia marked this conversation as resolved.
if (current_runtime_ == nullptr && !runtime_table_[static_cast<int>(Device::Type::CPU)].empty()) {
current_runtime_ = runtime_table_[static_cast<int>(Device::Type::CPU)][0].get();
}
Expand Down
20 changes: 11 additions & 9 deletions src/infinicore/context/runtime/runtime.cc
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,9 @@
namespace infinicore {
Runtime::Runtime(Device device) : device_(device), graph_manager_(std::make_unique<graph::GraphManager>()) {
activate();
INFINICORE_CHECK_ERROR(bridge::infini::rt::translate(infini::rt::runtime::StreamCreate(bridge::infini::rt::translate(&stream_))));
infini::rt::runtime::Stream stream = nullptr;
INFINICORE_CHECK_ERROR(bridge::infini::rt::translate(infini::rt::runtime::StreamCreate(&stream)));
stream_ = bridge::infini::rt::translate_from(stream);
INFINICORE_CHECK_ERROR(infiniopCreateHandle(&infiniop_handle_));
if (device_.getType() == Device::Type::CPU) {
device_memory_allocator_ = std::make_unique<PinnableBlockAllocator>(device);
Expand All @@ -27,11 +29,11 @@ Runtime::~Runtime() {
}
device_memory_allocator_.reset();
infiniopDestroyHandle(infiniop_handle_);
(void)infini::rt::runtime::StreamDestroy(bridge::infini::rt::translate(stream_));
(void)infini::rt::runtime::StreamDestroy(bridge::infini::rt::translate_to(stream_));
}

Runtime *Runtime::activate() {
auto rt_device = bridge::infini::rt::translate(static_cast<infiniDevice_t>(device_.getType()));
auto rt_device = bridge::infini::rt::translate_to(static_cast<infiniDevice_t>(device_.getType()));
INFINICORE_ASSERT(rt_device != infini::rt::Device::Type::kCount);
infini::rt::set_runtime_device_type(rt_device);
INFINICORE_CHECK_ERROR(bridge::infini::rt::translate(infini::rt::runtime::SetDevice(static_cast<int>(device_.getIndex()))));
Expand All @@ -51,7 +53,7 @@ infiniopHandle_t Runtime::infiniopHandle() const {
}

void Runtime::syncStream() {
INFINICORE_CHECK_ERROR(bridge::infini::rt::translate(infini::rt::runtime::StreamSynchronize(bridge::infini::rt::translate(stream_))));
INFINICORE_CHECK_ERROR(bridge::infini::rt::translate(infini::rt::runtime::StreamSynchronize(bridge::infini::rt::translate_to(stream_))));
}

void Runtime::syncDevice() {
Expand Down Expand Up @@ -108,7 +110,7 @@ std::shared_ptr<Memory> Runtime::reinstantiateBlob(std::shared_ptr<Memory> blob)

void Runtime::memcpyH2D(void *dst, const void *src, size_t size, bool async) {
if (async && device_.getType() != Device::Type::CPU) {
INFINICORE_CHECK_ERROR(bridge::infini::rt::translate(infini::rt::runtime::MemcpyAsync(dst, src, size, infini::rt::runtime::kMemcpyHostToDevice, bridge::infini::rt::translate(stream_))));
INFINICORE_CHECK_ERROR(bridge::infini::rt::translate(infini::rt::runtime::MemcpyAsync(dst, src, size, infini::rt::runtime::kMemcpyHostToDevice, bridge::infini::rt::translate_to(stream_))));
} else {
INFINICORE_CHECK_ERROR(bridge::infini::rt::translate(infini::rt::runtime::Memcpy(dst, src, size, infini::rt::runtime::kMemcpyHostToDevice)));
}
Expand All @@ -120,7 +122,7 @@ void Runtime::memcpyD2H(void *dst, const void *src, size_t size) {

void Runtime::memcpyD2D(void *dst, const void *src, size_t size, bool async) {
if (async && device_.getType() != Device::Type::CPU) {
INFINICORE_CHECK_ERROR(bridge::infini::rt::translate(infini::rt::runtime::MemcpyAsync(dst, src, size, infini::rt::runtime::kMemcpyDeviceToDevice, bridge::infini::rt::translate(stream_))));
INFINICORE_CHECK_ERROR(bridge::infini::rt::translate(infini::rt::runtime::MemcpyAsync(dst, src, size, infini::rt::runtime::kMemcpyDeviceToDevice, bridge::infini::rt::translate_to(stream_))));
} else {
INFINICORE_CHECK_ERROR(bridge::infini::rt::translate(infini::rt::runtime::Memcpy(dst, src, size, infini::rt::runtime::kMemcpyDeviceToDevice)));
}
Expand All @@ -132,7 +134,7 @@ void Runtime::setDeviceMemory(void *ptr, int value, size_t count) {

void Runtime::setDeviceMemoryAsync(void *ptr, int value, size_t count, infinirtStream_t stream) {
if (device_.getType() != Device::Type::CPU) {
INFINICORE_CHECK_ERROR(bridge::infini::rt::translate(infini::rt::runtime::MemsetAsync(ptr, value, count, bridge::infini::rt::translate(stream))));
INFINICORE_CHECK_ERROR(bridge::infini::rt::translate(infini::rt::runtime::MemsetAsync(ptr, value, count, bridge::infini::rt::translate_to(stream))));
} else {
INFINICORE_CHECK_ERROR(bridge::infini::rt::translate(infini::rt::runtime::Memset(ptr, value, count)));
}
Expand All @@ -155,7 +157,7 @@ void Runtime::recordEvent(infinirtEvent_t event, infinirtStream_t stream) {
if (stream == nullptr) {
stream = stream_;
}
INFINICORE_CHECK_ERROR(bridge::infini::rt::translate(infini::rt::runtime::EventRecord(event, bridge::infini::rt::translate(stream))));
INFINICORE_CHECK_ERROR(bridge::infini::rt::translate(infini::rt::runtime::EventRecord(event, bridge::infini::rt::translate_to(stream))));
}

bool Runtime::queryEvent(infinirtEvent_t event) {
Expand All @@ -181,7 +183,7 @@ void Runtime::streamWaitEvent(infinirtStream_t stream, infinirtEvent_t event) {
if (stream == nullptr) {
stream = stream_;
}
INFINICORE_CHECK_ERROR(bridge::infini::rt::translate(infini::rt::runtime::StreamWaitEvent(bridge::infini::rt::translate(stream), event, 0)));
INFINICORE_CHECK_ERROR(bridge::infini::rt::translate(infini::rt::runtime::StreamWaitEvent(bridge::infini::rt::translate_to(stream), event, 0)));
}

bool Runtime::isGraphRecording() const {
Expand Down
84 changes: 56 additions & 28 deletions src/infinicore/graph/graph.cc
Original file line number Diff line number Diff line change
@@ -1,14 +1,19 @@
#include "graph_manager.hpp"

#include "../../bridge/infini/rt.hpp"
#include "../utils.hpp"
#include "infinicore/context/context.hpp"

#ifdef USE_INFINIRT_GRAPH
#include <infinirt.h>
#include <infini/rt.h>
#endif

namespace infinicore::graph {

#ifdef USE_INFINIRT_GRAPH
namespace rt_runtime = ::infini::rt::runtime;
#endif

/* =========================
* GraphTensor
* ========================= */
Expand Down Expand Up @@ -36,26 +41,25 @@ DispatchableGraphOperator::~DispatchableGraphOperator() {

#ifdef USE_INFINIRT_GRAPH
struct Graph::DeviceGraph {
infinirtGraph_t graph;
infinirtGraphExec_t exec;
infinirtGraphNode_t node;
std::vector<char> log_buffer;

DeviceGraph() {
log_buffer.resize(4 * 1024);
}
rt_runtime::Graph graph = nullptr;
rt_runtime::GraphExec exec = nullptr;
rt_runtime::Stream stream = nullptr;
::infini::rt::Device::Type device_type = ::infini::rt::Device::Type::kCount;
int device_index = 0;

~DeviceGraph() {
if (exec) {
infinirtGraphExecDestroy(exec);
(void)rt_runtime::GraphExecDestroy(exec);
}
if (graph) {
infinirtGraphDestroy(graph);
(void)rt_runtime::GraphDestroy(graph);
Comment thread
voltjia marked this conversation as resolved.
}
}

void launch() {
INFINICORE_CHECK_ERROR(infinirtGraphLuanch(exec, context::getStream()));
::infini::rt::set_runtime_device_type(device_type);
INFINICORE_CHECK_ERROR(bridge::infini::rt::translate(rt_runtime::SetDevice(device_index)));
INFINICORE_CHECK_ERROR(bridge::infini::rt::translate(rt_runtime::GraphLaunch(exec, stream)));
}
};
#else
Expand Down Expand Up @@ -85,42 +89,66 @@ void Graph::instantiate() {
#ifdef USE_INFINIRT_GRAPH
// Reset device graph
device_graph_ = std::make_unique<DeviceGraph>();
auto current_device = context::getDevice();
device_graph_->device_type = bridge::infini::rt::translate_to(static_cast<infiniDevice_t>(current_device.getType()));
device_graph_->device_index = static_cast<int>(current_device.getIndex());
device_graph_->stream = bridge::infini::rt::translate_to(context::getStream());
if (device_graph_->device_type == ::infini::rt::Device::Type::kCount) {
spdlog::warn("InfiniRT graph runtime does not support the current device. Falling back to eager execution.");
device_graph_.reset();
return;
}
::infini::rt::set_runtime_device_type(device_graph_->device_type);
auto set_device_status = bridge::infini::rt::translate(rt_runtime::SetDevice(device_graph_->device_index));
if (set_device_status != INFINI_STATUS_SUCCESS) {
spdlog::warn("InfiniRT graph runtime failed to select the current device. Falling back to eager execution.");
device_graph_.reset();
return;
}

// warmup
for (size_t iter = 0; iter < 5; ++iter) {
this->run();
}
infinicore::context::syncStream();

if (infinirtStreamBeginCapture(
context::getStream(),
INFINIRT_STREAM_CAPTURE_MODE_RELAXED)
!= INFINI_STATUS_SUCCESS) {
auto begin_status = bridge::infini::rt::translate(rt_runtime::StreamBeginCapture(
device_graph_->stream,
rt_runtime::StreamCaptureMode::kStreamCaptureModeRelaxed));
if (begin_status != INFINI_STATUS_SUCCESS) {
spdlog::warn("Fail to begin device graph capture.");
device_graph_.reset();
return;
}

// Run and record
this->run();

if (infinirtStreamEndCapture(
context::getStream(),
&device_graph_.get()->graph)
!= INFINI_STATUS_SUCCESS) {
auto end_status = bridge::infini::rt::translate(rt_runtime::StreamEndCapture(
device_graph_->stream,
&device_graph_->graph));
if (end_status != INFINI_STATUS_SUCCESS) {
spdlog::warn("Fail to end device graph capture.");
device_graph_.reset();
return;
}

if (infinirtGraphInstantiate(
&device_graph_.get()->exec,
device_graph_.get()->graph,
&device_graph_.get()->node,
device_graph_.get()->log_buffer.data(),
device_graph_.get()->log_buffer.size())
!= INFINI_STATUS_SUCCESS) {
auto instantiate_status = bridge::infini::rt::translate(rt_runtime::GraphInstantiate(
&device_graph_->exec,
device_graph_->graph));
if (instantiate_status != INFINI_STATUS_SUCCESS) {
static bool warned_once = false;
if (!warned_once) {
warned_once = true;
spdlog::warn("Fail to instantiate device graph: {}", std::string(device_graph_.get()->log_buffer.data()));
spdlog::warn("Fail to instantiate device graph.");
}
device_graph_.reset();
return;
}
static bool logged_once = false;
if (!logged_once) {
logged_once = true;
spdlog::info("Using InfiniRT C++ graph runtime API for graph capture and replay.");
}
#endif
}
Expand Down
6 changes: 4 additions & 2 deletions src/infinicore/tensor/copy.cc
Original file line number Diff line number Diff line change
Expand Up @@ -43,11 +43,13 @@ void TensorImpl::copy_from(Tensor src) {
}
} else if (src->device().getType() == Device::Type::CPU) {
context::setDevice(this->device());
// Keep Ascend H2D synchronous because copy_from does not retain the host source.
const bool async = this->device().getType() != Device::Type::ASCEND;
if (this->is_contiguous()) {
context::memcpyH2D(this->data(), src->data(), copy_size);
context::memcpyH2D(this->data(), src->data(), copy_size, async);
} else {
auto local_src = Tensor::empty(this->shape(), this->dtype(), this->device());
context::memcpyH2D(local_src->data(), src->data(), copy_size);
context::memcpyH2D(local_src->data(), src->data(), copy_size, async);
op::rearrange_(Tensor(const_cast<TensorImpl *>(this)->shared_from_this()), local_src);
}
}
Expand Down
4 changes: 2 additions & 2 deletions src/infiniop/devices/handle.cc
Original file line number Diff line number Diff line change
Expand Up @@ -25,7 +25,7 @@
#endif

__INFINI_C infiniStatus_t infiniopSetRuntimeDevice(infiniDevice_t device, int device_id) {
auto rt_device = infinicore::bridge::infini::rt::translate(device);
auto rt_device = infinicore::bridge::infini::rt::translate_to(device);
if (rt_device == infini::rt::Device::Type::kCount) {
return INFINI_STATUS_DEVICE_TYPE_NOT_SUPPORTED;
}
Expand All @@ -40,7 +40,7 @@ __INFINI_C infiniStatus_t infiniopCreateHandle(infiniopHandle_t *handle_ptr) {
return INFINI_STATUS_NULL_POINTER;
}

infiniDevice_t device = infinicore::bridge::infini::rt::translate(infini::rt::runtime_device_type());
infiniDevice_t device = infinicore::bridge::infini::rt::translate_from(infini::rt::runtime_device_type());
if (device == INFINI_DEVICE_TYPE_COUNT) {
return INFINI_STATUS_DEVICE_TYPE_NOT_SUPPORTED;
}
Expand Down
Loading
Loading