diff --git a/tpu_sync/core/kv_cache_manager_with_transfer.cc b/tpu_sync/core/kv_cache_manager_with_transfer.cc index 31d459f10..15d90b260 100644 --- a/tpu_sync/core/kv_cache_manager_with_transfer.cc +++ b/tpu_sync/core/kv_cache_manager_with_transfer.cc @@ -1399,6 +1399,9 @@ absl::Status KVCacheManagerWithTransfer::WaitForPendingWork() { bool recv_pending = false; for (const auto& [uuid, session] : active_recv_sessions_) { (void)uuid; + if (!session->IsDraining() && session->IsReadyToComplete()) { + session->Finish(); + } if (!session->Done()) { recv_pending = true; break; @@ -1407,6 +1410,9 @@ absl::Status KVCacheManagerWithTransfer::WaitForPendingWork() { if (!recv_pending) { for (const auto& [uuid, session] : active_pool_reshard_recvs_) { (void)uuid; + if (!session->IsDraining() && session->IsReadyToComplete()) { + session->Finish(); + } if (session->HasPendingWork()) { recv_pending = true; break; diff --git a/tpu_sync/core/kv_cache_manager_with_transfer_send_drain_test.cc b/tpu_sync/core/kv_cache_manager_with_transfer_send_drain_test.cc index 5212531c1..9cd5760f8 100644 --- a/tpu_sync/core/kv_cache_manager_with_transfer_send_drain_test.cc +++ b/tpu_sync/core/kv_cache_manager_with_transfer_send_drain_test.cc @@ -742,5 +742,37 @@ TEST(RecvLifecycleTest, EXPECT_EQ(consumer.free_slots(), kSlots); } +TEST(RecvLifecycleTest, + WaitForPendingWorkSettlesReadyZeroLayerReceiveWithoutHanging) { + RecvTestManager consumer(/*num_layers=*/0, /*timeout_s=*/1.0); + consumer.AddRecv("req_zero_layers", /*uuid=*/89, /*blocks_per_layer=*/1); + EXPECT_TRUE(consumer.has_recv(89)); + + EXPECT_THAT(consumer.WaitForPendingWork(), ::absl_testing::IsOk()); + EXPECT_FALSE(consumer.has_recv(89)); + EXPECT_EQ(consumer.free_slots(), kSlots); + + Reports reports = consumer.CompleteReadRaw(); + EXPECT_THAT(DoneReceiving(reports), Contains("req_zero_layers")); + EXPECT_THAT(FailedRecving(reports), IsEmpty()); +} + +TEST(RecvLifecycleTest, ZeroBlockActivePlanSettlesImmediately) { + RecvTestManager consumer(/*num_layers=*/1, /*timeout_s=*/1.0); + ::tpu_sync::rpc::StartTransferRequest request; + request.set_uuid(90); + request.set_req_id("req_zero_blocks"); + absl::flat_hash_map + host_block_of; + auto session_or = TransferReceiveSession::CreateFromActivePlan( + consumer.base(), consumer.staging_allocator(), /*uuid=*/90, request, + /*generation=*/1, + std::chrono::steady_clock::now() + std::chrono::seconds(10), + &host_block_of); + ASSERT_THAT(session_or, ::absl_testing::IsOk()); + EXPECT_TRUE((*session_or)->Done()); + EXPECT_THAT((*session_or)->GetStatus(), ::absl_testing::IsOk()); +} + } // namespace } // namespace tpu_raiden diff --git a/tpu_sync/core/transfer_receive_session.cc b/tpu_sync/core/transfer_receive_session.cc index 8b40d11ce..749803a88 100644 --- a/tpu_sync/core/transfer_receive_session.cc +++ b/tpu_sync/core/transfer_receive_session.cc @@ -203,7 +203,7 @@ absl::Status TransferReceiveSession::InitFromActivePlan( h2d_copy_ = TransferSendSession::BuildCoalescedCopySpec(h2d_host_block_ids, h2d_local_block_ids); if (total_blocks_ == 0) { - ReleaseStagingLocked(); + FinishLocked(); } return absl::OkStatus(); } @@ -425,7 +425,8 @@ bool TransferReceiveSession::AllH2dDoneLocked() const { bool TransferReceiveSession::IsReadyToComplete() const { absl::MutexLock lock(mu_); const size_t total_layers = base_ != nullptr ? base_->num_layers() : 0; - return (network_completed_ || + return in_flight_ == 0 && + (network_completed_ || num_completed_layers_ == static_cast(total_layers)) && AllH2dDoneLocked(); } @@ -501,6 +502,8 @@ void TransferReceiveSession::ExecutePullRequest( if (!pull_status.ok()) { self->Finish(pull_status); + } else if (self->base_ != nullptr && self->base_->num_layers() == 0) { + self->Finish(); } }); }