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
13 changes: 12 additions & 1 deletion tpu_sync/api/torch/kv_cache_store.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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],
Expand Down
6 changes: 6 additions & 0 deletions tpu_sync/frameworks/jax/kv_cache_store.pyi
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
import enum
from collections.abc import Callable
from typing import Any, overload

class BlockStatus(enum.Enum):
Expand Down Expand Up @@ -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],
Expand Down
38 changes: 38 additions & 0 deletions tpu_sync/frameworks/jax/tpu_raiden_jax_module.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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<nb::callable>(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<nb::callable>(
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<const std::string> evicted) {
nb::gil_scoped_acquire acquire;
std::vector<nb::bytes> 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,
Expand Down
40 changes: 39 additions & 1 deletion tpu_sync/frameworks/torch/tpu_raiden_torch_module.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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<nb::callable>(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<nb::callable>(
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<const std::string> evicted) {
nb::gil_scoped_acquire acquire;
std::vector<nb::bytes> 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.
Expand Down
1 change: 1 addition & 0 deletions tpu_sync/kv_cache/BUILD
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
19 changes: 16 additions & 3 deletions tpu_sync/kv_cache/kv_cache_store.cc
Original file line number Diff line number Diff line change
Expand Up @@ -23,7 +23,6 @@
#include <memory>
#include <optional>
#include <string>
#include <thread> // NOLINT(build/c++11)
#include <tuple>
#include <utility>
#include <vector>
Expand All @@ -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"
Expand Down Expand Up @@ -1723,6 +1720,22 @@ KVCacheStore::PollRemoteReadStatus() {
std::move(res.pending));
}

void KVCacheStore::SetEvictionCallback(EvictionCallback callback) {
std::shared_ptr<const EvictionCallback> next;
if (callback != nullptr) {
next = std::make_shared<const EvictionCallback>(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<const EvictionCallback> previous;
{
absl::MutexLock lock(mutex_);
previous = std::exchange(eviction_callback_, std::move(next));
}
}

absl::StatusOr<size_t> KVCacheStore::RecoverFromLocalManifest() {
if (!raiden_controller_) {
return absl::FailedPreconditionError(
Expand Down
23 changes: 23 additions & 0 deletions tpu_sync/kv_cache/kv_cache_store.h
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down Expand Up @@ -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<void(absl::Span<const std::string>) 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().
Expand Down Expand Up @@ -569,6 +588,10 @@ class KVCacheStore {
const std::vector<std::string>& 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<const EvictionCallback> eviction_callback_
ABSL_GUARDED_BY(mutex_);
std::vector<std::shared_ptr<KVCacheStoreBackend>> backends_;
std::vector<BackendConfig> backend_configs_;
std::shared_ptr<global_registry::GlobalRegistryClient> registry_client_;
Expand Down
17 changes: 17 additions & 0 deletions tpu_sync/kv_cache/kv_cache_store_test.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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<const std::string>) {});

// 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<int>(0);
controller.SetEvictionCallback(
[state = std::move(owned_state)](absl::Span<const std::string>) {});

// 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();
Expand Down
Loading