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
59 changes: 59 additions & 0 deletions tpu_sync/core/kv_cache_manager_with_transfer.cc
Original file line number Diff line number Diff line change
Expand Up @@ -168,6 +168,65 @@ void KVCacheManagerWithTransfer::InitializeBaseHooks() {
return UnregisterActivePlan(uuid);
};
hooks.get_node_id = [this]() { return node_id(); };
hooks.begin_incoming_push = [this](uint64_t uuid) -> absl::Status {
std::shared_ptr<TransferReceiveSession> recv_session;
std::shared_ptr<ReshardReceiveSession> reshard_session;
{
absl::MutexLock lock(mu_);
if (auto it = active_recv_sessions_.find(uuid);
it != active_recv_sessions_.end()) {
recv_session = it->second;
} else if (auto reshard_it = active_pool_reshard_recvs_.find(uuid);
reshard_it != active_pool_reshard_recvs_.end()) {
reshard_session = reshard_it->second;
} else if (uuid == 0 || base_->HasActivePlan(uuid)) {
return absl::OkStatus();
} else {
return absl::NotFoundError(
absl::StrCat("No active receive session for uuid=", uuid));
}
}
const bool started = recv_session != nullptr
? recv_session->TryBeginRecvOp()
: reshard_session->TryBeginRecvOp();
return started ? absl::OkStatus()
: absl::CancelledError(absl::StrCat(
"Receive session for uuid=", uuid, " is draining"));
};
hooks.end_incoming_push = [this](uint64_t uuid) -> absl::Status {
std::shared_ptr<TransferReceiveSession> recv_session;
std::shared_ptr<ReshardReceiveSession> reshard_session;
{
absl::MutexLock lock(mu_);
if (auto it = active_recv_sessions_.find(uuid);
it != active_recv_sessions_.end()) {
recv_session = it->second;
} else if (auto reshard_it = active_pool_reshard_recvs_.find(uuid);
reshard_it != active_pool_reshard_recvs_.end()) {
reshard_session = reshard_it->second;
} else {
return absl::OkStatus();
}
}
if (recv_session != nullptr) {
recv_session->EndRecvOp();
MaybeUnregisterSettledRecv(uuid, *recv_session);
return recv_session->IsDraining()
? absl::CancelledError(absl::StrCat(
"Receive session for uuid=", uuid, " is draining"))
: absl::OkStatus();
}
reshard_session->EndRecvOp();
uint64_t generation = 0;
if (reshard_session->Done() &&
reshard_session->TakePendingUnregister(&generation)) {
UnregisterSettledPlan(uuid, generation);
}
return reshard_session->IsDraining()
? absl::CancelledError(absl::StrCat(
"ReshardReceiveSession for uuid=", uuid, " is draining"))
: absl::OkStatus();
};
base_->SetTransferEventHooks(std::move(hooks));
}

Expand Down
99 changes: 99 additions & 0 deletions tpu_sync/core/kv_cache_manager_with_transfer_send_drain_test.cc
Original file line number Diff line number Diff line change
Expand Up @@ -206,6 +206,17 @@ class RecvTestManager : public KVCacheManagerWithTransfer {
return it != active_recv_sessions_.end() && !it->second->Done();
}

void FailRecv(uint64_t uuid, absl::Status status) {
std::shared_ptr<TransferReceiveSession> session;
{
absl::MutexLock lock(mu_);
auto it = active_recv_sessions_.find(uuid);
if (it == active_recv_sessions_.end()) return;
session = it->second;
}
session->Finish(status);
}

void BlockH2dDispatch() { block_dispatch_.store(true); }

