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
21 changes: 19 additions & 2 deletions mooncake-transfer-engine/src/transfer_engine_impl.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -776,19 +776,34 @@ int TransferEngineImpl::registerLocalMemoryBatch(
}

std::vector<MemoryRegion> regions;
std::vector<void*> addr_list;
regions.reserve(buffer_list.size());
addr_list.reserve(buffer_list.size());
for (const auto& buffer : buffer_list) {
regions.push_back({buffer.addr, buffer.length, location, true});
addr_list.push_back(buffer.addr);
}
if (!tryReserveMemoryRegions(regions)) {
LOG(ERROR)
<< "Transfer Engine does not support overlapped memory region";
return ERR_ADDRESS_OVERLAPPED;
}

std::vector<Transport*> attempted_transports;
for (auto transport : multi_transports_->listTransports()) {
attempted_transports.push_back(transport);
int ret = transport->registerLocalMemoryBatch(buffer_list, location);
if (ret < 0) {
if (ret) {
for (auto it = attempted_transports.rbegin();
it != attempted_transports.rend(); ++it) {
int rollback_ret = (*it)->unregisterLocalMemoryBatch(addr_list);
if (rollback_ret != 0 &&
rollback_ret != ERR_ADDRESS_NOT_REGISTERED) {
LOG(WARNING)
<< "Failed to roll back batch registration for "
<< (*it)->getName() << ", ret=" << rollback_ret;
}
}
releaseMemoryRegions(regions);
return ret;
}
Expand All @@ -800,10 +815,12 @@ int TransferEngineImpl::registerLocalMemoryBatch(

int TransferEngineImpl::unregisterLocalMemoryBatch(
const std::vector<void*>& addr_list) {
int first_error = 0;
for (auto transport : multi_transports_->listTransports()) {
int ret = transport->unregisterLocalMemoryBatch(addr_list);
if (ret < 0) return ret;
if (ret && !first_error) first_error = ret;
}
if (first_error) return first_error;

std::unique_lock<std::shared_mutex> lock(mutex_);
for (auto& addr : addr_list) {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -420,16 +420,18 @@ int AscendDirectTransport::unregisterLocalMemoryBatch(
"with addr count: "
<< addr_list.size();

int first_error = 0;
for (void *addr : addr_list) {
int ret = unregisterLocalMemory(addr, false);
if (ret != 0) {
LOG(ERROR) << "Failed to unregister memory in batch, addr: "
<< addr;
return ret;
if (!first_error) first_error = ret;
}
}

// Update metadata once for the entire batch
return metadata_->updateLocalSegmentDesc();
int metadata_ret = metadata_->updateLocalSegmentDesc();
return first_error ? first_error : metadata_ret;
}
} // namespace mooncake
Original file line number Diff line number Diff line change
Expand Up @@ -607,15 +607,23 @@ int HcclTransport::allocateLocalSegmentID() {
int HcclTransport::registerLocalMemoryBatch(
const std::vector<Transport::BufferEntry> &buffer_list,
const std::string &location) {
for (auto &buffer : buffer_list)
registerLocalMemory(buffer.addr, buffer.length, location, true, -1);
for (auto &buffer : buffer_list) {
int ret = registerLocalMemory(buffer.addr, buffer.length, location,
true, false);
if (ret) return ret;
}
return metadata_->updateLocalSegmentDesc();
}

int HcclTransport::unregisterLocalMemoryBatch(
const std::vector<void *> &addr_list) {
for (auto &addr : addr_list) unregisterLocalMemory(addr, -1);
return metadata_->updateLocalSegmentDesc();
int first_error = 0;
for (auto &addr : addr_list) {
int ret = unregisterLocalMemory(addr, false);
if (ret && !first_error) first_error = ret;
}
int metadata_ret = metadata_->updateLocalSegmentDesc();
return first_error ? first_error : metadata_ret;
}

} // namespace mooncake
Original file line number Diff line number Diff line change
Expand Up @@ -208,16 +208,18 @@ int HeterogeneousRdmaTransport::registerLocalMemoryBatch(

int HeterogeneousRdmaTransport::unregisterLocalMemoryBatch(
const std::vector<void *> &addr_list) {
int first_error = 0;
for (auto &addr : addr_list) {
int ret = unregisterLocalMemory(addr, false);
if (ret) {
LOG(ERROR) << "HeterogeneousRdmaTransport "
"unregisterLocalMemoryBatch error, ret: "
<< ret;
return ret;
if (!first_error) first_error = ret;
}
}
return metadata_->updateLocalSegmentDesc();
int metadata_ret = metadata_->updateLocalSegmentDesc();
return first_error ? first_error : metadata_ret;
}

int HeterogeneousRdmaTransport::checkAndCreateStreamCopy() {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -722,7 +722,7 @@ int UBShmemTransport::registerLocalMemoryBatch(
for (auto &buffer : buffer_list) {
int rc = registerLocalMemory(buffer.addr, buffer.length, location, true,
false);
if (rc < 0) {
if (rc) {
return rc;
}
}
Expand All @@ -731,13 +731,13 @@ int UBShmemTransport::registerLocalMemoryBatch(

int UBShmemTransport::unregisterLocalMemoryBatch(
const std::vector<void *> &addr_list) {
int first_error = 0;
for (auto &addr : addr_list) {
int rc = unregisterLocalMemory(addr, false);
if (rc < 0) {
return rc;
}
if (rc && !first_error) first_error = rc;
}
return metadata_->updateLocalSegmentDesc();
int metadata_ret = metadata_->updateLocalSegmentDesc();
return first_error ? first_error : metadata_ret;
}

void *UBShmemTransport::allocatePinnedLocalMemory(size_t size) {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -306,7 +306,7 @@ int BarexTransport::registerLocalMemoryBatch(
if (ret) {
LOG(ERROR) << "BarexTransport: Failed to register memory: addr "
<< buffer.addr << " length " << buffer.length;
return ERR_ADDRESS_NOT_REGISTERED;
return ret;
}
}

Expand All @@ -323,13 +323,17 @@ int BarexTransport::unregisterLocalMemoryBatch(
}));
}

int first_error = 0;
for (size_t i = 0; i < addr_list.size(); ++i) {
if (results[i].get())
int ret = results[i].get();
if (ret) {
LOG(WARNING) << "BarexTransport: Failed to unregister memory: addr "
<< addr_list[i];
if (!first_error) first_error = ret;
}
}

return metadata_->updateLocalSegmentDesc();
int metadata_ret = metadata_->updateLocalSegmentDesc();
return first_error ? first_error : metadata_ret;
}

Status BarexTransport::submitTransfer(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -552,13 +552,17 @@ int CxiTransport::registerLocalMemoryBatch(
}));
}

int first_error = 0;
for (size_t i = 0; i < buffer_list.size(); ++i) {
if (results[i].get()) {
int ret = results[i].get();
if (ret) {
LOG(WARNING) << "CxiTransport: Failed to register memory: addr "
<< buffer_list[i].addr << " length "
<< buffer_list[i].length;
if (!first_error) first_error = ret;
}
}
if (first_error) return first_error;

return metadata_->updateLocalSegmentDesc();
}
Expand All @@ -573,13 +577,17 @@ int CxiTransport::unregisterLocalMemoryBatch(
}));
}

int first_error = 0;
for (size_t i = 0; i < addr_list.size(); ++i) {
if (results[i].get())
int ret = results[i].get();
if (ret) {
LOG(WARNING) << "CxiTransport: Failed to unregister memory: addr "
<< addr_list[i];
if (!first_error) first_error = ret;
}
}

return metadata_->updateLocalSegmentDesc();
int metadata_ret = metadata_->updateLocalSegmentDesc();
return first_error ? first_error : metadata_ret;
}

