diff --git a/mooncake-transfer-engine/include/transport/kunpeng_transport/ub_context.h b/mooncake-transfer-engine/include/transport/kunpeng_transport/ub_context.h index 07d0cb5711..1f119e7679 100644 --- a/mooncake-transfer-engine/include/transport/kunpeng_transport/ub_context.h +++ b/mooncake-transfer-engine/include/transport/kunpeng_transport/ub_context.h @@ -195,15 +195,18 @@ class UbContext { // Polls one JFC and processes the slices internally: // * Aggregates each completion's jetty depth into jetty_depth_set. - // * Successful slices have markSuccess() called in place and are NOT - // returned (they may be recycled by the submitting thread the + // * Successful normal slices have markSuccess() called in place and are + // NOT returned (they may be recycled by the submitting thread the // moment markSuccess() runs). + // * Successful slices that need post-URMA staging work are returned in + // deferred_success_slices and are NOT marked successful yet. // * Failed slices are returned in failed_slices[0..num_failed-1] for // the caller to apply retry / markFailed. // Returns the total number of completions polled (>= 0), or a negative // error code. virtual int poll(int num_entries, Transport::Slice** failed_slices, int& num_failed, + std::vector& deferred_success_slices, std::unordered_map& jetty_depth_set, int jfc_index = 0) = 0; diff --git a/mooncake-transfer-engine/include/transport/kunpeng_transport/ub_transport.h b/mooncake-transfer-engine/include/transport/kunpeng_transport/ub_transport.h index d38c72083e..1ba12c0815 100644 --- a/mooncake-transfer-engine/include/transport/kunpeng_transport/ub_transport.h +++ b/mooncake-transfer-engine/include/transport/kunpeng_transport/ub_transport.h @@ -14,8 +14,15 @@ #ifndef UB_TRANSPORT_H #define UB_TRANSPORT_H +#include +#include +#include +#include +#include #include #include +#include +#include #include #include "topology.h" #include "transfer_metadata.h" @@ -24,6 +31,7 @@ namespace mooncake { class UbContext; class UbEndPoint; +class UrmaContext; class TransferMetadata; class UbWorkerpool; @@ -36,6 +44,7 @@ enum UB_ENDPOINT_TYPE { URMA_ENDPOINT = 0, OBMM_ENDPOINT = 1 }; class UbTransport : public Transport { friend class UbContext; friend class UbEndPoint; + friend class UrmaContext; friend class UbWorkerPool; public: @@ -82,6 +91,55 @@ class UbTransport : public Transport { private: int allocateLocalSegmentID(); + struct StagingLease { + void* host_ptr = nullptr; + size_t size = 0; + }; + + struct DeviceRegion { + uint64_t addr = 0; + size_t length = 0; + std::string location; + bool remote_accessible = true; + }; + + struct StagingState { + void* original_device_ptr = nullptr; + void* staging_ptr = nullptr; + size_t size = 0; + size_t lease_size = 0; + TransferRequest::OpCode opcode = TransferRequest::WRITE; + std::atomic completed_slices{0}; + uint64_t total_slices = 0; + std::atomic failed{false}; + std::mutex deferred_mutex; + std::vector deferred_success_slices; + }; + + bool stagingEnabled() const; + bool isDevicePointer(const void* ptr) const; + bool isLogicalDeviceRange(const void* ptr, size_t length) const; + int registerLogicalDeviceRegion(void* addr, size_t length, + const std::string& location, + bool remote_accessible); + int unregisterLogicalDeviceRegion(void* addr); + bool copyDeviceToHost(void* dst, const void* src, size_t size) const; + bool copyHostToDevice(void* dst, const void* src, size_t size) const; + Status acquireStaging(size_t size, StagingLease& lease); + void releaseStaging(const StagingLease& lease); + bool isStagedSlice(Slice* slice); + bool shouldDeferSuccess(Slice* slice); + void attachStaging(TransferTask* task, + std::shared_ptr state); + void attachStagingSlice(Slice* slice, + const std::shared_ptr& state); + void detachStagingSlice(Slice* slice); + std::shared_ptr stagingStateForSlice(Slice* slice); + void cleanupStagingForTask(TransferTask* task, + bool detach_all_slices = false); + void onStagedSliceSuccess(Slice* slice); + void onStagedSliceFinalFailure(Slice* slice); + public: int onSetupConnections(const HandShakeDesc& peer_desc, HandShakeDesc& local_desc); @@ -118,6 +176,21 @@ class UbTransport : public Transport { std::shared_ptr local_topology_; UB_ENDPOINT_TYPE endpoint_type_; bool runtime_initialized_ = false; + + mutable std::once_flag staging_config_once_; + mutable bool staging_enabled_ = false; + mutable std::mutex device_region_mutex_; + std::vector device_regions_; + std::mutex staging_pool_mutex_; + void* staging_pool_base_ = nullptr; + size_t staging_pool_size_ = 0; + size_t staging_pool_offset_ = 0; + std::deque> staging_free_list_; + std::mutex staging_state_mutex_; + std::unordered_map> + task_staging_map_; + std::unordered_map> + slice_staging_map_; }; } // namespace mooncake diff --git a/mooncake-transfer-engine/include/transport/kunpeng_transport/urma/urma_endpoint.h b/mooncake-transfer-engine/include/transport/kunpeng_transport/urma/urma_endpoint.h index 846ac640c9..59fb599600 100644 --- a/mooncake-transfer-engine/include/transport/kunpeng_transport/urma/urma_endpoint.h +++ b/mooncake-transfer-engine/include/transport/kunpeng_transport/urma/urma_endpoint.h @@ -62,6 +62,7 @@ class UrmaContext : public UbContext { int doProcessContextEvents() override; void* retrieveRemoteSeg(const std::string& value) override; int poll(int num_entries, Transport::Slice** failed_slices, int& num_failed, + std::vector& deferred_success_slices, std::unordered_map& jetty_depth_set, int jfc_index) override; volatile int* outstandingCount(int jfc_index) override; diff --git a/mooncake-transfer-engine/src/transport/kunpeng_transport/ub_context.cpp b/mooncake-transfer-engine/src/transport/kunpeng_transport/ub_context.cpp index cda2b97cd6..61b8108882 100644 --- a/mooncake-transfer-engine/src/transport/kunpeng_transport/ub_context.cpp +++ b/mooncake-transfer-engine/src/transport/kunpeng_transport/ub_context.cpp @@ -442,18 +442,19 @@ void UbWorkerPool::performPostSend(int thread_id) { void UbWorkerPool::performPoll(int thread_id) { int processed_slice_count = 0; const static size_t kPollCount = 64; - // context_.poll() aggregates each completion's jetty_depth here and - // calls markSuccess() on successful slices in place. Successful slices - // are NOT returned from poll(), so this worker never dereferences them - // after they may have been recycled by the submitting thread. + // context_.poll() aggregates each completion's jetty_depth here. Normal + // successful slices are published in place; staged READ successes are + // returned in deferred_success_slices so H2D can finish before publishing. std::unordered_map jetty_depth_set; std::vector failed_slices; + std::vector deferred_success_slices; for (int jfc_index = thread_id; jfc_index < context_.jfcCount(); jfc_index += kTransferWorkerCount) { UbTransport::Slice* failed[kPollCount]; int num_failed = 0; int nr_poll = context_.poll(kPollCount, failed, num_failed, - jetty_depth_set, jfc_index); + deferred_success_slices, jetty_depth_set, + jfc_index); if (nr_poll < 0) { LOG(ERROR) << "Worker: Failed to poll jetty for complete"; continue; @@ -488,6 +489,10 @@ void UbWorkerPool::performPoll(int thread_id) { for (auto& entry : jetty_depth_set) __sync_fetch_and_sub(entry.first, entry.second); + for (auto& slice : deferred_success_slices) { + context_.engine().onStagedSliceSuccess(slice); + } + // Slices that hit max_retry: final markFailed() after all reads (and the // jetty depth returns above) are done. Failed slices were never published // by poll(), so they remained safe to deref up to this point. @@ -496,7 +501,11 @@ void UbWorkerPool::performPoll(int thread_id) { auto ptr = static_cast(slice->ub.endpoint); context_.deleteEndpointByPtr(ptr); } - slice->markFailed(); + if (context_.engine().isStagedSlice(slice)) { + context_.engine().onStagedSliceFinalFailure(slice); + } else { + slice->markFailed(); + } processed_slice_count_++; } @@ -633,4 +642,4 @@ void UbWorkerPool::monitorWorker() { int UbWorkerPool::doProcessContextEvents() { return context_.doProcessContextEvents(); } -} // namespace mooncake \ No newline at end of file +} // namespace mooncake diff --git a/mooncake-transfer-engine/src/transport/kunpeng_transport/ub_transport.cpp b/mooncake-transfer-engine/src/transport/kunpeng_transport/ub_transport.cpp index 78a5c205dd..6016f828d1 100644 --- a/mooncake-transfer-engine/src/transport/kunpeng_transport/ub_transport.cpp +++ b/mooncake-transfer-engine/src/transport/kunpeng_transport/ub_transport.cpp @@ -13,10 +13,18 @@ // limitations under the License. #include +#include +#include +#include +#include #include +#include +#include #include "config.h" +#include "cuda_alike.h" #include "memory_location.h" #include +#include "ub_allocator.h" #include "transport/kunpeng_transport/ub_context.h" #include "transport/kunpeng_transport/ub_transport.h" #include "transport/kunpeng_transport/ub_endpoint.h" @@ -26,6 +34,37 @@ namespace mooncake { namespace { constexpr uint64_t kNumaAffinitySampleInterval = 10000; +constexpr size_t kDefaultUbStagingPoolSize = 1ull << 30; + +size_t alignUp(size_t value, size_t alignment) { + return (value + alignment - 1) / alignment * alignment; +} + +uint64_t countUbSlices(size_t length, size_t block_size, + size_t fragment_size) { + uint64_t count = 0; + for (uint64_t offset = 0; offset < length; offset += block_size) { + ++count; + if (length - offset <= block_size + fragment_size) break; + } + return count; +} + +size_t getUbStagingPoolSize() { + static const size_t pool_size = [] { + const char* env = std::getenv("MC_UB_STAGING_POOL_SIZE"); + if (!env) return kDefaultUbStagingPoolSize; + try { + size_t value = std::stoull(env); + if (value > 0) return value; + } catch (const std::exception& e) { + LOG(WARNING) << "Invalid MC_UB_STAGING_POOL_SIZE value: " << env + << ", error: " << e.what(); + } + return kDefaultUbStagingPoolSize; + }(); + return pool_size; +} } // namespace UbTransport::UbTransport(UB_ENDPOINT_TYPE endpoint_type) @@ -35,11 +74,323 @@ UbTransport::~UbTransport() { #ifdef CONFIG_USE_BATCH_DESC_SET batch_desc_set_.clear(); #endif + if (staging_pool_base_) { + unregisterLocalMemory(staging_pool_base_, true); + ub_free_memory(staging_pool_base_); + staging_pool_base_ = nullptr; + staging_pool_size_ = 0; + } metadata_->removeSegmentDesc(local_server_name_); batch_desc_set_.clear(); context_list_.clear(); } +bool UbTransport::stagingEnabled() const { + std::call_once(staging_config_once_, [this] { + const char* env = std::getenv("MC_UB_TRANSPORT_CPU_STAGING"); + staging_enabled_ = !env || std::string(env) != "0"; + LOG(INFO) << "UbTransport CPU staging " + << (staging_enabled_ ? "enabled" : "disabled"); + }); + return staging_enabled_; +} + +bool UbTransport::isDevicePointer(const void* ptr) const { + if (!ptr) return false; +#if defined(USE_CUDA) || defined(USE_MUSA) || defined(USE_HIP) || \ + defined(USE_MLU) || defined(USE_MACA) || defined(USE_HYGON) || \ + defined(USE_COREX) + cudaPointerAttributes attributes; + auto status = cudaPointerGetAttributes(&attributes, ptr); + if (status != cudaSuccess) { + return false; + } + return attributes.type == cudaMemoryTypeDevice; +#else + return false; +#endif +} + +bool UbTransport::isLogicalDeviceRange(const void* ptr, size_t length) const { + if (!ptr || length == 0) return false; + uint64_t addr = reinterpret_cast(ptr); + std::lock_guard lock(device_region_mutex_); + for (const auto& region : device_regions_) { + if (addr < region.addr || length > region.length) continue; + if (addr - region.addr <= region.length - length) return true; + } + return false; +} + +int UbTransport::registerLogicalDeviceRegion(void* addr, size_t length, + const std::string& location, + bool remote_accessible) { + if (!addr || length == 0) return ERR_INVALID_ARGUMENT; + uint64_t start = reinterpret_cast(addr); + if (start + length < start) return ERR_INVALID_ARGUMENT; + uint64_t end = start + length; + std::lock_guard lock(device_region_mutex_); + for (const auto& region : device_regions_) { + uint64_t region_start = region.addr; + uint64_t region_end = region.addr + region.length; + if (start < region_end && region_start < end) { + LOG(ERROR) << "UbTransport: logical device region overlaps, addr=" + << addr << " length=" << length; + return ERR_ADDRESS_OVERLAPPED; + } + } + device_regions_.push_back( + DeviceRegion{start, length, location, remote_accessible}); + LOG(INFO) << "UbTransport: registered logical device region addr=" << addr + << " length=" << length << " location=" << location; + return 0; +} + +int UbTransport::unregisterLogicalDeviceRegion(void* addr) { + if (!addr) return ERR_INVALID_ARGUMENT; + uint64_t start = reinterpret_cast(addr); + std::lock_guard lock(device_region_mutex_); + for (auto it = device_regions_.begin(); it != device_regions_.end(); ++it) { + if (it->addr != start) continue; + LOG(INFO) << "UbTransport: unregistered logical device region addr=" + << addr << " length=" << it->length; + device_regions_.erase(it); + return 0; + } + return ERR_ADDRESS_NOT_REGISTERED; +} + +bool UbTransport::copyDeviceToHost(void* dst, const void* src, + size_t size) const { +#if defined(USE_CUDA) || defined(USE_MUSA) || defined(USE_HIP) || \ + defined(USE_MLU) || defined(USE_MACA) || defined(USE_HYGON) || \ + defined(USE_COREX) + return cudaMemcpy(dst, src, size, cudaMemcpyDeviceToHost) == cudaSuccess; +#else + (void)dst; + (void)src; + (void)size; + return false; +#endif +} + +bool UbTransport::copyHostToDevice(void* dst, const void* src, + size_t size) const { +#if defined(USE_CUDA) || defined(USE_MUSA) || defined(USE_HIP) || \ + defined(USE_MLU) || defined(USE_MACA) || defined(USE_HYGON) || \ + defined(USE_COREX) + return cudaMemcpy(dst, src, size, cudaMemcpyHostToDevice) == cudaSuccess; +#else + (void)dst; + (void)src; + (void)size; + return false; +#endif +} + +Status UbTransport::acquireStaging(size_t size, StagingLease& lease) { + if (size == 0) return Status::InvalidArgument("zero-sized UB staging"); + const size_t alignment = 4096; + const size_t aligned_size = alignUp(size, alignment); + + std::lock_guard lock(staging_pool_mutex_); + if (!staging_pool_base_) { + staging_pool_size_ = std::max(getUbStagingPoolSize(), aligned_size); + staging_pool_base_ = ub_allocate_memory(alignment, staging_pool_size_); + if (!staging_pool_base_) { + return Status::Memory("UbTransport: allocate CPU staging pool"); + } + int ret = registerLocalMemory(staging_pool_base_, staging_pool_size_, + kWildcardLocation, false, true); + if (ret) { + ub_free_memory(staging_pool_base_); + staging_pool_base_ = nullptr; + staging_pool_size_ = 0; + return Status::Context( + "UbTransport: register CPU staging pool failed"); + } + LOG(INFO) << "UbTransport: registered CPU staging pool base=" + << staging_pool_base_ << " size=" << staging_pool_size_; + } + + for (auto it = staging_free_list_.begin(); it != staging_free_list_.end(); + ++it) { + if (it->second < aligned_size) continue; + lease.host_ptr = it->first; + lease.size = it->second; + staging_free_list_.erase(it); + return Status::OK(); + } + + if (staging_pool_offset_ + aligned_size > staging_pool_size_) { + return Status::Memory("UbTransport: CPU staging pool exhausted"); + } + lease.host_ptr = static_cast(staging_pool_base_) + + staging_pool_offset_; + lease.size = aligned_size; + staging_pool_offset_ += aligned_size; + return Status::OK(); +} + +void UbTransport::releaseStaging(const StagingLease& lease) { + if (!lease.host_ptr || lease.size == 0) return; + std::lock_guard lock(staging_pool_mutex_); + staging_free_list_.push_back({lease.host_ptr, lease.size}); +} + +void UbTransport::attachStaging(TransferTask* task, + std::shared_ptr state) { + if (!task || !state) return; + std::lock_guard lock(staging_state_mutex_); + task_staging_map_[task] = state; +} + +void UbTransport::attachStagingSlice( + Slice* slice, const std::shared_ptr& state) { + if (!slice || !state) return; + std::lock_guard lock(staging_state_mutex_); + slice_staging_map_[slice] = state; +} + +void UbTransport::detachStagingSlice(Slice* slice) { + if (!slice) return; + std::lock_guard lock(staging_state_mutex_); + slice_staging_map_.erase(slice); +} + +std::shared_ptr UbTransport::stagingStateForSlice( + Slice* slice) { + std::lock_guard lock(staging_state_mutex_); + auto it = slice_staging_map_.find(slice); + return it == slice_staging_map_.end() ? nullptr : it->second; +} + +bool UbTransport::isStagedSlice(Slice* slice) { + return stagingStateForSlice(slice) != nullptr; +} + +bool UbTransport::shouldDeferSuccess(Slice* slice) { + auto state = stagingStateForSlice(slice); + return state && state->opcode == TransferRequest::READ; +} + +void UbTransport::cleanupStagingForTask(TransferTask* task, + bool detach_all_slices) { + std::shared_ptr state; + { + std::lock_guard lock(staging_state_mutex_); + auto task_it = task_staging_map_.find(task); + if (task_it == task_staging_map_.end()) return; + state = task_it->second; + task_staging_map_.erase(task_it); + if (detach_all_slices) { + for (auto* slice : task->slice_list) { + slice_staging_map_.erase(slice); + } + } + } + releaseStaging(StagingLease{state->staging_ptr, state->lease_size}); +} + +void UbTransport::onStagedSliceSuccess(Slice* slice) { + auto state = stagingStateForSlice(slice); + if (!state) { + slice->markSuccess(); + return; + } + + if (state->opcode == TransferRequest::WRITE) { + auto completed = + state->completed_slices.fetch_add(1, std::memory_order_acq_rel) + + 1; + if (completed == state->total_slices) { + cleanupStagingForTask(slice->task); + } + detachStagingSlice(slice); + slice->markSuccess(); + return; + } + + auto completed = + state->completed_slices.fetch_add(1, std::memory_order_acq_rel) + 1; + { + std::lock_guard lock(state->deferred_mutex); + state->deferred_success_slices.push_back(slice); + } + if (completed != state->total_slices) return; + + if (state->failed.load(std::memory_order_acquire)) { + std::vector deferred; + { + std::lock_guard lock(state->deferred_mutex); + deferred.swap(state->deferred_success_slices); + } + cleanupStagingForTask(slice->task); + for (auto* deferred_slice : deferred) { + detachStagingSlice(deferred_slice); + deferred_slice->markFailed(); + } + return; + } + + if (!copyHostToDevice(state->original_device_ptr, state->staging_ptr, + state->size)) { + LOG(ERROR) << "UbTransport: H2D staging copy failed for READ, size=" + << state->size << " dst=" << state->original_device_ptr; + state->failed.store(true, std::memory_order_release); + std::vector deferred; + { + std::lock_guard lock(state->deferred_mutex); + deferred.swap(state->deferred_success_slices); + } + cleanupStagingForTask(slice->task); + for (auto* deferred_slice : deferred) { + detachStagingSlice(deferred_slice); + deferred_slice->markFailed(); + } + return; + } + + std::vector deferred; + { + std::lock_guard lock(state->deferred_mutex); + deferred.swap(state->deferred_success_slices); + } + cleanupStagingForTask(slice->task); + for (auto* deferred_slice : deferred) { + detachStagingSlice(deferred_slice); + deferred_slice->markSuccess(); + } +} + +void UbTransport::onStagedSliceFinalFailure(Slice* slice) { + auto state = stagingStateForSlice(slice); + if (!state) { + slice->markFailed(); + return; + } + state->failed.store(true, std::memory_order_release); + auto completed = + state->completed_slices.fetch_add(1, std::memory_order_acq_rel) + 1; + detachStagingSlice(slice); + if (completed != state->total_slices) { + slice->markFailed(); + return; + } + + std::vector deferred; + if (state->opcode == TransferRequest::READ) { + std::lock_guard lock(state->deferred_mutex); + deferred.swap(state->deferred_success_slices); + } + cleanupStagingForTask(slice->task); + slice->markFailed(); + for (auto* deferred_slice : deferred) { + detachStagingSlice(deferred_slice); + deferred_slice->markFailed(); + } +} + int UbTransport::install(std::string& local_server_name, std::shared_ptr meta, std::shared_ptr topo) { @@ -90,7 +441,17 @@ int UbTransport::registerLocalMemory(void* addr, size_t length, const std::string& name, bool remote_accessible, bool update_metadata) { - (void)remote_accessible; + if (isDevicePointer(addr)) { + if (!stagingEnabled()) { + LOG(ERROR) << "UbTransport: refusing to register device memory " + "while CPU staging is disabled, addr=" + << addr << " length=" << length; + return ERR_INVALID_ARGUMENT; + } + return registerLogicalDeviceRegion(addr, length, name, + remote_accessible); + } + BufferDesc buffer_desc; for (auto& context : context_list_) { int ret = context->registerMemoryRegion((uint64_t)addr, length); @@ -135,6 +496,9 @@ int UbTransport::registerLocalMemory(void* addr, size_t length, } int UbTransport::unregisterLocalMemory(void* addr, bool update_metadata) { + int logical_rc = unregisterLogicalDeviceRegion(addr); + if (logical_rc == 0) return 0; + int rc = metadata_->removeLocalMemoryBuffer(addr, update_metadata); if (rc) return rc; for (auto& context : context_list_) @@ -225,7 +589,6 @@ Status UbTransport::submitTransferTask( const std::vector& task_list) { std::unordered_map, std::vector> slices_to_post; - auto local_segment_desc = metadata_->getSegmentDescByID(LOCAL_SEGMENT_ID); const size_t kBlockSize = globalConfig().slice_size; const int kMaxRetryCount = globalConfig().retry_cnt; const size_t kFragmentSize = globalConfig().fragment_limit; @@ -238,9 +601,58 @@ Status UbTransport::submitTransferTask( nr_slices = 0; assert(task.request); auto& request = *task.request; + void* effective_source = request.source; + std::shared_ptr staging_state; + StagingLease staging_lease; + bool staged_request = false; + + if (isDevicePointer(request.source)) { + if (!stagingEnabled()) { + return Status::InvalidArgument( + "UbTransport: device pointer requires CPU staging"); + } + if (!isLogicalDeviceRange(request.source, request.length)) { + return Status::AddressNotRegistered( + "UbTransport: device pointer is not registered as a " + "logical UB device region, address: " + + std::to_string( + reinterpret_cast(request.source))); + } + auto staging_status = acquireStaging(request.length, staging_lease); + if (!staging_status.ok()) return staging_status; + + staging_state = std::make_shared(); + if (!staging_state) { + releaseStaging(staging_lease); + return Status::Memory("UbTransport: allocate staging state"); + } + staging_state->original_device_ptr = request.source; + staging_state->staging_ptr = staging_lease.host_ptr; + staging_state->size = request.length; + staging_state->lease_size = staging_lease.size; + staging_state->opcode = request.opcode; + staging_state->total_slices = + countUbSlices(request.length, kBlockSize, kFragmentSize); + effective_source = staging_state->staging_ptr; + staged_request = true; + + if (request.opcode == TransferRequest::WRITE && + !copyDeviceToHost(staging_state->staging_ptr, request.source, + request.length)) { + LOG(ERROR) + << "UbTransport: D2H staging copy failed for WRITE, size=" + << request.length << " src=" << request.source; + releaseStaging(staging_lease); + return Status::Memory("UbTransport: D2H staging copy failed"); + } + attachStaging(&task, staging_state); + } + + auto local_segment_desc = + metadata_->getSegmentDescByID(LOCAL_SEGMENT_ID); auto request_buffer_id = -1, request_device_id = -1; - if (selectDevice(local_segment_desc.get(), (uint64_t)request.source, + if (selectDevice(local_segment_desc.get(), (uint64_t)effective_source, request.length, request_buffer_id, request_device_id)) { request_buffer_id = -1; @@ -258,7 +670,7 @@ Status UbTransport::submitTransferTask( slice->dest_rkeys.clear(); bool merge_final_slice = request.length - offset <= kBlockSize + kFragmentSize; - slice->source_addr = (char*)request.source + offset; + slice->source_addr = (char*)effective_source + offset; slice->length = merge_final_slice ? request.length - offset : kBlockSize; slice->opcode = request.opcode; @@ -275,6 +687,7 @@ Status UbTransport::submitTransferTask( slice->ub.src_chip_id = INVALID_CHIP_ID; slice->ub.dst_chip_id = INVALID_CHIP_ID; task.slice_list.push_back(slice); + if (staged_request) attachStagingSlice(slice, staging_state); int buffer_id = -1, device_id = -1, retry_cnt = request.advise_retry_cnt; @@ -314,6 +727,7 @@ Status UbTransport::submitTransferTask( LOG(ERROR) << "UbTransport: Address not registered by any device(s) " << source_addr; + if (staged_request) cleanupStagingForTask(&task, true); return Status::AddressNotRegistered( "UbTransport: not registered by any device(s), " "address: " + @@ -323,6 +737,7 @@ Status UbTransport::submitTransferTask( auto& context = context_list_[device_id]; if (!context->active()) { LOG(ERROR) << "Device " << device_id << " is not active"; + if (staged_request) cleanupStagingForTask(&task, true); return Status::InvalidArgument( "Device " + std::to_string(device_id) + " is not active"); } @@ -360,7 +775,7 @@ Status UbTransport::submitTransferTask( slices_to_post[context].push_back(slice); task.total_bytes += slice->length; __sync_fetch_and_add(&task.slice_count, 1); - if (nr_slices >= kSubmitWatermark) { + if (!staged_request && nr_slices >= kSubmitWatermark) { for (auto& entry : slices_to_post) entry.first->submitPostSend(entry.second); slices_to_post.clear(); @@ -371,6 +786,9 @@ Status UbTransport::submitTransferTask( break; } } + if (staged_request && task.slice_count == 0) { + cleanupStagingForTask(&task, true); + } } for (auto& entry : slices_to_post) entry.first->submitPostSend(entry.second); diff --git a/mooncake-transfer-engine/src/transport/kunpeng_transport/urma/urma_endpoint.cpp b/mooncake-transfer-engine/src/transport/kunpeng_transport/urma/urma_endpoint.cpp index 607453645d..fa69ac3f01 100644 --- a/mooncake-transfer-engine/src/transport/kunpeng_transport/urma/urma_endpoint.cpp +++ b/mooncake-transfer-engine/src/transport/kunpeng_transport/urma/urma_endpoint.cpp @@ -635,6 +635,7 @@ bool UrmaContext::transEidFromString(const std::string& eid_str, int UrmaContext::poll(int num_entries, Transport::Slice** failed_slices, int& num_failed, + std::vector& deferred_success_slices, std::unordered_map& jetty_depth_set, int jfc_index) { num_failed = 0; @@ -662,9 +663,17 @@ int UrmaContext::poll(int num_entries, Transport::Slice** failed_slices, jetty_depth_set[depth] = 1; if (cr[i].status == URMA_CR_SUCCESS) { + if (engine().shouldDeferSuccess(slice)) { + deferred_success_slices.push_back(slice); + continue; + } // Safe to publish here — we are done with this slice and do not // return it to the caller, so no one else will deref it. - slice->markSuccess(); + if (engine().isStagedSlice(slice)) { + engine().onStagedSliceSuccess(slice); + } else { + slice->markSuccess(); + } continue; }