bool WaitForH2dDispatch(absl::Duration timeout) {
Expand Down Expand Up @@ -742,5 +753,93 @@ TEST(RecvLifecycleTest,
EXPECT_EQ(consumer.free_slots(), kSlots);
}

TEST(RecvLifecycleTest,
IncomingPushLeasePinsStagingDuringWriteAndRejectsWhenDraining) {
RecvTestManager consumer(/*num_layers=*/1, /*timeout_s=*/10.0);
consumer.AddRecv("req_ingress_lease", /*uuid=*/91, /*blocks_per_layer=*/1);
ASSERT_EQ(consumer.free_slots(), kSlots - 1);

EXPECT_THAT(consumer.base()->BeginIncomingPush(/*uuid=*/999),
::absl_testing::StatusIs(absl::StatusCode::kNotFound));

ASSERT_THAT(consumer.base()->BeginIncomingPush(/*uuid=*/91),
::absl_testing::IsOk());

// Failing/timing out mid-write marks the session draining and keeps staging
// pinned until EndIncomingPush finishes.
consumer.FailRecv(/*uuid=*/91, absl::InternalError("simulated failure"));
EXPECT_EQ(consumer.free_slots(), kSlots - 1);

EXPECT_THAT(consumer.base()->BeginIncomingPush(/*uuid=*/91),
::absl_testing::StatusIs(absl::StatusCode::kCancelled));

EXPECT_THAT(consumer.base()->EndIncomingPush(/*uuid=*/91),
::absl_testing::StatusIs(absl::StatusCode::kCancelled));
EXPECT_EQ(consumer.free_slots(), kSlots);
}

TEST(RecvLifecycleTest,
NonSessionTransfersSucceedOverBlockTransportWhileStaleUuidIsRejected) {
constexpr size_t kSliceBytes = 128;
KVCacheManagerWithTransfer sender(
/*num_layers=*/1, /*num_shards=*/1, kSliceBytes,
/*local_port=*/0, /*host_blocks_to_allocate=*/4,
/*parallelism=*/1, /*node_id=*/0, /*local_control_port=*/-1,
/*max_blocks=*/1, /*num_slots=*/1, /*timeout_s=*/10.0);
KVCacheManagerWithTransfer receiver(
/*num_layers=*/1, /*num_shards=*/1, kSliceBytes,
/*local_port=*/0, /*host_blocks_to_allocate=*/4,
/*parallelism=*/1, /*node_id=*/1, /*local_control_port=*/-1,
/*max_blocks=*/1, /*num_slots=*/1, /*timeout_s=*/10.0);

const std::string peer =
absl::StrCat("127.0.0.1:", *receiver.base()->local_port());

// 1. Raw H2H push with dynamic block allocation (op = 1, uuid = 0).
std::memset(sender.base()->GetBlockHostPointer(0, 0, 0), 0x5A, kSliceBytes);
absl::StatusOr<std::vector<int>> dyn_ids =
sender.base()->H2hWriteDirect(peer, /*src_block_ids=*/{0});
ASSERT_THAT(dyn_ids, ::absl_testing::IsOk());
ASSERT_EQ(dyn_ids->size(), 1);
EXPECT_EQ(receiver.base()->GetBlockHostPointer(0, 0, (*dyn_ids)[0])[0], 0x5A);

// 2. Raw H2H push with explicit destination block (op = 6, uuid = 0).
std::memset(sender.base()->GetBlockHostPointer(0, 0, 0), 0xA5, kSliceBytes);
absl::StatusOr<std::vector<int>> exp_ids = sender.base()->H2hWriteDirect(
peer, /*src_block_ids=*/{0}, /*dst_block_ids=*/{2}, /*uuid=*/0);
ASSERT_THAT(exp_ids, ::absl_testing::IsOk());
EXPECT_EQ(receiver.base()->GetBlockHostPointer(0, 0, 2)[0], 0xA5);

// 3. Host-only (MEMORY_TYPE_DRAM) active plan push (op = 6, uuid = 42,
// no TransferReceiveSession in active_recv_sessions_).
::tpu_sync::rpc::StartTransferRequest dram_plan;
dram_plan.set_uuid(42);
dram_plan.set_dst_mem_type(::tpu_sync::rpc::MEMORY_TYPE_DRAM);
dram_plan.set_use_block_chunks(true);
auto* entry = (*dram_plan.mutable_shard_push_schedules())[0].add_entries();
entry->set_dst_peer(peer);
entry->set_dst_shard_idx(0);
entry->set_src_block_id(0);
entry->set_dst_block_id(3);
entry->set_size_bytes(kSliceBytes);
entry->set_count(1);
ASSERT_THAT(receiver.RegisterActivePlan(42, dram_plan, /*is_sender=*/false),
::absl_testing::IsOk());

std::memset(sender.base()->GetBlockHostPointer(0, 0, 0), 0x3C, kSliceBytes);
absl::StatusOr<std::vector<int>> plan_ids = sender.base()->H2hWriteDirect(
peer, /*src_block_ids=*/{0}, /*dst_block_ids=*/{3}, /*uuid=*/42);
ASSERT_THAT(plan_ids, ::absl_testing::IsOk());
EXPECT_EQ(receiver.base()->GetBlockHostPointer(0, 0, 3)[0], 0x3C);

// 4. Once the DRAM plan is unregistered, pushes for uuid = 42 are rejected
// before writing to host memory, preserving the existing bytes at block 3.
ASSERT_THAT(receiver.UnregisterActivePlan(42), ::absl_testing::IsOk());
std::memset(sender.base()->GetBlockHostPointer(0, 0, 0), 0xFF, kSliceBytes);
absl::StatusOr<std::vector<int>> rejected = sender.base()->H2hWriteDirect(
peer, /*src_block_ids=*/{0}, /*dst_block_ids=*/{3}, /*uuid=*/42);
EXPECT_FALSE(rejected.ok());
EXPECT_EQ(receiver.base()->GetBlockHostPointer(0, 0, 3)[0], 0x3C);
}
} // namespace
} // namespace tpu_raiden
12 changes: 8 additions & 4 deletions tpu_sync/core/reshard_receive_session.h
Original file line number Diff line number Diff line change
Expand Up @@ -93,6 +93,14 @@ class ReshardReceiveSession
// times.
void ReleaseStaging();

