From 7da96ff13f98cddc62f29c96f6c48f4973ce887f Mon Sep 17 00:00:00 2001 From: Googler Date: Wed, 23 Sep 2026 13:27:26 -0700 Subject: [PATCH] [tpu-raiden] Add SetEvictionCallback API skeleton for KV cache eviction tracking. PiperOrigin-RevId: 986966283 --- tpu_sync/api/torch/kv_cache_store.py | 13 +++++- tpu_sync/frameworks/jax/kv_cache_store.pyi | 6 +++ .../frameworks/jax/tpu_raiden_jax_module.cc | 38 ++++++++++++++++++ .../torch/tpu_raiden_torch_module.cc | 40 ++++++++++++++++++- tpu_sync/kv_cache/BUILD | 1 + tpu_sync/kv_cache/kv_cache_store.cc | 19 +++++++-- tpu_sync/kv_cache/kv_cache_store.h | 23 +++++++++++ tpu_sync/kv_cache/kv_cache_store_test.cc | 17 ++++++++ 8 files changed, 152 insertions(+), 5 deletions(-) diff --git a/tpu_sync/api/torch/kv_cache_store.py b/tpu_sync/api/torch/kv_cache_store.py index 448338c4a..5e1439b98 100644 --- a/tpu_sync/api/torch/kv_cache_store.py +++ b/tpu_sync/api/torch/kv_cache_store.py @@ -14,7 +14,7 @@ """Raiden KV Cache Store API for PyTorch.""" -from typing import Any +from typing import Any, Callable from tpu_sync.api.torch import torch_tpu_common_loader @@ -497,6 +497,17 @@ def poll_load_status(self) -> tuple[list[bytes], list[bytes], list[bytes]]: """ return self._impl.poll_load_status() + def set_eviction_callback( + self, callback: Callable[[list[bytes]], None] | None + ) -> None: + """Registers a callback invoked upon host LRU cache eviction. + + Args: + callback: Function called with the list of evicted block hashes, or None + to unregister. + """ + self._impl.set_eviction_callback(callback) + def read_remote( self, block_hashes: list[bytes], diff --git a/tpu_sync/frameworks/jax/kv_cache_store.pyi b/tpu_sync/frameworks/jax/kv_cache_store.pyi index 1b4682372..3ad437726 100644 --- a/tpu_sync/frameworks/jax/kv_cache_store.pyi +++ b/tpu_sync/frameworks/jax/kv_cache_store.pyi @@ -1,4 +1,5 @@ import enum +from collections.abc import Callable from typing import Any, overload class BlockStatus(enum.Enum): @@ -169,6 +170,11 @@ class KVCacheStore: def poll_load_status(self) -> tuple[list[bytes], list[bytes], list[bytes]]: """Polls status of asynchronous Load operations.""" ... + def set_eviction_callback( + self, callback: Callable[[list[bytes]], None] | None + ) -> None: + """Registers a callback invoked upon host LRU cache eviction.""" + ... def read_remote( self, block_hashes: list[bytes], diff --git a/tpu_sync/frameworks/jax/tpu_raiden_jax_module.cc b/tpu_sync/frameworks/jax/tpu_raiden_jax_module.cc index 362bb5339..ee262cc53 100644 --- a/tpu_sync/frameworks/jax/tpu_raiden_jax_module.cc +++ b/tpu_sync/frameworks/jax/tpu_raiden_jax_module.cc @@ -820,6 +820,44 @@ NB_MODULE(_tpu_raiden_jax, m) { } return std::make_tuple(py_done, py_failed, py_pending); }) + .def( + "set_eviction_callback", + [](tpu_raiden::kv_cache::KVCacheStoreWrapper& self, + nb::object py_cb) { + if (py_cb.is_none()) { + self->SetEvictionCallback(nullptr); + return; + } + nb::callable cb_callable = nb::cast(py_cb); + // The callback outlives this call and is released by whichever + // thread drops it last -- possibly the background eviction + // thread, which holds no GIL. Own the handle through a deleter + // that takes the GIL, so its refcount is only touched with the + // GIL held; copying the shared_ptr itself is GIL-free. + auto callable = std::shared_ptr( + new nb::callable(std::move(cb_callable)), + [](nb::callable* held) { + nb::gil_scoped_acquire acquire; + delete held; + }); + self->SetEvictionCallback( + [callable = std::move(callable)]( + absl::Span evicted) { + nb::gil_scoped_acquire acquire; + std::vector py_evicted; + py_evicted.reserve(evicted.size()); + for (const auto& h : evicted) { + py_evicted.push_back(nb::bytes(h.data(), h.size())); + } + try { + (*callable)(py_evicted); + } catch (...) { + // Ignore Python exceptions to avoid terminating host + // process. + } + }); + }, + nb::arg("callback")) .def( "read_remote", [](tpu_raiden::kv_cache::KVCacheStoreWrapper& self, diff --git a/tpu_sync/frameworks/torch/tpu_raiden_torch_module.cc b/tpu_sync/frameworks/torch/tpu_raiden_torch_module.cc index fac8df50e..262885e7c 100644 --- a/tpu_sync/frameworks/torch/tpu_raiden_torch_module.cc +++ b/tpu_sync/frameworks/torch/tpu_raiden_torch_module.cc @@ -1134,7 +1134,45 @@ NB_MODULE(_tpu_raiden_torch, m) { py_pending.push_back(nb::bytes(h.data(), h.size())); } return std::make_tuple(py_done, py_failed, py_pending); - }); + }) + .def( + "set_eviction_callback", + [](tpu_raiden::kv_cache::KVCacheStoreWrapper& self, + nb::object py_cb) { + if (py_cb.is_none()) { + self->SetEvictionCallback(nullptr); + return; + } + nb::callable cb_callable = nb::cast(py_cb); + // The callback outlives this call and is released by whichever + // thread drops it last -- possibly the background eviction + // thread, which holds no GIL. Own the handle through a deleter + // that takes the GIL, so its refcount is only touched with the + // GIL held; copying the shared_ptr itself is GIL-free. + auto callable = std::shared_ptr( + new nb::callable(std::move(cb_callable)), + [](nb::callable* held) { + nb::gil_scoped_acquire acquire; + delete held; + }); + self->SetEvictionCallback( + [callable = std::move(callable)]( + absl::Span evicted) { + nb::gil_scoped_acquire acquire; + std::vector py_evicted; + py_evicted.reserve(evicted.size()); + for (const auto& h : evicted) { + py_evicted.push_back(nb::bytes(h.data(), h.size())); + } + try { + (*callable)(py_evicted); + } catch (...) { + // Ignore Python exceptions to avoid terminating host + // process. + } + }); + }, + nb::arg("callback")); // C++-owned reshard client. The facade-compatible surface is provided by // tpu_raiden/api/torch/reshard_client.py on top of this binding. diff --git a/tpu_sync/kv_cache/BUILD b/tpu_sync/kv_cache/BUILD index ef22aba2e..e3b536869 100644 --- a/tpu_sync/kv_cache/BUILD +++ b/tpu_sync/kv_cache/BUILD @@ -561,6 +561,7 @@ cc_library( "@com_google_absl//absl/cleanup", "@com_google_absl//absl/container:flat_hash_map", "@com_google_absl//absl/container:flat_hash_set", + "@com_google_absl//absl/functional:any_invocable", "@com_google_absl//absl/log", "@com_google_absl//absl/memory", "@com_google_absl//absl/status", diff --git a/tpu_sync/kv_cache/kv_cache_store.cc b/tpu_sync/kv_cache/kv_cache_store.cc index a9aa2f5a8..8c63f96fc 100644 --- a/tpu_sync/kv_cache/kv_cache_store.cc +++ b/tpu_sync/kv_cache/kv_cache_store.cc @@ -23,7 +23,6 @@ #include #include #include -#include // NOLINT(build/c++11) #include #include #include @@ -48,10 +47,8 @@ #include "grpcpp/security/credentials.h" #include "xla/tsl/concurrency/future.h" #include "tpu_sync/common/raiden_id.h" -#include "tpu_sync/core/buffer.h" #include "tpu_sync/core/controller/raiden_controller.h" #include "tpu_sync/core/host_memory_allocator.h" -#include "tpu_sync/kv_cache/completion_executor.h" #include "tpu_sync/kv_cache/global_registry/global_registry_client.h" #include "tpu_sync/kv_cache/host_offload_backend.h" #include "tpu_sync/kv_cache/kv_cache_metadata.h" @@ -1723,6 +1720,22 @@ KVCacheStore::PollRemoteReadStatus() { std::move(res.pending)); } +void KVCacheStore::SetEvictionCallback(EvictionCallback callback) { + std::shared_ptr next; + if (callback != nullptr) { + next = std::make_shared(std::move(callback)); + } + // Keeps the outgoing callback alive past the unlock. Destroying it here + // rather than under mutex_ matters for language bindings: the previous + // target may own a Python handle whose release needs the GIL, and this + // thread is the one that owns it. + std::shared_ptr previous; + { + absl::MutexLock lock(mutex_); + previous = std::exchange(eviction_callback_, std::move(next)); + } +} + absl::StatusOr KVCacheStore::RecoverFromLocalManifest() { if (!raiden_controller_) { return absl::FailedPreconditionError( diff --git a/tpu_sync/kv_cache/kv_cache_store.h b/tpu_sync/kv_cache/kv_cache_store.h index f30e3103b..26e1bdfa7 100644 --- a/tpu_sync/kv_cache/kv_cache_store.h +++ b/tpu_sync/kv_cache/kv_cache_store.h @@ -29,6 +29,7 @@ #include "absl/base/thread_annotations.h" #include "absl/container/flat_hash_map.h" #include "absl/container/flat_hash_set.h" +#include "absl/functional/any_invocable.h" #include "absl/status/status.h" #include "absl/status/statusor.h" #include "absl/strings/string_view.h" @@ -441,6 +442,24 @@ class KVCacheStore { using PollLoadStatusResult = BlockTracker::StatusResult; PollLoadStatusResult PollLoadStatus(); + // Registers a callback invoked whenever blocks are evicted from the host + // LRU cache, or clears it when `callback` is empty. + // + // The callback is move-only, so it may own its state outright (e.g. capture + // a std::unique_ptr) rather than being forced into shared or copyable + // captures. + // + // Threading contract: + // - It runs on whichever thread performed the eviction, which includes the + // background sweep thread, not just the caller of Evict(). + // - It is invoked with no KVCacheStore lock held, so it may call back into + // the store, including SetEvictionCallback to unregister itself. + // - Registration and invocation may race: a callback replaced concurrently + // with an eviction can still see one final notification. + using EvictionCallback = + absl::AnyInvocable) const>; + void SetEvictionCallback(EvictionCallback callback); + // Launches an async receiver-initiated read of REMOTE blocks from their // owning peers straight into local HBM. Returns as soon as the reads are // issued; poll with PollRemoteReadStatus(). @@ -569,6 +588,10 @@ class KVCacheStore { const std::vector& batch); mutable absl::Mutex mutex_; + // Held by pointer so an eviction can snapshot it under mutex_ and invoke it + // after unlocking: the callback must never run while a store lock is held. + std::shared_ptr eviction_callback_ + ABSL_GUARDED_BY(mutex_); std::vector> backends_; std::vector backend_configs_; std::shared_ptr registry_client_; diff --git a/tpu_sync/kv_cache/kv_cache_store_test.cc b/tpu_sync/kv_cache/kv_cache_store_test.cc index c5205e124..b068877f3 100644 --- a/tpu_sync/kv_cache/kv_cache_store_test.cc +++ b/tpu_sync/kv_cache/kv_cache_store_test.cc @@ -270,6 +270,23 @@ TEST(KVCacheStoreTest, EvictionTracking) { EXPECT_EQ(PeekLookup(controller, {"103"})->size(), 1); } +TEST(KVCacheStoreTest, SetEvictionCallback) { + KVCacheStore controller(2, "", {}, /*num_shards=*/1, + /*shard_size_bytes=*/512, + /*store_server_ip=*/"127.0.0.1"); + + controller.SetEvictionCallback([](absl::Span) {}); + + // The callback is move-only, so it can own its state outright instead of + // being forced into a copyable capture. + auto owned_state = std::make_unique(0); + controller.SetEvictionCallback( + [state = std::move(owned_state)](absl::Span) {}); + + // Passing an empty callback unregisters whatever was registered before. + controller.SetEvictionCallback(nullptr); +} + TEST(KVCacheStoreTest, GlobalLookupFallback) { // 1. Start a local registry server auto reg_server = global_registry::CreateTestGlobalRegistryServer();