Skip to content
Open
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
2 changes: 1 addition & 1 deletion CMakeLists.txt
Original file line number Diff line number Diff line change
@@ -1,3 +1,3 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
# Copyright 2024-2026 Arm Limited and/or its affiliates.
Expand Down Expand Up @@ -1013,7 +1013,7 @@
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)
Expand Down
5 changes: 2 additions & 3 deletions backends/cuda/batching/CMakeLists.txt
Original file line number Diff line number Diff line change
@@ -1,3 +1,3 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
#
Expand All @@ -18,9 +18,8 @@
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(
Expand Down
240 changes: 14 additions & 226 deletions backends/cuda/batching/cuda_executor.cpp
Original file line number Diff line number Diff line change
@@ -1,3 +1,3 @@
/*
* Copyright (c) Meta Platforms, Inc. and affiliates.
* All rights reserved.
Expand All @@ -11,34 +11,28 @@
#include <algorithm>
#include <cinttypes>
#include <limits>
#include <new>
#include <random>
#include <utility>
#include <vector>

#include <executorch/backends/cuda/batching/step_plan.h>
#include <executorch/backends/cuda/runtime/backend_options.h>
#include <executorch/extension/llm/runner/model_metadata.h>
#include <executorch/extension/llm/sampler/sampler.h>
#include <executorch/extension/llm/sampler/util.h>
#include <executorch/extension/tensor/tensor.h>
#include <executorch/runtime/backend/backend_options_map.h>
#include <executorch/runtime/backend/interface.h>
#include <executorch/runtime/backend/options.h>
#include <executorch/runtime/core/exec_aten/util/scalar_type_util.h>
#include <executorch/runtime/platform/log.h>

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;
Expand All @@ -52,23 +46,6 @@
// 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 ||
Expand Down Expand Up @@ -174,11 +151,10 @@
module_(std::move(module)),
ctl_(cache->as<llm_cache::BatchControl>()),
kv_(cache->as<CudaKVCache>()),
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;

Expand Down Expand Up @@ -333,185 +309,30 @@
return true;
}

// Copied from ModuleExecutor::build_step; see there for the reasoning.
Result<CudaExecutor::Step> 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<std::pair<std::int32_t, int>> rewinds;
std::unordered_map<std::int32_t, int> 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<std::int64_t>(input.position) +
static_cast<std::int64_t>(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<int>(start));
at = static_cast<int>(start);
}

const std::int64_t end = start + static_cast<std::int64_t>(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<std::int64_t>(slice[k]));
step.positions.push_back(start + static_cast<std::int64_t>(k));
}
step.seq_ids.insert(step.seq_ids.end(), input.size, seq_id);
at = static_cast<int>(end);
step.logit_indices.push_back(
input.produce_output ? static_cast<int>(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<SessionId> CudaExecutor::open_session() {
if (sessions_.size() >= static_cast<std::size_t>(max_sessions_) ||
next_session_ == 0) {
return std::nullopt;
}
#if ET_HAS_EXCEPTIONS
try {
#endif
const std::optional<std::int32_t> 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<SessionId> 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<SessionId>::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<SessionId> 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<std::size_t>(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<std::uint64_t> seed) {
const auto it = sessions_.find(session);
if (it == sessions_.end()) {
return;
}
it->second.sampler = std::make_unique<Sampler>(
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> step = build_step(batch);
const Result<llm_batching::util::PackedStep> step = sessions_.pack(batch);
if (!step.ok()) {
return false;
}
Expand Down Expand Up @@ -540,21 +361,9 @@
{n},
std::vector<std::int64_t>(
step->positions.begin() + off, step->positions.begin() + off + n));
std::vector<std::int64_t> selector_values;
std::vector<std::size_t> 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<int>(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<int>(selected.selector.size());
auto selector = make_tensor_ptr({rows}, std::move(selected.selector));

auto result = module_->execute(method, {tokens, positions, selector});
if (!result.ok()) {
Expand All @@ -580,11 +389,11 @@
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> token =
sample_row(logits, static_cast<int>(row), session);
sessions_.sample(session, logits, static_cast<int>(row));
if (!token) {
return false;
}
Expand All @@ -594,27 +403,6 @@
return true;
}

std::optional<Token> 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<std::uint8_t*>(logits.mutable_data_ptr()) +
static_cast<std::size_t>(row) * vocab_size_ *
::executorch::runtime::elementSize(logits.scalar_type()),
logits.scalar_type());
return static_cast<Token>(
::executorch::extension::llm::sample_from_logits(
*one_row, *it->second.sampler));
}

OffGraphKVMetrics CudaExecutor::kv_metrics() const {
return kv_->metrics();
}
Expand Down
Loading
Loading