bool TryBeginRecvOp() {
absl::MutexLock lock(mu_);
if (done_ || draining_) return false;
++in_flight_;
return true;
}
void EndRecvOp();

// Atomically claims any pending settle-unregister request and writes its
// plan generation (0 for pool-reshard plans) to |generation|.
bool TakePendingUnregister(uint64_t* generation);
Expand Down Expand Up @@ -135,10 +143,6 @@ class ReshardReceiveSession
ABSL_EXCLUSIVE_LOCKS_REQUIRED(mu_);
void EndRecvOpLocked() ABSL_EXCLUSIVE_LOCKS_REQUIRED(mu_);

// Decrements the count of in-flight operations; marks the receive done and
// releases its staging resources if it was draining and waiting for this op.
void EndRecvOp();

// Handles completion of |pool_idx|'s H2D upload and finalizes the
// pool-reshard receive when all pools have completed or on error.
void FinishPoolH2d(KVCacheManagerWithTransfer& manager, size_t pool_idx,
Expand Down
12 changes: 8 additions & 4 deletions tpu_sync/core/transfer_receive_session.h
Original file line number Diff line number Diff line change
Expand Up @@ -106,6 +106,14 @@ class TransferReceiveSession

void ReleaseStaging();

bool TryBeginRecvOp() {
absl::MutexLock lock(mu_);
if (done_ || draining_) return false;
++in_flight_;
return true;
}
void EndRecvOp();

// Marks the plan to be unregistered when the receive settles, returning true
// if the receive is still active and will unregister on settle, or false if
// it is already done.
Expand Down Expand Up @@ -207,10 +215,6 @@ class TransferReceiveSession
ABSL_EXCLUSIVE_LOCKS_REQUIRED(mu_);
void EndRecvOpLocked() ABSL_EXCLUSIVE_LOCKS_REQUIRED(mu_);

// Decrements the count of in-flight operations; marks the receive done and
// releases its staging resources if it was draining and waiting for this op.
void EndRecvOp();

bool AllH2dDoneLocked() const ABSL_EXCLUSIVE_LOCKS_REQUIRED(mu_);
bool RecordBlocksReceivedLocked(const std::vector<int>& block_ids,
bool* first_packet,
Expand Down
12 changes: 12 additions & 0 deletions tpu_sync/kv_cache/kv_cache_manager_base.cc
Original file line number Diff line number Diff line change
Expand Up @@ -2289,6 +2289,18 @@ bool KVCacheManagerBase::AcceptsPlanlessExplicitPush(uint64_t uuid) const {
return active_plans_.contains(uuid);
}

absl::Status KVCacheManagerBase::BeginIncomingPush(uint64_t uuid) {
return transfer_hooks_.begin_incoming_push
? transfer_hooks_.begin_incoming_push(uuid)
: absl::OkStatus();
}

absl::Status KVCacheManagerBase::EndIncomingPush(uint64_t uuid) {
return transfer_hooks_.end_incoming_push
? transfer_hooks_.end_incoming_push(uuid)
: absl::OkStatus();
}

absl::StatusOr<std::optional<tpu_raiden::transport::PoolPushProgressSpec>>
KVCacheManagerBase::GetPoolPushProgressSpec(size_t pool_idx,
uint64_t uuid) const {
Expand Down
4 changes: 4 additions & 0 deletions tpu_sync/kv_cache/kv_cache_manager_base.h
Original file line number Diff line number Diff line change
Expand Up @@ -186,6 +186,8 @@ class KVCacheManagerBase : public tpu_raiden::RaidenManagerBase {
register_active_plan;
std::function<absl::Status(uint64_t uuid)> unregister_active_plan;
std::function<int64_t()> get_node_id;
std::function<absl::Status(uint64_t uuid)> begin_incoming_push;
std::function<absl::Status(uint64_t uuid)> end_incoming_push;
};

void SetTransferEventHooks(TransferEventHooks hooks) {
Expand Down Expand Up @@ -533,6 +535,8 @@ class KVCacheManagerBase : public tpu_raiden::RaidenManagerBase {
std::optional<uint64_t> uuid = std::nullopt);

bool AcceptsPlanlessExplicitPush(uint64_t uuid) const override;
absl::Status BeginIncomingPush(uint64_t uuid) override;
absl::Status EndIncomingPush(uint64_t uuid) override;

absl::StatusOr<std::optional<tpu_raiden::transport::PoolPushProgressSpec>>
GetPoolPushProgressSpec(size_t pool_idx, uint64_t uuid) const override;
Expand Down
1 change: 1 addition & 0 deletions tpu_sync/transport/BUILD
Original file line number Diff line number Diff line change
Expand Up @@ -62,6 +62,7 @@ cc_library(
"//tpu_sync/transport/peregrine/src/api:socket_util",
"//tpu_sync/transport/peregrine/src/internal/control:service_cc_grpc",
"@com_google_absl//absl/base:core_headers",
"@com_google_absl//absl/cleanup",
"@com_google_absl//absl/container:flat_hash_map",
"@com_google_absl//absl/log",
"@com_google_absl//absl/log:absl_check",
Expand Down
11 changes: 11 additions & 0 deletions tpu_sync/transport/block_transport.cc
Original file line number Diff line number Diff line change
Expand Up @@ -36,6 +36,7 @@
#include <utility>
#include <vector>

#include "absl/cleanup/cleanup.h"
#include "absl/log/absl_check.h"
#include "absl/log/log.h"
#include "absl/status/status.h"
Expand Down Expand Up @@ -310,6 +311,14 @@ absl::Status BlockTransport::HandleIncomingPush(
std::vector<int> allocated_ids;

std::vector<int> src_block_ids;
ABSL_RETURN_IF_ERROR(block_delegate_->BeginIncomingPush(header.uuid));
bool incoming_push_lease_held = true;
absl::Cleanup end_incoming_push = [&]() {
if (incoming_push_lease_held) {
block_delegate_->EndIncomingPush(header.uuid).IgnoreError();
}
};

if (header.op == 1) {
ABSL_ASSIGN_OR_RETURN(
allocated_ids,
Expand Down Expand Up @@ -378,6 +387,8 @@ absl::Status BlockTransport::HandleIncomingPush(
}
return absl::OkStatus();
}));
incoming_push_lease_held = false;
ABSL_RETURN_IF_ERROR(block_delegate_->EndIncomingPush(header.uuid));

if (total_received_bytes > 0) {
// TODO: Add interface name (e.g. eth0, lo) using
Expand Down
7 changes: 7 additions & 0 deletions tpu_sync/transport/block_transport_delegate.h
Original file line number Diff line number Diff line change
Expand Up @@ -64,6 +64,13 @@ class BlockTransportDelegate : public lib::RawBufferTransportDelegate {
// host mirror.
virtual bool AcceptsPlanlessExplicitPush(uint64_t uuid) const { return true; }

virtual absl::Status BeginIncomingPush(uint64_t uuid) {
return absl::OkStatus();
}
virtual absl::Status EndIncomingPush(uint64_t uuid) {
return absl::OkStatus();
}

// The transport address space is historically one block array per manager
// layer. Explicit pool tables widen that address space to one block array
// per pool without changing the wire's integer index.
Expand Down
Loading