int CxiTransport::warmupSegment(const std::string& segment_name) {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -272,15 +272,23 @@ int CxlTransport::unregisterLocalMemory(void *addr, bool update_metadata) {
int CxlTransport::registerLocalMemoryBatch(
const std::vector<Transport::BufferEntry> &buffer_list,
const std::string &location) {
for (auto &buffer : buffer_list)
registerLocalMemory(buffer.addr, buffer.length, location, true, false);
for (auto &buffer : buffer_list) {
int ret = registerLocalMemory(buffer.addr, buffer.length, location,
true, false);
if (ret) return ret;
}
return metadata_->updateLocalSegmentDesc();
}

int CxlTransport::unregisterLocalMemoryBatch(
const std::vector<void *> &addr_list) {
for (auto &addr : addr_list) unregisterLocalMemory(addr, false);
return metadata_->updateLocalSegmentDesc();
int first_error = 0;
for (auto &addr : addr_list) {
int ret = unregisterLocalMemory(addr, false);
if (ret && !first_error) first_error = ret;
}
int metadata_ret = metadata_->updateLocalSegmentDesc();
return first_error ? first_error : metadata_ret;
}

Status CxlTransport::getTransferStatus(BatchID batch_id, size_t task_id,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -636,13 +636,17 @@ int EfaTransport::registerLocalMemoryBatch(
}));
}

