From f2c231d1b9353904ee4a6623664b6886e8206c4f Mon Sep 17 00:00:00 2001 From: gasoonjia Date: Fri, 2 Oct 2026 12:06:57 -0700 Subject: [PATCH] Update [ghstack-poisoned] --- CMakeLists.txt | 2 +- backends/cuda/batching/CMakeLists.txt | 5 +- backends/cuda/batching/cuda_executor.cpp | 240 +------------- backends/cuda/batching/cuda_executor.h | 38 +-- backends/cuda/batching/targets.bzl | 3 +- extension/llm/batching/CMakeLists.txt | 28 +- extension/llm/batching/module_executor.cpp | 262 +--------------- extension/llm/batching/module_executor.h | 33 +- extension/llm/batching/targets.bzl | 29 +- extension/llm/batching/test/CMakeLists.txt | 5 + .../llm/batching/test/session_table_test.cpp | 209 +++++++++++++ extension/llm/batching/test/targets.bzl | 13 + extension/llm/batching/util/session_table.cpp | 293 ++++++++++++++++++ extension/llm/batching/util/session_table.h | 121 ++++++++ 14 files changed, 726 insertions(+), 555 deletions(-) create mode 100644 extension/llm/batching/test/session_table_test.cpp create mode 100644 extension/llm/batching/util/session_table.cpp create mode 100644 extension/llm/batching/util/session_table.h diff --git a/CMakeLists.txt b/CMakeLists.txt index d8fd79ba98f..0c7cd2fd9c1 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -1013,7 +1013,7 @@ if(EXECUTORCH_BUILD_EXTENSION_LLM) list(APPEND _executorch_extensions extension_llm_cache) add_subdirectory(${CMAKE_CURRENT_SOURCE_DIR}/extension/llm/batching) list(APPEND _executorch_extensions extension_llm_batching - extension_llm_batching_module + extension_llm_batching_util extension_llm_batching_module ) add_subdirectory(${CMAKE_CURRENT_SOURCE_DIR}/extension/llm/serving) list(APPEND _executorch_extensions extension_llm_serving) diff --git a/backends/cuda/batching/CMakeLists.txt b/backends/cuda/batching/CMakeLists.txt index c665ab79485..ce063a41792 100644 --- a/backends/cuda/batching/CMakeLists.txt +++ b/backends/cuda/batching/CMakeLists.txt @@ -18,9 +18,8 @@ target_compile_options(cuda_batching PUBLIC ${_common_compile_options}) target_compile_definitions(cuda_batching PRIVATE CUDA_AVAILABLE=1) target_link_libraries( cuda_batching - PUBLIC aoti_cuda_backend extension_llm_batching extension_llm_cache - extension_module extension_tensor - PRIVATE extension_llm_sampler + PUBLIC aoti_cuda_backend extension_llm_batching extension_llm_batching_util + extension_llm_cache extension_module extension_tensor ) install( diff --git a/backends/cuda/batching/cuda_executor.cpp b/backends/cuda/batching/cuda_executor.cpp index b6c815fa490..904a30ad4b0 100644 --- a/backends/cuda/batching/cuda_executor.cpp +++ b/backends/cuda/batching/cuda_executor.cpp @@ -11,20 +11,16 @@ #include #include #include -#include -#include #include +#include #include #include #include -#include -#include #include #include #include #include -#include #include namespace executorch::backends::cuda::batching { @@ -32,13 +28,11 @@ namespace executorch::backends::cuda::batching { using ::executorch::extension::make_tensor_ptr; using ::executorch::extension::Module; using ::executorch::extension::llm::LogitsToKeepMode; -using ::executorch::extension::llm::Sampler; using ::executorch::runtime::Error; using ::executorch::runtime::MethodMeta; using ::executorch::runtime::Result; using llm_batching::BatchInput; using llm_batching::BatchOutput; -using llm_batching::Input; using llm_batching::Output; using llm_batching::Position; using llm_batching::SamplingParams; @@ -52,23 +46,6 @@ namespace { // A prefill forward carries at least this many tokens; one token is decode's. constexpr int kMinPrefillTokens = 2; -struct SequenceGuard { - llm_cache::BatchControl& control; - std::int32_t seq_id; - bool owned = true; - - ~SequenceGuard() { - if (owned) { - control.seq_rm(seq_id); - } - } -}; - -std::uint64_t nondeterministic_seed() { - std::random_device device; - return device(); -} - bool is_supported_logits_type(::executorch::aten::ScalarType type) { using ScalarType = ::executorch::aten::ScalarType; return type == ScalarType::Float || type == ScalarType::Half || @@ -174,11 +151,10 @@ CudaExecutor::CudaExecutor( module_(std::move(module)), ctl_(cache->as()), kv_(cache->as()), - max_sessions_(max_sessions), - max_session_tokens_(max_session_tokens), backend_id_(std::move(backend_id)), vocab_size_(vocab_size), - max_step_tokens_(max_step_tokens) {} + max_step_tokens_(max_step_tokens), + sessions_(*ctl_, max_sessions, max_session_tokens, vocab_size) {} CudaExecutor::~CudaExecutor() = default; @@ -333,185 +309,30 @@ bool CudaExecutor::initialize() { return true; } -// Copied from ModuleExecutor::build_step; see there for the reasoning. -Result CudaExecutor::build_step(const BatchInput& batch) { - Step step; - const std::size_t total = batch.size(); - step.tokens.reserve(total); - step.positions.reserve(total); - step.seq_ids.reserve(total); - step.logit_indices.reserve(batch.inputs.size()); - std::vector> rewinds; - std::unordered_map cursor; - - for (const Input& input : batch.inputs) { - const auto seq_it = sessions_.find(input.sid); - if (seq_it == sessions_.end()) { - ET_LOG(Error, "build_step: session %" PRId64 " is not open", input.sid); - return Error::InvalidArgument; - } - const std::int32_t seq_id = seq_it->second.seq_id; - if (input.size == 0 || !input.tokens || - input.offset > input.tokens->size() || - input.size > input.tokens->size() - input.offset) { - ET_LOG( - Error, - "build_step: session %" PRId64 " gave a slice its tokens do not hold", - input.sid); - return Error::InvalidArgument; - } - - const std::int64_t start = static_cast(input.position) + - static_cast(input.offset); - const auto [cursor_it, first_for_seq] = - cursor.try_emplace(seq_id, ctl_->pos(seq_id)); - int& at = cursor_it->second; - if (start > at) { - ET_LOG( - Error, - "build_step: session %" PRId64 " starts at %" PRId64 - " over a sequence holding %d", - input.sid, - start, - at); - return Error::InvalidArgument; - } - if (start < at) { - if (!first_for_seq) { - ET_LOG( - Error, - "build_step: session %" PRId64 " overlaps its earlier input", - input.sid); - return Error::InvalidArgument; - } - if (start == 0) { - ET_LOG( - Error, - "build_step: session %" PRId64 " reopens from the start", - input.sid); - return Error::InvalidArgument; - } - rewinds.emplace_back(seq_id, static_cast(start)); - at = static_cast(start); - } - - const std::int64_t end = start + static_cast(input.size); - if (end > max_session_tokens_) { - ET_LOG( - Error, - "build_step: session %" PRId64 " reaches %" PRId64 " of %d cells", - input.sid, - end, - max_session_tokens_); - return Error::OutOfResources; - } - - const Token* slice = input.tokens->data() + input.offset; - for (std::size_t k = 0; k < input.size; ++k) { - step.tokens.push_back(static_cast(slice[k])); - step.positions.push_back(start + static_cast(k)); - } - step.seq_ids.insert(step.seq_ids.end(), input.size, seq_id); - at = static_cast(end); - step.logit_indices.push_back( - input.produce_output ? static_cast(step.tokens.size()) - 1 : -1); - } - - for (const auto& [seq_id, from] : rewinds) { - if (!ctl_->rewind(seq_id, from)) { - ET_LOG(Error, "build_step: sequence %d would not truncate", seq_id); - return Error::Internal; - } - } - return step; -} - std::optional CudaExecutor::open_session() { - if (sessions_.size() >= static_cast(max_sessions_) || - next_session_ == 0) { - return std::nullopt; - } -#if ET_HAS_EXCEPTIONS - try { -#endif - const std::optional seq_id = ctl_->seq_new(); - return seq_id ? publish_session(*seq_id, 0) : std::nullopt; -#if ET_HAS_EXCEPTIONS - } catch (const std::bad_alloc&) { - return std::nullopt; - } -#endif -} - -std::optional CudaExecutor::publish_session( - std::int32_t seq_id, - Position position) { - SequenceGuard guard{*ctl_, seq_id}; - if (ctl_->pos(seq_id) != position) { - return std::nullopt; - } - const SessionId session = next_session_; - if (!sessions_.emplace(session, SessionState{seq_id, nullptr}).second) { - return std::nullopt; - } - guard.owned = false; - next_session_ = - session == std::numeric_limits::max() ? 0 : session + 1; - return session; + return sessions_.open(); } void CudaExecutor::close_session(SessionId session) { - const auto it = sessions_.find(session); - if (it == sessions_.end()) { - return; - } - ctl_->seq_rm(it->second.seq_id); - sessions_.erase(it); + sessions_.close(session); } std::optional CudaExecutor::clone(SessionId source, Position upto) { - const auto it = sessions_.find(source); - if (it == sessions_.end() || upto < 0 || upto > max_session_tokens_ || - sessions_.size() >= static_cast(max_sessions_) || - next_session_ == 0) { - return std::nullopt; - } -#if ET_HAS_EXCEPTIONS - try { -#endif - if (upto > ctl_->pos(it->second.seq_id)) { - return std::nullopt; - } - const auto seq_id = ctl_->seq_clone(it->second.seq_id, upto); - return seq_id ? publish_session(*seq_id, upto) : std::nullopt; -#if ET_HAS_EXCEPTIONS - } catch (const std::bad_alloc&) { - return std::nullopt; - } -#endif + return sessions_.clone(source, upto); } void CudaExecutor::set_sampling( SessionId session, const SamplingParams& params, std::optional seed) { - const auto it = sessions_.find(session); - if (it == sessions_.end()) { - return; - } - it->second.sampler = std::make_unique( - vocab_size_, - params.temperature, - params.top_p, - seed.value_or(nondeterministic_seed())); - it->second.sampler->set_topk(params.top_k); + sessions_.set_sampling(session, params, seed); } bool CudaExecutor::execute(const BatchInput& batch, BatchOutput& out) { out.outputs.clear(); out.outputs.resize(batch.inputs.size()); - const Result step = build_step(batch); + const Result step = sessions_.pack(batch); if (!step.ok()) { return false; } @@ -540,21 +361,9 @@ bool CudaExecutor::execute(const BatchInput& batch, BatchOutput& out) { {n}, std::vector( step->positions.begin() + off, step->positions.begin() + off + n)); - std::vector selector_values; - std::vector selected_inputs; - for (std::size_t i = 0; i < step->logit_indices.size(); ++i) { - const int row = step->logit_indices[i]; - if (row >= off && row < off + n) { - selector_values.push_back(row - off); - selected_inputs.push_back(i); - } - } - if (selector_values.empty()) { - // The forward still has to produce a row; nothing reads it. - selector_values.push_back(n - 1); - } - const int rows = static_cast(selector_values.size()); - auto selector = make_tensor_ptr({rows}, std::move(selector_values)); + auto selected = llm_batching::util::select_rows(*step, off, n); + const int rows = static_cast(selected.selector.size()); + auto selector = make_tensor_ptr({rows}, std::move(selected.selector)); auto result = module_->execute(method, {tokens, positions, selector}); if (!result.ok()) { @@ -580,11 +389,11 @@ bool CudaExecutor::execute(const BatchInput& batch, BatchOutput& out) { vocab_size_); return false; } - for (std::size_t row = 0; row < selected_inputs.size(); ++row) { - const std::size_t input_index = selected_inputs[row]; + for (std::size_t row = 0; row < selected.inputs.size(); ++row) { + const std::size_t input_index = selected.inputs[row]; const SessionId session = batch.inputs[input_index].sid; const std::optional token = - sample_row(logits, static_cast(row), session); + sessions_.sample(session, logits, static_cast(row)); if (!token) { return false; } @@ -594,27 +403,6 @@ bool CudaExecutor::execute(const BatchInput& batch, BatchOutput& out) { return true; } -std::optional CudaExecutor::sample_row( - ::executorch::aten::Tensor& logits, - int row, - SessionId session) { - const auto it = sessions_.find(session); - if (it == sessions_.end() || it->second.sampler == nullptr) { - ET_LOG( - Error, "CudaExecutor: session %" PRId64 " has no sampling policy", session); - return std::nullopt; - } - auto one_row = make_tensor_ptr( - {vocab_size_}, - static_cast(logits.mutable_data_ptr()) + - static_cast(row) * vocab_size_ * - ::executorch::runtime::elementSize(logits.scalar_type()), - logits.scalar_type()); - return static_cast( - ::executorch::extension::llm::sample_from_logits( - *one_row, *it->second.sampler)); -} - OffGraphKVMetrics CudaExecutor::kv_metrics() const { return kv_->metrics(); } diff --git a/backends/cuda/batching/cuda_executor.h b/backends/cuda/batching/cuda_executor.h index d9a7fd71074..ce609c6fd38 100644 --- a/backends/cuda/batching/cuda_executor.h +++ b/backends/cuda/batching/cuda_executor.h @@ -21,25 +21,16 @@ #include #include #include -#include -#include #include #include +#include #include #include #include #include #include // ET_EXPERIMENTAL -namespace executorch { -namespace extension { -namespace llm { -class Sampler; -} // namespace llm -} // namespace extension -} // namespace executorch - namespace executorch::backends::cuda::batching { namespace llm_batching = ::executorch::extension::llm::batching; @@ -108,18 +99,6 @@ class ET_EXPERIMENTAL CudaExecutor : public llm_batching::Executor { OffGraphKVMetrics kv_metrics() const; private: - struct SessionState { - std::int32_t seq_id; - std::unique_ptr<::executorch::extension::llm::Sampler> sampler; - }; - - struct Step { - std::vector tokens; - std::vector positions; - std::vector seq_ids; - std::vector logit_indices; - }; - CudaExecutor( std::unique_ptr<::executorch::extension::Module> module, std::shared_ptr cache, @@ -129,15 +108,6 @@ class ET_EXPERIMENTAL CudaExecutor : public llm_batching::Executor { std::int32_t vocab_size, int max_step_tokens); - ::executorch::runtime::Result build_step( - const llm_batching::BatchInput& batch); - std::optional publish_session( - std::int32_t seq_id, - llm_batching::Position position); - std::optional sample_row( - ::executorch::aten::Tensor& logits, - int row, - llm_batching::SessionId session); // Ordered so the module dies first, releasing the delegates that resolved // the cache before the registry entry naming it goes. @@ -145,14 +115,10 @@ class ET_EXPERIMENTAL CudaExecutor : public llm_batching::Executor { std::unique_ptr<::executorch::extension::Module> module_; llm_cache::BatchControl* const ctl_; const CudaKVCache* const kv_; - int max_sessions_; - int max_session_tokens_; std::string backend_id_; std::int32_t vocab_size_; int max_step_tokens_; - - llm_batching::SessionId next_session_ = 1; - std::unordered_map sessions_; + llm_batching::util::SessionTable sessions_; }; } // namespace executorch::backends::cuda::batching diff --git a/backends/cuda/batching/targets.bzl b/backends/cuda/batching/targets.bzl index 44ef58dad36..1e1cc32cd72 100644 --- a/backends/cuda/batching/targets.bzl +++ b/backends/cuda/batching/targets.bzl @@ -29,15 +29,14 @@ def define_common_targets(is_fbcode = False): ":step_plan", "//executorch/backends/cuda/runtime:cuda_backend", "//executorch/extension/llm/batching:batching", + "//executorch/extension/llm/batching:session_table", "//executorch/extension/llm/cache:kv_cache", "//executorch/extension/module:module", ], deps = [ "//executorch/extension/llm/runner:stats", - "//executorch/extension/llm/sampler:sampler", "//executorch/extension/tensor:tensor", "//executorch/runtime/backend:interface", - "//executorch/runtime/core/exec_aten/util:scalar_type_util", "//executorch/runtime/platform:platform", ], external_deps = [ diff --git a/extension/llm/batching/CMakeLists.txt b/extension/llm/batching/CMakeLists.txt index fac3cfff5f0..fbde92bcb0d 100644 --- a/extension/llm/batching/CMakeLists.txt +++ b/extension/llm/batching/CMakeLists.txt @@ -10,9 +10,10 @@ # types; the runner owns a thread, so this is a static library rather than an # INTERFACE target. # -# extension_llm_batching_module is a separate target because it implements that -# seam against a program and a KV cache, and so carries the runtime types the -# seam itself is kept clear of. +# extension_llm_batching_util is what every executor over a multi-sequence +# cache shares: the session table, batch packing and per-session sampling. +# extension_llm_batching_module implements the seam against a program and a KV +# cache, and so carries the runtime types the seam itself is kept clear of. if(NOT EXECUTORCH_ROOT) set(EXECUTORCH_ROOT ${CMAKE_CURRENT_SOURCE_DIR}/../../..) @@ -30,12 +31,24 @@ target_compile_options(extension_llm_batching PUBLIC ${_common_compile_options}) find_package(Threads REQUIRED) target_link_libraries(extension_llm_batching PUBLIC Threads::Threads) +add_library(extension_llm_batching_util util/session_table.cpp) +target_link_libraries( + extension_llm_batching_util + PUBLIC extension_llm_batching extension_llm_cache executorch_core + PRIVATE extension_llm_sampler extension_tensor +) +target_include_directories( + extension_llm_batching_util PUBLIC ${_common_include_directories} +) +target_compile_options( + extension_llm_batching_util PUBLIC ${_common_compile_options} +) + add_library(extension_llm_batching_module module_executor.cpp) target_link_libraries( extension_llm_batching_module - PUBLIC extension_llm_batching extension_llm_cache extension_module - extension_tensor - PRIVATE extension_llm_sampler + PUBLIC extension_llm_batching extension_llm_batching_util extension_llm_cache + extension_module extension_tensor ) target_include_directories( extension_llm_batching_module PUBLIC ${_common_include_directories} @@ -45,7 +58,8 @@ target_compile_options( ) install( - TARGETS extension_llm_batching extension_llm_batching_module + TARGETS extension_llm_batching extension_llm_batching_util + extension_llm_batching_module EXPORT ExecuTorchTargets DESTINATION ${CMAKE_INSTALL_LIBDIR} INCLUDES diff --git a/extension/llm/batching/module_executor.cpp b/extension/llm/batching/module_executor.cpp index c0ac6dd703a..a16c2b487f4 100644 --- a/extension/llm/batching/module_executor.cpp +++ b/extension/llm/batching/module_executor.cpp @@ -10,12 +10,8 @@ #include #include -#include -#include #include -#include -#include #include #include #include @@ -33,140 +29,14 @@ using ::executorch::runtime::Result; namespace { -struct SequenceGuard { - cache::BatchControl& control; - std::int32_t seq_id; - bool owned = true; - - ~SequenceGuard() { - if (owned) { - control.seq_rm(seq_id); - } - } -}; - bool is_supported_logits_type(::executorch::aten::ScalarType type) { using ScalarType = ::executorch::aten::ScalarType; return type == ScalarType::Float || type == ScalarType::Half || type == ScalarType::BFloat16 || type == ScalarType::UInt16; } -std::uint64_t nondeterministic_seed() { - std::random_device device; - return device(); -} - } // namespace -Result ModuleExecutor::build_step( - const BatchInput& batch) { - // Flatten the batch and truncate whatever it reopens; execute() declares each - // slice to the cache as it runs it. A per-sequence cursor carries the batch's - // own writes, so consecutive chunks of one prompt abut and only the first can - // reopen committed ground. Every input is checked before any is truncated, so - // a refusal leaves the cache untouched. - Step step; - const std::size_t total = batch.size(); - step.tokens.reserve(total); - step.positions.reserve(total); - step.logit_indices.reserve(batch.inputs.size()); - - step.seq_ids.reserve(total); - // Truncations the batch asks for, held until every input has been checked. - std::vector> rewinds; - // Where each sequence stands mid-batch: the cache still reports what it held - // before the step, so the batch's own writes live here. - std::unordered_map cursor; - - for (const Input& input : batch.inputs) { - const auto seq_it = sessions_.find(input.sid); - if (seq_it == sessions_.end()) { - ET_LOG(Error, "build_step: session %" PRId64 " is not open", input.sid); - return Error::InvalidArgument; - } - const std::int32_t seq_id = seq_it->second.seq_id; - if (input.size == 0 || !input.tokens || - input.offset > input.tokens->size() || - input.size > input.tokens->size() - input.offset) { - ET_LOG( - Error, - "build_step: session %" PRId64 " gave a slice its tokens do not hold", - input.sid); - return Error::InvalidArgument; - } - - const std::int64_t start = static_cast(input.position) + - static_cast(input.offset); - const auto [cursor_it, first_for_seq] = - cursor.try_emplace(seq_id, ctl_->pos(seq_id)); - int& at = cursor_it->second; - if (start > at) { - // Positions nothing attended, and nothing later reaches back to fill. - ET_LOG( - Error, - "build_step: session %" PRId64 " starts at %" PRId64 - " over a sequence holding %d", - input.sid, - start, - at); - return Error::InvalidArgument; - } - if (start < at) { - if (!first_for_seq) { - // Its predecessor in this batch has already been laid down, so a - // rewind now would truncate committed cells for a step whose - // positions repeat and cannot be placed. - ET_LOG( - Error, - "build_step: session %" PRId64 " overlaps its earlier input", - input.sid); - return Error::InvalidArgument; - } - if (start == 0) { - // Emptying a sequence hands its id back, and the step names it. - ET_LOG( - Error, - "build_step: session %" PRId64 " reopens from the start", - input.sid); - return Error::InvalidArgument; - } - rewinds.emplace_back(seq_id, static_cast(start)); - at = static_cast(start); - } - - const std::int64_t end = start + static_cast(input.size); - if (end > max_session_tokens_) { - ET_LOG( - Error, - "build_step: session %" PRId64 " reaches %" PRId64 " of %d cells", - input.sid, - end, - max_session_tokens_); - return Error::OutOfResources; - } - - const Token* slice = input.tokens->data() + input.offset; - for (std::size_t k = 0; k < input.size; ++k) { - step.tokens.push_back(static_cast(slice[k])); - } - for (std::size_t k = 0; k < input.size; ++k) { - step.positions.push_back(start + static_cast(k)); - } - step.seq_ids.insert(step.seq_ids.end(), input.size, seq_id); - at = static_cast(end); - step.logit_indices.push_back( - input.produce_output ? static_cast(step.tokens.size()) - 1 : -1); - } - - for (const auto& [seq_id, from] : rewinds) { - if (!ctl_->rewind(seq_id, from)) { - ET_LOG(Error, "build_step: sequence %d would not truncate", seq_id); - return Error::Internal; - } - } - return step; -} - ModuleExecutor::ModuleExecutor( std::unique_ptr module, std::shared_ptr cache, @@ -180,13 +50,12 @@ ModuleExecutor::ModuleExecutor( : install_guard_(cache), module_(std::move(module)), ctl_(cache->as()), - max_sessions_(max_sessions), - max_session_tokens_(max_session_tokens), backend_id_(std::move(backend_id)), method_(std::move(method)), vocab_size_(vocab_size), max_step_tokens_(max_step_tokens), - logits_to_keep_mode_(logits_to_keep_mode) {} + logits_to_keep_mode_(logits_to_keep_mode), + sessions_(*ctl_, max_sessions, max_session_tokens, vocab_size) {} ModuleExecutor::~ModuleExecutor() = default; @@ -425,95 +294,31 @@ bool ModuleExecutor::initialize() { } std::optional ModuleExecutor::open_session() { - if (sessions_.size() >= static_cast(max_sessions_) || - next_session_ == 0) { - return std::nullopt; - } -#if ET_HAS_EXCEPTIONS - try { -#endif - const std::optional seq_id = ctl_->seq_new(); - return seq_id ? publish_session(*seq_id, 0) : std::nullopt; -#if ET_HAS_EXCEPTIONS - } catch (const std::bad_alloc&) { - return std::nullopt; - } -#endif -} - -std::optional ModuleExecutor::publish_session( - std::int32_t seq_id, - Position position) { - SequenceGuard guard{*ctl_, seq_id}; - if (ctl_->pos(seq_id) != position) { - return std::nullopt; - } - const SessionId session = next_session_; - if (!sessions_.emplace(session, SessionState{seq_id, nullptr}).second) { - return std::nullopt; - } - guard.owned = false; - next_session_ = - session == std::numeric_limits::max() ? 0 : session + 1; - return session; + return sessions_.open(); } void ModuleExecutor::close_session(SessionId session) { - const auto it = sessions_.find(session); - if (it == sessions_.end()) { - return; - } - // Frees the cells and hands the sequence id back. The session id is not. - ctl_->seq_rm(it->second.seq_id); - sessions_.erase(it); + sessions_.close(session); } std::optional ModuleExecutor::clone( SessionId source, Position upto) { - const auto it = sessions_.find(source); - if (it == sessions_.end() || upto < 0 || upto > max_session_tokens_ || - sessions_.size() >= static_cast(max_sessions_) || - next_session_ == 0) { - return std::nullopt; - } -#if ET_HAS_EXCEPTIONS - try { -#endif - if (upto > ctl_->pos(it->second.seq_id)) { - return std::nullopt; - } - const auto seq_id = ctl_->seq_clone(it->second.seq_id, upto); - return seq_id ? publish_session(*seq_id, upto) : std::nullopt; -#if ET_HAS_EXCEPTIONS - } catch (const std::bad_alloc&) { - return std::nullopt; - } -#endif + return sessions_.clone(source, upto); } void ModuleExecutor::set_sampling( SessionId session, const SamplingParams& params, std::optional seed) { - const auto it = sessions_.find(session); - if (it == sessions_.end()) { - return; - } - // One sampler per generation, carrying its own generator state from here on. - it->second.sampler = std::make_unique( - vocab_size_, - params.temperature, - params.top_p, - seed.value_or(nondeterministic_seed())); - it->second.sampler->set_topk(params.top_k); + sessions_.set_sampling(session, params, seed); } bool ModuleExecutor::execute(const BatchInput& batch, BatchOutput& out) { out.outputs.clear(); out.outputs.resize(batch.inputs.size()); - const Result step = build_step(batch); + const Result step = sessions_.pack(batch); if (!step.ok()) { return false; } @@ -539,28 +344,18 @@ bool ModuleExecutor::execute(const BatchInput& batch, BatchOutput& out) { {n}, std::vector( step->positions.begin() + off, step->positions.begin() + off + n)); - std::vector selector_values; - std::vector selected_inputs; + util::SliceRows selected; if (logits_to_keep_mode_ == LogitsToKeepMode::Selected) { - for (std::size_t i = 0; i < step->logit_indices.size(); ++i) { - const int row = step->logit_indices[i]; - if (row >= off && row < off + n) { - selector_values.push_back(row - off); - selected_inputs.push_back(i); - } - } - if (selector_values.empty()) { - selector_values.push_back(n - 1); - } + selected = util::select_rows(*step, off, n); } const int expected_rows = logits_to_keep_mode_ == LogitsToKeepMode::Selected - ? static_cast(selector_values.size()) + ? static_cast(selected.selector.size()) : n; auto result = [&]() -> Result> { if (logits_to_keep_mode_ == LogitsToKeepMode::Selected) { auto selector = - make_tensor_ptr({expected_rows}, std::move(selector_values)); + make_tensor_ptr({expected_rows}, std::move(selected.selector)); return module_->execute(method_, {tokens, positions, selector}); } return module_->execute(method_, {tokens, positions}); @@ -592,11 +387,11 @@ bool ModuleExecutor::execute(const BatchInput& batch, BatchOutput& out) { } if (logits_to_keep_mode_ == LogitsToKeepMode::Selected) { - for (std::size_t row = 0; row < selected_inputs.size(); ++row) { - const std::size_t input_index = selected_inputs[row]; + for (std::size_t row = 0; row < selected.inputs.size(); ++row) { + const std::size_t input_index = selected.inputs[row]; const SessionId session = batch.inputs[input_index].sid; const std::optional token = - sample_row(logits, static_cast(row), session); + sessions_.sample(session, logits, static_cast(row)); if (!token) { return false; } @@ -610,7 +405,7 @@ bool ModuleExecutor::execute(const BatchInput& batch, BatchOutput& out) { } const SessionId session = batch.inputs[i].sid; const std::optional token = - sample_row(logits, row - off, session); + sessions_.sample(session, logits, row - off); if (!token) { return false; } @@ -621,33 +416,6 @@ bool ModuleExecutor::execute(const BatchInput& batch, BatchOutput& out) { return true; } -std::optional ModuleExecutor::sample_row( - ::executorch::aten::Tensor& logits, - int row, - SessionId session) { - const auto it = sessions_.find(session); - if (it == sessions_.end() || it->second.sampler == nullptr) { - ET_LOG( - Error, - "ModuleExecutor: session %" PRId64 " has no sampling policy", - session); - return std::nullopt; - } - if (row >= logits.numel() / vocab_size_) { - ET_LOG(Error, "ModuleExecutor: logits hold no row %d", row); - return std::nullopt; - } - // A one-row view over the model's own output: sample_from_logits reduces in - // place and reads the last dimension. - auto one_row = make_tensor_ptr( - {vocab_size_}, - static_cast(logits.mutable_data_ptr()) + - static_cast(row) * vocab_size_ * - ::executorch::runtime::elementSize(logits.scalar_type()), - logits.scalar_type()); - return static_cast(sample_from_logits(*one_row, *it->second.sampler)); -} - } // namespace batching } // namespace llm } // namespace extension diff --git a/extension/llm/batching/module_executor.h b/extension/llm/batching/module_executor.h index 1cc908e9231..ef199e21523 100644 --- a/extension/llm/batching/module_executor.h +++ b/extension/llm/batching/module_executor.h @@ -17,10 +17,10 @@ #include #include #include -#include #include #include +#include #include #include #include @@ -31,9 +31,6 @@ namespace executorch { namespace extension { namespace llm { - -class Sampler; - namespace batching { namespace cache = ::executorch::extension::llm::cache; @@ -90,23 +87,6 @@ class ET_EXPERIMENTAL ModuleExecutor : public Executor { bool execute(const BatchInput& batch, BatchOutput& out) override; private: - struct SessionState { - std::int32_t seq_id; - std::unique_ptr sampler; - }; - - struct Step { - std::vector tokens; - std::vector positions; - std::vector seq_ids; - std::vector logit_indices; - }; - - ::executorch::runtime::Result build_step(const BatchInput& batch); - std::optional publish_session( - std::int32_t seq_id, - Position position); - ModuleExecutor( std::unique_ptr module, std::shared_ptr cache, @@ -118,27 +98,18 @@ class ET_EXPERIMENTAL ModuleExecutor : public Executor { int max_step_tokens, LogitsToKeepMode logits_to_keep_mode); - // Draw the token an input produced from its row of `logits`, which the - // session's sampler consumes in place. - std::optional - sample_row(::executorch::aten::Tensor& logits, int row, SessionId session); - // Ordered so the module dies first, releasing the delegate that resolved the // cache before the registry entry naming it goes. cache::InstallGuard install_guard_; std::unique_ptr module_; cache::BatchControl* const ctl_; - int max_sessions_; - int max_session_tokens_; std::string backend_id_; std::string method_; // The method's logits width, so a sampler can be built by its policy. std::int32_t vocab_size_; int max_step_tokens_; LogitsToKeepMode logits_to_keep_mode_; - - SessionId next_session_ = 1; // never reused, unlike the cache's sequence ids - std::unordered_map sessions_; + util::SessionTable sessions_; }; } // namespace batching diff --git a/extension/llm/batching/targets.bzl b/extension/llm/batching/targets.bzl index ec5c6ad0818..667d0f2f5e2 100644 --- a/extension/llm/batching/targets.bzl +++ b/extension/llm/batching/targets.bzl @@ -4,8 +4,9 @@ def define_common_targets(): """Mirrors the extension_llm_batching CMake targets. `batching` is the scheduler, the executor seam and the runner, free of - ExecuTorch runtime types. `module_executor` implements the seam against a - program and a KV cache. + ExecuTorch runtime types. `session_table` is what every executor over a + multi-sequence cache shares. `module_executor` implements the seam against + a program and a KV cache. """ runtime.cxx_library( name = "batching", @@ -28,6 +29,29 @@ def define_common_targets(): ], ) + runtime.cxx_library( + name = "session_table", + srcs = [ + "util/session_table.cpp", + ], + exported_headers = [ + "util/session_table.h", + ], + visibility = ["PUBLIC"], + exported_deps = [ + ":batching", + "//executorch/extension/llm/cache:kv_cache", + "//executorch/runtime/core:core", + "//executorch/runtime/core/exec_aten:lib", + ], + deps = [ + "//executorch/extension/llm/sampler:sampler", + "//executorch/extension/tensor:tensor", + "//executorch/runtime/core/exec_aten/util:scalar_type_util", + "//executorch/runtime/platform:platform", + ], + ) + runtime.cxx_library( name = "module_executor", srcs = [ @@ -39,6 +63,7 @@ def define_common_targets(): visibility = ["PUBLIC"], exported_deps = [ ":batching", + ":session_table", "//executorch/extension/llm/cache:kv_cache", "//executorch/extension/llm/runner:stats", "//executorch/extension/module:module", diff --git a/extension/llm/batching/test/CMakeLists.txt b/extension/llm/batching/test/CMakeLists.txt index 456c6403390..7ad202af8fb 100644 --- a/extension/llm/batching/test/CMakeLists.txt +++ b/extension/llm/batching/test/CMakeLists.txt @@ -16,3 +16,8 @@ et_cxx_test( extension_llm_batching_test SOURCES ${_test_srcs} EXTRA_LIBS extension_llm_batching ) + +et_cxx_test( + extension_llm_batching_util_test SOURCES session_table_test.cpp EXTRA_LIBS + extension_llm_batching_util extension_tensor +) diff --git a/extension/llm/batching/test/session_table_test.cpp b/extension/llm/batching/test/session_table_test.cpp new file mode 100644 index 00000000000..2c5e9cd379a --- /dev/null +++ b/extension/llm/batching/test/session_table_test.cpp @@ -0,0 +1,209 @@ +/* + * Copyright (c) Meta Platforms, Inc. and affiliates. + * All rights reserved. + * + * This source code is licensed under the BSD-style license found in the + * LICENSE file in the root directory of this source tree. + */ + +#include + +#include +#include +#include + +#include + +#include +#include + +namespace batching = ::executorch::extension::llm::batching; +namespace cache = ::executorch::extension::llm::cache; +namespace util = ::executorch::extension::llm::batching::util; +using ::executorch::extension::make_tensor_ptr; +using ::executorch::runtime::Error; + +namespace { + +constexpr int kVocab = 8; + +// The neutral cell cache is a complete BatchControl without any bytes, so the +// table runs against the real sequence semantics on the CPU. +class SessionTableTest : public ::testing::Test { + protected: + void SetUp() override { + ::executorch::runtime::runtime_init(); + cache::CacheGeometry geometry; + geometry.layers = {{{cache::LayerPolicy::Kind::Flat, 0}, 1, 1}}; + cache::CacheConfig cfg{/*capacity=*/64, /*kv_dtype=*/6}; + cells_ = std::make_unique(geometry, cfg); + table_ = std::make_unique( + *cells_, /*max_sessions=*/3, /*max_session_tokens=*/16, kVocab); + } + + // Declares and places a step, as a forward would. + void run(const util::PackedStep& step) { + ASSERT_TRUE(cells_->declare_step(step.seq_ids)); + std::vector positions(step.positions.begin(), step.positions.end()); + ASSERT_NE( + cells_->place_step(0, positions.data(), static_cast(positions.size())), + nullptr); + } + + static batching::Input input( + batching::SessionId sid, + std::vector tokens, + batching::Position position, + bool produce_output = true, + std::size_t offset = 0, + std::size_t size = 0) { + auto shared = std::make_shared>(std::move(tokens)); + batching::Input in; + in.sid = sid; + in.tokens = shared; + in.offset = offset; + in.size = size == 0 ? shared->size() - offset : size; + in.position = position; + in.produce_output = produce_output; + return in; + } + + std::unique_ptr cells_; + std::unique_ptr table_; +}; + +} // namespace + +TEST_F(SessionTableTest, PacksInputsInOrderWithTheirRows) { + const auto a = *table_->open(); + const auto b = *table_->open(); + batching::BatchInput batch; + batch.inputs = {input(a, {1, 2, 3}, 0), input(b, {7}, 0, false)}; + const auto step = table_->pack(batch); + ASSERT_EQ(step.error(), Error::Ok); + EXPECT_EQ(step->tokens, std::vector({1, 2, 3, 7})); + EXPECT_EQ(step->positions, std::vector({0, 1, 2, 0})); + EXPECT_EQ(step->seq_ids[0], step->seq_ids[2]); + EXPECT_NE(step->seq_ids[0], step->seq_ids[3]); + EXPECT_EQ(step->logit_indices, std::vector({2, -1})); +} + +TEST_F(SessionTableTest, ConsecutiveChunksOfOnePromptAbut) { + const auto a = *table_->open(); + batching::BatchInput batch; + const std::vector prompt{1, 2, 3, 4, 5}; + batch.inputs = { + input(a, prompt, 0, false, 0, 2), input(a, prompt, 0, true, 2, 3)}; + const auto step = table_->pack(batch); + ASSERT_EQ(step.error(), Error::Ok); + EXPECT_EQ(step->positions, std::vector({0, 1, 2, 3, 4})); + EXPECT_EQ(step->logit_indices, std::vector({-1, 4})); +} + +TEST_F(SessionTableTest, RewindsOnlyAfterEveryInputChecks) { + const auto a = *table_->open(); + const auto b = *table_->open(); + batching::BatchInput first; + first.inputs = {input(a, {1, 2, 3, 4}, 0), input(b, {5, 6}, 0)}; + const auto step = table_->pack(first); + ASSERT_EQ(step.error(), Error::Ok); + run(*step); + + // a reopens at 2; b skips ahead, which refuses the batch before a rewinds. + batching::BatchInput bad; + bad.inputs = {input(a, {9}, 2), input(b, {9}, 5)}; + EXPECT_EQ(table_->pack(bad).error(), Error::InvalidArgument); + EXPECT_EQ(cells_->pos(step->seq_ids[0]), 4); + + batching::BatchInput good; + good.inputs = {input(a, {9}, 2)}; + const auto rewound = table_->pack(good); + ASSERT_EQ(rewound.error(), Error::Ok); + EXPECT_EQ(cells_->pos(rewound->seq_ids[0]), 2); + EXPECT_EQ(rewound->positions, std::vector({2})); +} + +TEST_F(SessionTableTest, RefusesOverlapsRestartsAndOverruns) { + const auto a = *table_->open(); + batching::BatchInput first; + first.inputs = {input(a, {1, 2, 3}, 0)}; + const auto step = table_->pack(first); + ASSERT_EQ(step.error(), Error::Ok); + run(*step); + + batching::BatchInput overlap; + overlap.inputs = {input(a, {4}, 3), input(a, {5}, 2)}; + EXPECT_EQ(table_->pack(overlap).error(), Error::InvalidArgument); + batching::BatchInput restart; + restart.inputs = {input(a, {4}, 0)}; + EXPECT_EQ(table_->pack(restart).error(), Error::InvalidArgument); + batching::BatchInput overrun; + overrun.inputs = {input(a, std::vector(14, 1), 3)}; + EXPECT_EQ(table_->pack(overrun).error(), Error::OutOfResources); + batching::BatchInput unknown; + unknown.inputs = {input(a + 100, {4}, 0)}; + EXPECT_EQ(table_->pack(unknown).error(), Error::InvalidArgument); +} + +TEST_F(SessionTableTest, SessionsAreBoundedAndIdsNeverReused) { + const auto a = *table_->open(); + const auto b = *table_->open(); + const auto c = *table_->open(); + EXPECT_FALSE(table_->open().has_value()); + table_->close(b); + const auto d = *table_->open(); + EXPECT_NE(d, b); + EXPECT_EQ(table_->size(), 3u); + (void)a; + (void)c; +} + +TEST_F(SessionTableTest, CloneSharesTheSourcePrefix) { + const auto a = *table_->open(); + batching::BatchInput batch; + batch.inputs = {input(a, {1, 2, 3, 4}, 0)}; + const auto step = table_->pack(batch); + ASSERT_EQ(step.error(), Error::Ok); + run(*step); + + const auto c = table_->clone(a, 3); + ASSERT_TRUE(c.has_value()); + EXPECT_FALSE(table_->clone(a, 5).has_value()); // past the source + batching::BatchInput next; + next.inputs = {input(*c, {8}, 3)}; + const auto cloned = table_->pack(next); + ASSERT_EQ(cloned.error(), Error::Ok); + EXPECT_EQ(cloned->positions, std::vector({3})); +} + +TEST_F(SessionTableTest, SamplesEachSessionsRowWithItsOwnPolicy) { + const auto a = *table_->open(); + const auto b = *table_->open(); + batching::SamplingParams greedy; + greedy.temperature = 0.0f; + table_->set_sampling(a, greedy, 0); + std::vector values(2 * kVocab, 0.0f); + values[3] = 5.0f; // row 0 peaks at 3 + values[kVocab + 6] = 5.0f; // row 1 peaks at 6 + auto logits = make_tensor_ptr({2, kVocab}, values); + EXPECT_EQ(table_->sample(a, *logits, 1), 6u); + EXPECT_EQ(table_->sample(a, *logits, 0), 3u); + // No policy yet, and no such row. + EXPECT_FALSE(table_->sample(b, *logits, 0).has_value()); + EXPECT_FALSE(table_->sample(a, *logits, 2).has_value()); +} + +TEST(SelectRowsTest, KeepsOnlyTheSlicesWantedRows) { + util::PackedStep step; + step.logit_indices = {2, -1, 5, 9}; + auto rows = util::select_rows(step, 0, 6); + EXPECT_EQ(rows.selector, std::vector({2, 5})); + EXPECT_EQ(rows.inputs, std::vector({0, 2})); + rows = util::select_rows(step, 6, 4); + EXPECT_EQ(rows.selector, std::vector({3})); + EXPECT_EQ(rows.inputs, std::vector({3})); + // Nothing wanted: the last token still runs, and nothing reads it. + rows = util::select_rows(step, 10, 3); + EXPECT_EQ(rows.selector, std::vector({2})); + EXPECT_TRUE(rows.inputs.empty()); +} diff --git a/extension/llm/batching/test/targets.bzl b/extension/llm/batching/test/targets.bzl index fbab7d5326d..c58362dc7ae 100644 --- a/extension/llm/batching/test/targets.bzl +++ b/extension/llm/batching/test/targets.bzl @@ -16,3 +16,16 @@ def define_common_targets(): "//executorch/extension/llm/batching:batching", ], ) + + runtime.cxx_test( + name = "session_table_test", + srcs = [ + "session_table_test.cpp", + ], + deps = [ + "//executorch/extension/llm/batching:session_table", + "//executorch/extension/llm/cache:kv_cache", + "//executorch/extension/tensor:tensor", + "//executorch/runtime/platform:platform", + ], + ) diff --git a/extension/llm/batching/util/session_table.cpp b/extension/llm/batching/util/session_table.cpp new file mode 100644 index 00000000000..0b47df00c98 --- /dev/null +++ b/extension/llm/batching/util/session_table.cpp @@ -0,0 +1,293 @@ +/* + * Copyright (c) Meta Platforms, Inc. and affiliates. + * All rights reserved. + * + * This source code is licensed under the BSD-style license found in the + * LICENSE file in the root directory of this source tree. + */ + +#include + +#include +#include +#include +#include +#include + +// sampler/util.h switches on dtype with the macros this defines. +#include + +#include +#include +#include +#include + +namespace executorch { +namespace extension { +namespace llm { +namespace batching { +namespace util { + +using ::executorch::extension::make_tensor_ptr; +using ::executorch::runtime::Error; +using ::executorch::runtime::Result; + +namespace { + +// Releases a sequence unless ownership passes to a session. +struct SequenceGuard { + cache::BatchControl& control; + std::int32_t seq_id; + bool owned = true; + + ~SequenceGuard() { + if (owned) { + control.seq_rm(seq_id); + } + } +}; + +std::uint64_t nondeterministic_seed() { + std::random_device device; + return device(); +} + +} // namespace + +SliceRows select_rows(const PackedStep& step, int offset, int length) { + SliceRows rows; + for (std::size_t i = 0; i < step.logit_indices.size(); ++i) { + const int row = step.logit_indices[i]; + if (row >= offset && row < offset + length) { + rows.selector.push_back(row - offset); + rows.inputs.push_back(i); + } + } + if (rows.selector.empty()) { + rows.selector.push_back(length - 1); + } + return rows; +} + +SessionTable::SessionTable( + cache::BatchControl& control, + int max_sessions, + int max_session_tokens, + std::int32_t vocab_size) + : ctl_(control), + max_sessions_(max_sessions), + max_session_tokens_(max_session_tokens), + vocab_size_(vocab_size) {} + +SessionTable::~SessionTable() = default; + +std::optional SessionTable::open() { + if (sessions_.size() >= static_cast(max_sessions_) || + next_session_ == 0) { + return std::nullopt; + } +#if ET_HAS_EXCEPTIONS + try { +#endif + const std::optional seq_id = ctl_.seq_new(); + return seq_id ? publish(*seq_id, 0) : std::nullopt; +#if ET_HAS_EXCEPTIONS + } catch (const std::bad_alloc&) { + return std::nullopt; + } +#endif +} + +std::optional SessionTable::publish( + std::int32_t seq_id, + Position position) { + SequenceGuard guard{ctl_, seq_id}; + if (ctl_.pos(seq_id) != position) { + return std::nullopt; + } + const SessionId session = next_session_; + if (!sessions_.emplace(session, Session{seq_id, nullptr}).second) { + return std::nullopt; + } + guard.owned = false; + next_session_ = + session == std::numeric_limits::max() ? 0 : session + 1; + return session; +} + +void SessionTable::close(SessionId session) { + const auto it = sessions_.find(session); + if (it == sessions_.end()) { + return; + } + ctl_.seq_rm(it->second.seq_id); + sessions_.erase(it); +} + +std::optional SessionTable::clone(SessionId source, Position upto) { + const auto it = sessions_.find(source); + if (it == sessions_.end() || upto < 0 || upto > max_session_tokens_ || + sessions_.size() >= static_cast(max_sessions_) || + next_session_ == 0) { + return std::nullopt; + } +#if ET_HAS_EXCEPTIONS + try { +#endif + if (upto > ctl_.pos(it->second.seq_id)) { + return std::nullopt; + } + const auto seq_id = ctl_.seq_clone(it->second.seq_id, upto); + return seq_id ? publish(*seq_id, upto) : std::nullopt; +#if ET_HAS_EXCEPTIONS + } catch (const std::bad_alloc&) { + return std::nullopt; + } +#endif +} + +void SessionTable::set_sampling( + SessionId session, + const SamplingParams& params, + std::optional seed) { + const auto it = sessions_.find(session); + if (it == sessions_.end()) { + return; + } + // One sampler per generation, carrying its own generator state from here on. + it->second.sampler = std::make_unique( + vocab_size_, + params.temperature, + params.top_p, + seed.value_or(nondeterministic_seed())); + it->second.sampler->set_topk(params.top_k); +} + +Result SessionTable::pack(const BatchInput& batch) { + PackedStep step; + const std::size_t total = batch.size(); + step.tokens.reserve(total); + step.positions.reserve(total); + step.seq_ids.reserve(total); + step.logit_indices.reserve(batch.inputs.size()); + // Truncations the batch asks for, held until every input has been checked. + std::vector> rewinds; + // Where each sequence stands mid-batch: the cache still reports what it held + // before the step, so the batch's own writes live here. + std::unordered_map cursor; + + for (const Input& input : batch.inputs) { + const auto seq_it = sessions_.find(input.sid); + if (seq_it == sessions_.end()) { + ET_LOG(Error, "pack: session %" PRId64 " is not open", input.sid); + return Error::InvalidArgument; + } + const std::int32_t seq_id = seq_it->second.seq_id; + if (input.size == 0 || !input.tokens || + input.offset > input.tokens->size() || + input.size > input.tokens->size() - input.offset) { + ET_LOG( + Error, + "pack: session %" PRId64 " gave a slice its tokens do not hold", + input.sid); + return Error::InvalidArgument; + } + + const std::int64_t start = static_cast(input.position) + + static_cast(input.offset); + const auto [cursor_it, first_for_seq] = + cursor.try_emplace(seq_id, ctl_.pos(seq_id)); + int& at = cursor_it->second; + if (start > at) { + // Positions nothing attended, and nothing later reaches back to fill. + ET_LOG( + Error, + "pack: session %" PRId64 " starts at %" PRId64 + " over a sequence holding %d", + input.sid, + start, + at); + return Error::InvalidArgument; + } + if (start < at) { + if (!first_for_seq) { + // Its predecessor in this batch has already been laid down, so a + // rewind now would truncate committed cells for a step whose + // positions repeat and cannot be placed. + ET_LOG( + Error, + "pack: session %" PRId64 " overlaps its earlier input", + input.sid); + return Error::InvalidArgument; + } + if (start == 0) { + // Emptying a sequence hands its id back, and the step names it. + ET_LOG( + Error, "pack: session %" PRId64 " reopens from the start", input.sid); + return Error::InvalidArgument; + } + rewinds.emplace_back(seq_id, static_cast(start)); + at = static_cast(start); + } + + const std::int64_t end = start + static_cast(input.size); + if (end > max_session_tokens_) { + ET_LOG( + Error, + "pack: session %" PRId64 " reaches %" PRId64 " of %d cells", + input.sid, + end, + max_session_tokens_); + return Error::OutOfResources; + } + + const Token* slice = input.tokens->data() + input.offset; + for (std::size_t k = 0; k < input.size; ++k) { + step.tokens.push_back(static_cast(slice[k])); + step.positions.push_back(start + static_cast(k)); + } + step.seq_ids.insert(step.seq_ids.end(), input.size, seq_id); + at = static_cast(end); + step.logit_indices.push_back( + input.produce_output ? static_cast(step.tokens.size()) - 1 : -1); + } + + for (const auto& [seq_id, from] : rewinds) { + if (!ctl_.rewind(seq_id, from)) { + ET_LOG(Error, "pack: sequence %d would not truncate", seq_id); + return Error::Internal; + } + } + return step; +} + +std::optional SessionTable::sample( + SessionId session, + ::executorch::aten::Tensor& logits, + int row) { + const auto it = sessions_.find(session); + if (it == sessions_.end() || it->second.sampler == nullptr) { + ET_LOG( + Error, "sample: session %" PRId64 " has no sampling policy", session); + return std::nullopt; + } + if (row < 0 || row >= logits.numel() / vocab_size_) { + ET_LOG(Error, "sample: logits hold no row %d", row); + return std::nullopt; + } + // A one-row view over the model's own output: sample_from_logits reduces in + // place and reads the last dimension. + auto one_row = make_tensor_ptr( + {vocab_size_}, + static_cast(logits.mutable_data_ptr()) + + static_cast(row) * vocab_size_ * + ::executorch::runtime::elementSize(logits.scalar_type()), + logits.scalar_type()); + return static_cast(sample_from_logits(*one_row, *it->second.sampler)); +} + +} // namespace util +} // namespace batching +} // namespace llm +} // namespace extension +} // namespace executorch diff --git a/extension/llm/batching/util/session_table.h b/extension/llm/batching/util/session_table.h new file mode 100644 index 00000000000..ef2fd9546af --- /dev/null +++ b/extension/llm/batching/util/session_table.h @@ -0,0 +1,121 @@ +/* + * Copyright (c) Meta Platforms, Inc. and affiliates. + * All rights reserved. + * + * This source code is licensed under the BSD-style license found in the + * LICENSE file in the root directory of this source tree. + */ + +#pragma once + +// What every Executor over a multi-sequence cache does the same way, whatever +// runs its forwards: mapping sessions to cache sequences, packing a batch onto +// one token axis, and sampling each session's rows with its own policy. + +#include +#include +#include +#include +#include +#include + +#include +#include +#include +#include +#include // ET_EXPERIMENTAL + +namespace executorch { +namespace extension { +namespace llm { + +class Sampler; + +namespace batching { +namespace util { + +namespace cache = ::executorch::extension::llm::cache; + +// A batch flattened onto one token axis, in input order. +struct ET_EXPERIMENTAL PackedStep { + std::vector tokens; + std::vector positions; + std::vector seq_ids; + // Per input: the token whose logits it wants, or -1 for none. + std::vector logit_indices; +}; + +// The logits rows a forward over tokens [offset, offset + length) must select: +// `selector` indexes that forward's tokens, `inputs` names the batch input +// each row belongs to. A forward with no wanted row still selects its last +// token, which nothing reads, since the program produces at least one row. +struct ET_EXPERIMENTAL SliceRows { + std::vector selector; + std::vector inputs; +}; + +ET_EXPERIMENTAL SliceRows +select_rows(const PackedStep& step, int offset, int length); + +class ET_EXPERIMENTAL SessionTable { + public: + // `control` must outlive the table. `max_sessions` counts every resident + // session, clones included; `max_session_tokens` bounds each one. + SessionTable( + cache::BatchControl& control, + int max_sessions, + int max_session_tokens, + std::int32_t vocab_size); + ~SessionTable(); + + SessionTable(const SessionTable&) = delete; + SessionTable& operator=(const SessionTable&) = delete; + + std::optional open(); + // Frees the session's cells and hands its sequence id back. Session ids + // are never reused. + void close(SessionId session); + // A new session holding the source's [0, upto). + std::optional clone(SessionId source, Position upto); + void set_sampling( + SessionId session, + const SamplingParams& params, + std::optional seed); + + // Flattens the batch and truncates whatever it reopens. A per-sequence + // cursor carries the batch's own writes, so consecutive chunks of one prompt + // abut and only the first can reopen committed ground. Every input is + // checked before any sequence is truncated, so a refusal leaves the cache + // untouched. Declaring each forward's tokens is the caller's, per forward. + ::executorch::runtime::Result pack(const BatchInput& batch); + + // Draws the session's token from `row` of `logits`, [rows, vocab], which + // the sampler reduces in place. + std::optional + sample(SessionId session, ::executorch::aten::Tensor& logits, int row); + + std::size_t size() const { + return sessions_.size(); + } + + private: + struct Session { + std::int32_t seq_id; + std::unique_ptr sampler; + }; + + std::optional publish(std::int32_t seq_id, Position position); + + cache::BatchControl& ctl_; + const int max_sessions_; + const int max_session_tokens_; + const std::int32_t vocab_size_; + SessionId next_session_ = 1; + std::unordered_map sessions_; +}; + +} // namespace util +} // namespace batching +} // namespace llm +} // namespace extension +} // namespace executorch