From 7db538fa46e63a7eaabce6de78b26e1edb8588db Mon Sep 17 00:00:00 2001 From: gasoonjia Date: Fri, 2 Oct 2026 12:06:36 -0700 Subject: [PATCH] Update [ghstack-poisoned] --- backends/cuda/runtime/cuda_kv_cache.cpp | 388 +++++++++++++++++- backends/cuda/runtime/cuda_kv_cache.h | 11 + .../cuda/runtime/test/test_cuda_kv_cache.cpp | 270 ++++++++++++ extension/cuda/runtime_api.h | 12 + 4 files changed, 669 insertions(+), 12 deletions(-) diff --git a/backends/cuda/runtime/cuda_kv_cache.cpp b/backends/cuda/runtime/cuda_kv_cache.cpp index a0ade97b598..d38cd15585d 100644 --- a/backends/cuda/runtime/cuda_kv_cache.cpp +++ b/backends/cuda/runtime/cuda_kv_cache.cpp @@ -8,14 +8,20 @@ #include +#include #include #include #include +#include +#include +#include +#include #include #include #include #include +#include #include #include @@ -235,6 +241,330 @@ class CudaSequenceKVCache final : public cache::SequenceCache, Error error_{Error::Ok}; }; + +// The window each layer attends over, 0 = its whole history. Layers agreeing +// on a window read one mask, as the lowering pass declares it. +int layer_window(const cache::LayerGeometry& layer) { + return is_ring(layer) ? layer.policy.window : 0; +} + +std::string cell_mask_fqn(int window) { + return "__et_offgraph_kv_mask_w" + std::to_string(window); +} + +// Side buffers of the cell layout, in the order CudaCellCache addresses them: +// cells, read_len, then one mask per distinct window. Names, dtypes and shapes +// match LowerOffGraphKVPass's cell layout. +constexpr size_t kCellsBuffer = 0; +constexpr size_t kReadLenBuffer = 1; +constexpr size_t kFirstMaskBuffer = 2; + +// Many sequences over one pool of per-token cells on one CUDA device. +// +// The neutral base owns the cell table: placement, ownership, and the verbs a +// runner drives between forwards. The pool owns the bytes. Before each +// forward, prepare_step places the declared tokens and writes the placement +// (cells), the extent (read_len) and one visibility mask per window into +// buffers the program reads at fixed addresses -- so a captured CUDA graph +// serves any mix of sequences, and only growth forces a recapture. +class CudaCellCache final : public cache::CellCache, public CudaKVCache { + public: + CudaCellCache( + const cache::CacheGeometry& geometry, + const cache::CacheConfig& cfg, + slimc10::ScalarType storage_dtype, + std::vector windows) + : cache::CellCache(geometry, cfg), + max_write_(*cfg.max_write), + windows_(std::move(windows)), + pool_( + pool_layers(geometry, cfg), + side_buffers(cfg, windows_), + storage_dtype, + cfg.initial_capacity) { + // The first layer of each window: placing through it memoizes that + // window's step for the forward. + for (const int window : windows_) { + for (size_t index = 0; index < geometry.layers.size(); ++index) { + if (layer_window(geometry.layers[index]) == window) { + window_layers_.push_back(static_cast(index)); + break; + } + } + } + } + + CudaCellCache(const CudaCellCache&) = delete; + CudaCellCache& operator=(const CudaCellCache&) = delete; + CudaCellCache(CudaCellCache&&) = delete; + CudaCellCache& operator=(CudaCellCache&&) = delete; + + ~CudaCellCache() override = default; + + // -- CacheControl / BatchControl, serialized against the delegate. Keeps the + // storage: a reset reuses the grown pools. + + void clear() override { + std::lock_guard guard(mutex_); + cache::CellCache::clear(); + step_seq_ids_.clear(); + } + + bool declare_step(const std::vector& seq_ids) override { + std::lock_guard guard(mutex_); + if (!cache::CellCache::declare_step(seq_ids)) { + step_seq_ids_.clear(); + return false; + } + step_seq_ids_ = seq_ids; + return true; + } + + std::optional seq_new() override { + std::lock_guard guard(mutex_); + return cache::CellCache::seq_new(); + } + + std::optional seq_clone(int32_t src, std::optional upto) + override { + std::lock_guard guard(mutex_); + return cache::CellCache::seq_clone(src, upto); + } + + bool seq_rm(int32_t seq_id) override { + std::lock_guard guard(mutex_); + return cache::CellCache::seq_rm(seq_id); + } + + bool rewind(int32_t seq_id, int position) override { + std::lock_guard guard(mutex_); + return cache::CellCache::rewind(seq_id, position); + } + + // -- CudaKVCache. + + runtime::Result note_handle(CudaDelegateHandle* handle) override { + std::lock_guard guard(mutex_); + auto serves = pool_.note_handle(handle); + if (!serves.ok()) { + error_ = serves.error(); + } + return serves; + } + + void forget_handle(CudaDelegateHandle* handle) override { + std::lock_guard guard(mutex_); + pool_.forget_handle(handle); + } + + Error rebind_for_execute(CudaDelegateHandle* handle) override { + std::lock_guard guard(mutex_); + if (!pool_.serves(handle)) { + return Error::Ok; + } + ET_CHECK_OK_OR_RETURN_ERROR(error_); + ET_CHECK_OR_RETURN_ERROR( + pool_.allocated(), + InvalidState, + "offgraph_kv: prepare_step must run before execute"); + return pool_.bind(handle); + } + + // Places the declared step and writes it where the program reads it. Runs on + // the host between executes, so its copies never land inside a capture. + Error prepare_step(int64_t write_length, cudaStream_t stream) override { + std::lock_guard guard(mutex_); + ET_CHECK_OK_OR_RETURN_ERROR(error_); + if (!validated_) { + const Error valid = pool_.validate(); + if (valid != Error::Ok) { + error_ = valid; + return valid; + } + validated_ = true; + } + const int width = static_cast(write_length); + ET_CHECK_OR_RETURN_ERROR( + write_length > 0 && write_length <= max_write_, + InvalidArgument, + "offgraph_kv: a %lld-token step is outside [1, %d]", + static_cast(write_length), + max_write_); + ET_CHECK_OR_RETURN_ERROR( + static_cast(width) == step_seq_ids_.size(), + InvalidState, + "offgraph_kv: the step carries %d tokens, declare_step declared %zu", + width, + step_seq_ids_.size()); + + // Every declared token continues its sequence, so the positions follow + // from the declaration; nothing is read back from the device. + positions_.resize(width); + std::vector next(kMaxSeqs, -1); + for (int i = 0; i < width; ++i) { + const int32_t seq_id = step_seq_ids_[i]; + if (next[seq_id] < 0) { + next[seq_id] = cache::CellCache::pos(seq_id); + } + positions_[i] = next[seq_id]++; + } + + // Grow before placing, to where placement could reach at most, so a + // failed allocation leaves the table as it was. Lowest-free placement + // never takes a cell past used_end + width. + const int live = used_end(); + ET_CHECK_OK_OR_RETURN_ERROR(pool_.prepare( + std::min(capacity(), static_cast(live) + width), + live, + stream)); + + const cache::CellStep* first = nullptr; + for (size_t index = 0; index < windows_.size(); ++index) { + const cache::CellStep* step = + place_step(window_layers_[index], positions_.data(), width); + ET_CHECK_OR_RETURN_ERROR( + step != nullptr, + InvalidArgument, + "offgraph_kv: the declared step does not place"); + if (first == nullptr) { + first = step; + ET_CHECK_OK_OR_RETURN_ERROR(write_placement(*step, stream)); + } + ET_CHECK_OK_OR_RETURN_ERROR( + write_mask(kFirstMaskBuffer + index, *step, stream)); + } + step_seq_ids_.clear(); + return Error::Ok; + } + + // The cells were claimed when the step was placed; the forward only filled + // them. + Error commit_step(int64_t write_length) override { + std::lock_guard guard(mutex_); + ET_CHECK_OR_RETURN_ERROR( + write_length > 0, InvalidArgument, "write length must be positive"); + ET_CHECK_OK_OR_RETURN_ERROR(error_); + return pool_.mark_step_done(); + } + + // logical_length is the read extent: every occupied cell is below it. + OffGraphKVMetrics metrics() const override { + std::lock_guard guard(mutex_); + OffGraphKVMetrics metrics; + metrics.logical_length = used_end(); + metrics.flat_capacity = pool_.rows(); + metrics.growth_count = pool_.growth_count(); + metrics.allocated_bytes = pool_.allocated_bytes(); + return metrics; + } + + protected: + void* face(cache::FaceId id) override { + if (void* p = cache::CellCache::face(id)) { + return p; + } + return cache::expose(this, id); + } + + private: + static std::vector pool_layers( + const cache::CacheGeometry& geometry, + const cache::CacheConfig& cfg) { + std::vector layers; + layers.reserve(geometry.layers.size()); + for (const cache::LayerGeometry& layer : geometry.layers) { + layers.push_back(CudaKVPool::Layer{ + layer.n_kv_heads, + layer.head_dim, + static_cast(cfg.capacity), + /*growable=*/true}); + } + return layers; + } + + static std::vector side_buffers( + const cache::CacheConfig& cfg, + const std::vector& windows) { + const int64_t max_write = *cfg.max_write; + std::vector buffers{ + {"__et_offgraph_kv_cells", slimc10::ScalarType::Long, {max_write}}, + {"__et_offgraph_kv_read_len", slimc10::ScalarType::Long, {1}}, + }; + for (const int window : windows) { + buffers.push_back( + {cell_mask_fqn(window), + slimc10::ScalarType::Bool, + {1, 1, max_write, static_cast(cfg.capacity)}}); + } + return buffers; + } + + // Pageable sources: the copy has consumed them when the call returns, so + // the host vectors may be reused by the next step at once. + Error write_placement(const cache::CellStep& step, cudaStream_t stream) { + staged_cells_.assign(step.cells.begin(), step.cells.end()); + const int64_t read_len = step.read_len; + for (const auto& [index, src, bytes] : + {std::tuple{ + kCellsBuffer, + static_cast(staged_cells_.data()), + staged_cells_.size() * sizeof(int64_t)}, + std::tuple{ + kReadLenBuffer, + static_cast(&read_len), + sizeof(int64_t)}}) { + const cudaError_t error = cudaMemcpyAsync( + pool_.side_buffer(index), src, bytes, cudaMemcpyHostToDevice, stream); + ET_CHECK_OR_RETURN_ERROR( + error == cudaSuccess, + Internal, + "offgraph_kv: cannot write the step placement: %s", + cudaGetErrorString(error)); + } + return Error::Ok; + } + + // Rows [0, length) over columns [0, read_len); the program bounds its sweep + // by read_len, so columns past it keep whatever an earlier step left. + Error write_mask(size_t buffer, const cache::CellStep& step, cudaStream_t stream) { + if (step.read_len == 0) { + return Error::Ok; + } + const size_t row = static_cast(capacity()); + const cudaError_t error = cudaMemcpy2DAsync( + pool_.side_buffer(buffer), + row, + step.mask_bits.data(), + static_cast(step.read_len), + static_cast(step.read_len), + static_cast(step.length), + cudaMemcpyHostToDevice, + stream); + ET_CHECK_OR_RETURN_ERROR( + error == cudaSuccess, + Internal, + "offgraph_kv: cannot write the step mask: %s", + cudaGetErrorString(error)); + return Error::Ok; + } + + // Guards everything below and the base's table. The runner drives the verbs + // and the delegate the steps, from whichever threads they run on. Recursive + // because the base's verbs call each other virtually: seq_clone takes its + // new id through seq_new. + mutable std::recursive_mutex mutex_; + + const int max_write_; + const std::vector windows_; + std::vector window_layers_; + CudaKVPool pool_; + std::vector step_seq_ids_; + std::vector positions_; + std::vector staged_cells_; + bool validated_{false}; + Error error_{Error::Ok}; +}; + } // namespace std::shared_ptr make_cuda_sequence_kv_cache( @@ -263,6 +593,33 @@ std::shared_ptr make_cuda_sequence_kv_cache( return std::make_shared(geometry, cfg, storage_dtype); } +std::shared_ptr make_cuda_cell_kv_cache( + const cache::CacheGeometry& geometry, + const cache::CacheConfig& cfg) { + slimc10::ScalarType storage_dtype = slimc10::ScalarType::BFloat16; + if (!cache::valid(geometry, cfg) || cfg.initial_capacity <= 0 || + cfg.initial_capacity > cfg.capacity || + !storage_dtype_of(cfg.kv_dtype, storage_dtype)) { + ET_LOG(Error, "offgraph_kv: invalid cache geometry or config"); + return nullptr; + } + // The step buffers are declared [max_write] and [max_write, max_cells]; a + // cache without the program's widest step cannot address them. + if (!cfg.max_write || *cfg.max_write <= 0 || *cfg.max_write > cfg.capacity) { + ET_LOG(Error, "offgraph_kv: the cell layout requires max_write"); + return nullptr; + } + std::set windows; + for (const cache::LayerGeometry& layer : geometry.layers) { + windows.insert(layer_window(layer)); + } + return std::make_shared( + geometry, + cfg, + storage_dtype, + std::vector(windows.begin(), windows.end())); +} + runtime::Error attach_offgraph_kv_cache( CudaDelegateHandle& handle, const char* cache_key, @@ -304,18 +661,25 @@ namespace { // naming any CUDA type. Lives beside the cache so a build without the LLM // extension, which drops this file, registers nothing. const bool cuda_cache_builders_registered = [] { - const auto error = cache::CacheFactory::global().register_builder( - kCudaBackendId, - cache::kind::kSingle, - [](const cache::CacheGeometry& geometry, const cache::CacheConfig& cfg) { - return make_cuda_sequence_kv_cache(geometry, cfg); - }); - if (error != runtime::Error::Ok) { - ET_LOG( - Error, - "Failed to register cache builder for %s:%s", - kCudaBackendId, - cache::kind::kSingle); + const struct { + const char* kind; + cache::CacheBuilder builder; + } builders[] = { + {cache::kind::kSingle, make_cuda_sequence_kv_cache}, + {cache::kind::kBatchedCell, make_cuda_cell_kv_cache}, + // The cell layout is the batch layout this backend serves. + {cache::kind::kBatched, make_cuda_cell_kv_cache}, + }; + for (const auto& entry : builders) { + const auto error = cache::CacheFactory::global().register_builder( + kCudaBackendId, entry.kind, entry.builder); + if (error != runtime::Error::Ok) { + ET_LOG( + Error, + "Failed to register cache builder for %s:%s", + kCudaBackendId, + entry.kind); + } } return true; }(); diff --git a/backends/cuda/runtime/cuda_kv_cache.h b/backends/cuda/runtime/cuda_kv_cache.h index 322839f3ddb..3ccbade7c80 100644 --- a/backends/cuda/runtime/cuda_kv_cache.h +++ b/backends/cuda/runtime/cuda_kv_cache.h @@ -206,4 +206,15 @@ std::shared_ptr make_cuda_sequence_kv_cache( const cache::CacheGeometry& geometry, const cache::CacheConfig& cfg); +// Builder for cache::kind::kBatchedCell (and kBatched): many sequences over +// one pool of per-token cells, for a program lowered in the cell layout. +// +// cfg.capacity is the program's max_cells and cfg.max_write its widest step; +// both fix the shapes the program declared for its step buffers, so they must +// match the export exactly. Pools start at cfg.initial_capacity rows and grow +// geometrically up to cfg.capacity. +std::shared_ptr make_cuda_cell_kv_cache( + const cache::CacheGeometry& geometry, + const cache::CacheConfig& cfg); + } // namespace executorch::backends::cuda diff --git a/backends/cuda/runtime/test/test_cuda_kv_cache.cpp b/backends/cuda/runtime/test/test_cuda_kv_cache.cpp index f78c801705f..f9a00c9291e 100644 --- a/backends/cuda/runtime/test/test_cuda_kv_cache.cpp +++ b/backends/cuda/runtime/test/test_cuda_kv_cache.cpp @@ -9,6 +9,9 @@ #include #include #include +#include +#include +#include #include #include @@ -1295,3 +1298,270 @@ TEST_F(CudaKVPoolTest, FailedSideBufferAllocationLeavesThePoolRetryable) { EXPECT_EQ(pool.allocated_bytes(), 0); pool.forget_handle(&handle); } + +namespace { + +// A flat layer and a window-2 layer over 16 cells, steps of up to 8 tokens. +class CudaCellCacheTest : public CudaKVCacheTest { + protected: + static constexpr int kCells = 16; + static constexpr int kMaxWrite = 8; + static constexpr int kWindow = 2; + + static std::shared_ptr make(int initial = 4) { + return cu::make_cuda_cell_kv_cache( + flat_and_ring(kWindow), config(kCells, initial, kMaxWrite)); + } + + // What a program lowered in the cell layout for make() declares: every + // layer at kCells rows, and the step buffers at the widest step. + static FakeContainer container() { + FakeContainer container{ + {"f_k", "f_v", "r_k", "r_v", "cells", "read_len", "mask0", "mask2"}, + {"__et_offgraph_kv_layer_0_k", + "__et_offgraph_kv_layer_0_v", + "__et_offgraph_kv_layer_1_k", + "__et_offgraph_kv_layer_1_v", + "__et_offgraph_kv_cells", + "__et_offgraph_kv_read_len", + "__et_offgraph_kv_mask_w0", + "__et_offgraph_kv_mask_w2"}, + {}, + 0}; + for (size_t index = 0; index < 4; ++index) { + declare( + container, + container.fqns[index], + slimc10::ScalarType::BFloat16, + {1, kCells, kHeads, kDim}); + } + declare(container, "__et_offgraph_kv_cells", slimc10::ScalarType::Long, {kMaxWrite}); + declare(container, "__et_offgraph_kv_read_len", slimc10::ScalarType::Long, {1}); + for (const char* mask : + {"__et_offgraph_kv_mask_w0", "__et_offgraph_kv_mask_w2"}) { + declare( + container, mask, slimc10::ScalarType::Bool, {1, 1, kMaxWrite, kCells}); + } + return container; + } + + static std::vector read_longs(void* device, int count) { + std::vector values(count); + EXPECT_EQ( + cudaMemcpy( + values.data(), + device, + count * sizeof(int64_t), + cudaMemcpyDeviceToHost), + cudaSuccess); + return values; + } + + // Rows [0, rows) over columns [0, cols) of a [kMaxWrite, kCells] mask. + static std::vector> read_mask(void* device, int rows, int cols) { + std::vector flat(static_cast(kMaxWrite) * kCells); + EXPECT_EQ( + cudaMemcpy(flat.data(), device, flat.size(), cudaMemcpyDeviceToHost), + cudaSuccess); + std::vector> out(rows, std::vector(cols)); + for (int i = 0; i < rows; ++i) { + for (int j = 0; j < cols; ++j) { + out[i][j] = flat[static_cast(i) * kCells + j]; + } + } + return out; + } + + // declare + prepare + bind, as the executor and the delegate do per forward. + static Error step( + cache::Cache& cache, + cu::CudaDelegateHandle& handle, + const std::vector& seq_ids) { + auto* control = cache.as(); + auto* kv = cache.as(); + if (!control->declare_step(seq_ids)) { + return Error::InvalidArgument; + } + const auto width = static_cast(seq_ids.size()); + ET_CHECK_OK_OR_RETURN_ERROR(kv->prepare_step(width, cudaStreamPerThread)); + ET_CHECK_OK_OR_RETURN_ERROR(kv->rebind_for_execute(&handle)); + return kv->commit_step(width); + } +}; + +} // namespace + +TEST_F(CudaCellCacheTest, InterleavedSequencesWritePlacementAndMasks) { + auto cache_ptr = make(); + ASSERT_NE(cache_ptr, nullptr); + auto* control = cache_ptr->as(); + auto* kv = cache_ptr->as(); + ASSERT_NE(control, nullptr); + ASSERT_NE(kv, nullptr); + auto fake = container(); + auto handle = make_handle(fake); + ASSERT_TRUE(kv->note_handle(&handle).get()); + const int32_t a = *control->seq_new(); + const int32_t b = *control->seq_new(); + + // Both prefill in one forward. + ASSERT_EQ(step(*cache_ptr, handle, {a, a, a, b, b}), Error::Ok); + EXPECT_EQ(read_longs(fake.bound["cells"].data, 5), + std::vector({0, 1, 2, 3, 4})); + EXPECT_EQ(read_longs(fake.bound["read_len"].data, 1), std::vector({5})); + using Mask = std::vector>; + EXPECT_EQ( + read_mask(fake.bound["mask0"].data, 5, 5), + Mask({{1, 0, 0, 0, 0}, + {1, 1, 0, 0, 0}, + {1, 1, 1, 0, 0}, + {0, 0, 0, 1, 0}, + {0, 0, 0, 1, 1}})); + // The window-2 layer drops a's position 0 for its position-2 query. + EXPECT_EQ( + read_mask(fake.bound["mask2"].data, 5, 5), + Mask({{1, 0, 0, 0, 0}, + {1, 1, 0, 0, 0}, + {0, 1, 1, 0, 0}, + {0, 0, 0, 1, 0}, + {0, 0, 0, 1, 1}})); + EXPECT_EQ(fake.bound["mask0"].sizes, std::vector({1, 1, kMaxWrite, kCells})); + EXPECT_EQ(fake.bound["cells"].dtype, slimc10::ScalarType::Long); + EXPECT_EQ(fake.bound["mask0"].dtype, slimc10::ScalarType::Bool); + // Pools are declared at every cell; allocated past the first step's reach. + EXPECT_EQ(fake.bound["r_k"].sizes, std::vector({1, kCells, kHeads, kDim})); + EXPECT_EQ(kv->metrics().flat_capacity, 5); + + // Then they decode together, in the other order. + ASSERT_EQ(step(*cache_ptr, handle, {b, a}), Error::Ok); + EXPECT_EQ(read_longs(fake.bound["cells"].data, 2), std::vector({5, 6})); + EXPECT_EQ(read_longs(fake.bound["read_len"].data, 1), std::vector({7})); + EXPECT_EQ( + read_mask(fake.bound["mask0"].data, 2, 7), + Mask({{0, 0, 0, 1, 1, 1, 0}, {1, 1, 1, 0, 0, 0, 1}})); + EXPECT_EQ(control->pos(a), 4); + EXPECT_EQ(control->pos(b), 3); + const auto metrics = kv->metrics(); + EXPECT_EQ(metrics.logical_length, 7); + EXPECT_EQ(metrics.flat_capacity, 10); + EXPECT_EQ(metrics.growth_count, 1); + kv->forget_handle(&handle); +} + +TEST_F(CudaCellCacheTest, SwitchingSequencesKeepsACapturedGraph) { + // Room for every step below but the last: the pool grows ahead of placement + // to where the step could reach, used_end + width. + auto cache_ptr = make(/*initial=*/12); + auto* control = cache_ptr->as(); + auto* kv = cache_ptr->as(); + auto prefill_fake = container(); + auto decode_fake = container(); + auto prefill = make_handle(prefill_fake); + auto decode = make_handle(decode_fake); + ASSERT_TRUE(kv->note_handle(&prefill).get()); + ASSERT_TRUE(kv->note_handle(&decode).get()); + const int32_t a = *control->seq_new(); + const int32_t b = *control->seq_new(); + ASSERT_EQ(step(*cache_ptr, prefill, {a, a, b, b}), Error::Ok); + + ASSERT_EQ(step(*cache_ptr, decode, {a}), Error::Ok); + auto& graph = decode.cuda_graph_state; + graph.phase = cu::CudaGraphPhase::Replay; + const auto bound = decode_fake.bound; + + // Another sequence, then both: the graph reads new placements from the same + // addresses, so it stays captured. + ASSERT_EQ(step(*cache_ptr, decode, {b}), Error::Ok); + EXPECT_EQ(graph.phase, cu::CudaGraphPhase::Replay); + ASSERT_EQ(step(*cache_ptr, prefill, {b, a}), Error::Ok); + ASSERT_EQ(step(*cache_ptr, decode, {a}), Error::Ok); + EXPECT_EQ(graph.phase, cu::CudaGraphPhase::Replay); + for (const auto& [name, binding] : bound) { + EXPECT_EQ(decode_fake.bound[name].data, binding.data) << name; + } + // The two methods share every buffer. + for (const auto& [name, binding] : decode_fake.bound) { + EXPECT_EQ(prefill_fake.bound[name].data, binding.data) << name; + } + + // Growth moves the pools: the graph is captured again, and only the pools' + // bindings change. + ASSERT_EQ(step(*cache_ptr, prefill, {a, a, a, b, b, b}), Error::Ok); + EXPECT_NE(graph.phase, cu::CudaGraphPhase::Replay); + ASSERT_EQ(step(*cache_ptr, decode, {a}), Error::Ok); + EXPECT_NE(decode_fake.bound["f_k"].data, bound.at("f_k").data); + EXPECT_EQ(decode_fake.bound["cells"].data, bound.at("cells").data); + EXPECT_EQ(decode_fake.bound["mask2"].data, bound.at("mask2").data); + kv->forget_handle(&prefill); + kv->forget_handle(&decode); +} + +TEST_F(CudaCellCacheTest, ClonedSequenceSharesCellsUntilItsLastOwnerGoes) { + auto cache_ptr = make(/*initial=*/8); + auto* control = cache_ptr->as(); + auto* kv = cache_ptr->as(); + auto fake = container(); + auto handle = make_handle(fake); + ASSERT_TRUE(kv->note_handle(&handle).get()); + const int32_t a = *control->seq_new(); + ASSERT_EQ(step(*cache_ptr, handle, {a, a, a}), Error::Ok); + + const auto c = control->seq_clone(a, std::nullopt); + ASSERT_TRUE(c.has_value()); + EXPECT_EQ(control->pos(*c), 3); + ASSERT_EQ(step(*cache_ptr, handle, {*c}), Error::Ok); + EXPECT_EQ(read_longs(fake.bound["cells"].data, 1), std::vector({3})); + EXPECT_EQ( + read_mask(fake.bound["mask0"].data, 1, 4), + std::vector>({{1, 1, 1, 1}})); + + auto* cells = static_cast(cache_ptr.get()); + ASSERT_TRUE(control->seq_rm(a)); + EXPECT_EQ(cells->free_cells(), kCells - 4); + ASSERT_TRUE(control->seq_rm(*c)); + EXPECT_EQ(cells->free_cells(), kCells); + EXPECT_EQ(cells->used_end(), 0); + kv->forget_handle(&handle); +} + +TEST_F(CudaCellCacheTest, RejectsStepsThatDisagreeWithTheDeclaration) { + auto cache_ptr = make(); + auto* control = cache_ptr->as(); + auto* kv = cache_ptr->as(); + auto fake = container(); + auto handle = make_handle(fake); + ASSERT_TRUE(kv->note_handle(&handle).get()); + const int32_t a = *control->seq_new(); + + // Nothing declared. + EXPECT_EQ(kv->prepare_step(1, cudaStreamPerThread), Error::InvalidState); + // Declared two, the program's input carries three. + ASSERT_TRUE(control->declare_step({a, a})); + EXPECT_EQ(kv->prepare_step(3, cudaStreamPerThread), Error::InvalidState); + // Wider than the step buffers the program declared. + EXPECT_EQ( + kv->prepare_step(kMaxWrite + 1, cudaStreamPerThread), + Error::InvalidArgument); + // More tokens than free cells never declares. + EXPECT_FALSE(control->declare_step(std::vector(kCells + 1, a))); + // A rejected step leaves the table untouched. + EXPECT_EQ(control->pos(a), 0); + kv->forget_handle(&handle); +} + +TEST_F(CudaCellCacheTest, BuilderValidatesAndIsRegisteredForBatchedKinds) { + auto no_max_write = config(kCells, 4, kMaxWrite); + no_max_write.max_write.reset(); + EXPECT_EQ(cu::make_cuda_cell_kv_cache(flat_and_ring(kWindow), no_max_write), nullptr); + EXPECT_EQ( + cu::make_cuda_cell_kv_cache( + flat_and_ring(kWindow), config(kCells, 4, kCells + 1)), + nullptr); + for (const char* kind : {cache::kind::kBatchedCell, cache::kind::kBatched}) { + auto built = cache::CacheFactory::global().build( + cu::kCudaBackendId, kind, flat_and_ring(kWindow), config(kCells, 4, kMaxWrite)); + ASSERT_TRUE(built.ok()) << kind; + EXPECT_NE(built.get()->as(), nullptr) << kind; + EXPECT_NE(built.get()->as(), nullptr) << kind; + } +} diff --git a/extension/cuda/runtime_api.h b/extension/cuda/runtime_api.h index fc69f183946..1245cca80c0 100644 --- a/extension/cuda/runtime_api.h +++ b/extension/cuda/runtime_api.h @@ -134,6 +134,18 @@ inline cudaError_t cudaMemcpyAsync( return hipMemcpyAsync(dst, src, size, kind, stream); } +inline cudaError_t cudaMemcpy2DAsync( + void* dst, + size_t dpitch, + const void* src, + size_t spitch, + size_t width, + size_t height, + cudaMemcpyKind kind, + cudaStream_t stream) { + return hipMemcpy2DAsync(dst, dpitch, src, spitch, width, height, kind, stream); +} + inline cudaError_t cudaMemGetInfo(size_t* free, size_t* total) { return hipMemGetInfo(free, total); }