int first_error = 0;
for (size_t i = 0; i < buffer_list.size(); ++i) {
if (results[i].get()) {
int ret = results[i].get();
if (ret) {
LOG(WARNING) << "EfaTransport: Failed to register memory: addr "
<< buffer_list[i].addr << " length "
<< buffer_list[i].length;
if (!first_error) first_error = ret;
}
}
if (first_error) return first_error;

return metadata_->updateLocalSegmentDesc();
}
Expand All @@ -657,13 +661,17 @@ int EfaTransport::unregisterLocalMemoryBatch(
}));
}

int first_error = 0;
for (size_t i = 0; i < addr_list.size(); ++i) {
if (results[i].get())
int ret = results[i].get();
if (ret) {
LOG(WARNING) << "EfaTransport: Failed to unregister memory: addr "
<< addr_list[i];
if (!first_error) first_error = ret;
}
}

return metadata_->updateLocalSegmentDesc();
int metadata_ret = metadata_->updateLocalSegmentDesc();
return first_error ? first_error : metadata_ret;
}

int EfaTransport::warmupSegment(const std::string& segment_name) {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -823,18 +823,20 @@ int HipTransport::registerLocalMemoryBatch(
for (auto& buffer : buffer_list) {
int rc = registerLocalMemory(buffer.addr, buffer.length, location, true,
false);
if (rc < 0) return rc;
if (rc) return rc;
}
return metadata_->updateLocalSegmentDesc();
}

int HipTransport::unregisterLocalMemoryBatch(
const std::vector<void*>& addr_list) {
int first_error = 0;
for (auto& addr : addr_list) {
int rc = unregisterLocalMemory(addr, false);
if (rc < 0) return rc;
if (rc && !first_error) first_error = rc;
}
return metadata_->updateLocalSegmentDesc();
int metadata_ret = metadata_->updateLocalSegmentDesc();
return first_error ? first_error : metadata_ret;
}

void* HipTransport::allocatePinnedLocalMemory(size_t size) {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -762,15 +762,23 @@ int IntraNodeNvlinkTransport::relocateSharedMemoryAddress(uint64_t &dest_addr,
int IntraNodeNvlinkTransport::registerLocalMemoryBatch(
const std::vector<Transport::BufferEntry> &buffer_list,
const std::string &location) {
for (auto &buffer : buffer_list)
registerLocalMemory(buffer.addr, buffer.length, location, true, false);
for (auto &buffer : buffer_list) {
int ret = registerLocalMemory(buffer.addr, buffer.length, location,
true, false);
if (ret) return ret;
}
return metadata_->updateLocalSegmentDesc();
}

int IntraNodeNvlinkTransport::unregisterLocalMemoryBatch(
const std::vector<void *> &addr_list) {
for (auto &addr : addr_list) unregisterLocalMemory(addr, false);
return metadata_->updateLocalSegmentDesc();
int first_error = 0;
for (auto &addr : addr_list) {
int ret = unregisterLocalMemory(addr, false);
if (ret && !first_error) first_error = ret;
}
int metadata_ret = metadata_->updateLocalSegmentDesc();
return first_error ? first_error : metadata_ret;
}

void *IntraNodeNvlinkTransport::allocatePinnedLocalMemory(size_t size) {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -142,13 +142,17 @@ int UbTransport::registerLocalMemoryBatch(
}));
}

int first_error = 0;
for (size_t i = 0; i < buffer_list.size(); ++i) {
if (results[i].get()) {
int ret = results[i].get();
if (ret) {
LOG(WARNING) << "UbTransport: Failed to register memory: addr "
<< buffer_list[i].addr << " length "
<< buffer_list[i].length;
if (!first_error) first_error = ret;
}
}
if (first_error) return first_error;

return metadata_->updateLocalSegmentDesc();
}
Expand All @@ -164,13 +168,17 @@ int UbTransport::unregisterLocalMemoryBatch(
}));
}

int first_error = 0;
for (size_t i = 0; i < addr_list.size(); ++i) {
if (results[i].get())
int ret = results[i].get();
if (ret) {
LOG(WARNING) << "UbTransport: Failed to unregister memory: addr "
<< addr_list[i];
if (!first_error) first_error = ret;
}
}

return metadata_->updateLocalSegmentDesc();
int metadata_ret = metadata_->updateLocalSegmentDesc();
return first_error ? first_error : metadata_ret;
}

Status UbTransport::submitTransfer(
Expand Down
Loading
Loading