Skip to content
Open
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
57 changes: 56 additions & 1 deletion mooncake-transfer-engine/include/transport/transport.h
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,7 @@
#include <functional>
#include <mutex>
#include <condition_variable>
#include <utility>

#include "common/base/status.h"
#include "transfer_metadata.h"
Expand Down Expand Up @@ -284,6 +285,30 @@ class Transport {
};

struct TransferTask {
TransferTask() = default;

TransferTask(const TransferTask &) = delete;
TransferTask &operator=(const TransferTask &) = delete;

TransferTask(TransferTask &&other) noexcept {
moveFrom(std::move(other));
}

TransferTask &operator=(TransferTask &&other) noexcept {
if (this != &other) {
releaseSlices();
moveFrom(std::move(other));
}
return *this;
}

~TransferTask() { releaseSlices(); }

void setRequest(const TransferRequest &transfer_request) {
request_storage = transfer_request;
request = &request_storage;
}

volatile uint64_t slice_count = 0;
volatile uint64_t success_slice_count = 0;
volatile uint64_t failed_slice_count = 0;
Expand All @@ -306,6 +331,8 @@ class Transport {
volatile uint64_t completed_slice_count = 0;
#endif

TransferRequest request_storage{};

// record the origin request
#ifdef USE_ASCEND_HETEROGENEOUS
// need to modify the request's source address, changing it from an NPU
Expand All @@ -316,9 +343,37 @@ class Transport {
#endif
// record the slice list for freeing objects
std::vector<Slice *> slice_list;
~TransferTask() {

private:
void releaseSlices() {
for (auto &slice : slice_list)
Transport::getSliceCache().deallocate(slice);
slice_list.clear();
}

void moveFrom(TransferTask &&other) noexcept {
slice_count = other.slice_count;
success_slice_count = other.success_slice_count;
failed_slice_count = other.failed_slice_count;
transferred_bytes = other.transferred_bytes;
is_finished = other.is_finished;
total_bytes = other.total_bytes;
batch_id = other.batch_id;
transport_ = other.transport_;

#ifdef WITH_METRICS
start_time = other.start_time;
#endif

#ifdef USE_EVENT_DRIVEN_COMPLETION
completed_slice_count = other.completed_slice_count;
#endif

request_storage = other.request_storage;
request = other.request == nullptr ? nullptr : &request_storage;
slice_list = std::move(other.slice_list);
other.slice_list.clear();
other.request = nullptr;
}
};

Expand Down
12 changes: 2 additions & 10 deletions mooncake-transfer-engine/src/multi_transport.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -125,11 +125,7 @@ Status MultiTransport::submitTransfer(
auto& task = batch_desc.task_list[task_id];
task.batch_id = batch_id;
task.transport_ = transport;
#ifdef USE_ASCEND_HETEROGENEOUS
task.request = const_cast<Transport::TransferRequest*>(&request);
#else
task.request = &request;
#endif
task.setRequest(request);
++task_id;
submit_tasks[transport].push_back(&task);
}
Expand Down Expand Up @@ -167,11 +163,7 @@ Status MultiTransport::mp_submitTransfer(
assert(transport);
auto& task = batch_desc.task_list[task_id];
task.batch_id = batch_id;
#ifdef USE_ASCEND_HETEROGENEOUS
task.request = const_cast<Transport::TransferRequest*>(&request);
#else
task.request = &request;
#endif
task.setRequest(request);
++task_id;
submit_tasks[transport].push_back(&task);
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -422,9 +422,11 @@ void UbWorkerPool::performPoll(int thread_id) {
redispatch_counter_++;
}
} else {
// slice->markSuccess();
processed_slice_count++;
success_nr_polls++;
// Publish completion only after this worker has finished all
// local accounting that reads fields from the slice.
slice->markSuccess();
}
}
if (nr_poll)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -572,7 +572,7 @@ int UrmaContext::poll(int num_entries, Transport::Slice** slices,
}
slices[i] = slice;
if (cr[i].status == URMA_CR_SUCCESS) {
slice->markSuccess();
slice->status = Transport::Slice::SUCCESS;
continue;
}
if (cr[i].status != URMA_CR_WR_FLUSH_ERR ||
Expand Down
28 changes: 27 additions & 1 deletion mooncake-transfer-engine/tests/transport_uint_test.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -171,9 +171,35 @@ TEST_F(TransportTest, ReadEmptyFile) {

close(fd);
}

TEST_F(TransportTest, TransferTaskKeepsOwnedRequestCopy) {
Transport::TransferTask task;
Transport::TransferRequest request;
request.opcode = Transport::TransferRequest::READ;
request.source = reinterpret_cast<void *>(0x1234);
request.target_id = 7;
request.target_offset = 4096;
request.length = 8192;
request.advise_retry_cnt = 3;

task.setRequest(request);

ASSERT_NE(task.request, nullptr);
EXPECT_NE(task.request, &request);

request.source = reinterpret_cast<void *>(0x4321);
request.length = 16;

EXPECT_EQ(task.request->opcode, Transport::TransferRequest::READ);
EXPECT_EQ(task.request->source, reinterpret_cast<void *>(0x1234));
EXPECT_EQ(task.request->target_id, 7);
EXPECT_EQ(task.request->target_offset, 4096);
EXPECT_EQ(task.request->length, 8192);
EXPECT_EQ(task.request->advise_retry_cnt, 3);
}
} // namespace mooncake

int main(int argc, char** argv) {
::testing::InitGoogleTest(&argc, argv);
return RUN_ALL_TESTS();
}
}
Loading