From cf57e1ceb7bec3021d211a1f382d396170749e3f Mon Sep 17 00:00:00 2001 From: ajmalrasi Date: Mon, 14 Sep 2026 21:09:27 +0530 Subject: [PATCH 1/7] feat: add persistent sequence step runtime Signed-off-by: ajmalrasi --- cpp/runtime/llmRankRuntime.h | 3 + cpp/runtime/sequenceStepRuntime.cpp | 296 +++++++++ cpp/runtime/sequenceStepRuntime.h | 64 ++ cpp/runtime/state/sequenceSlots.cpp | 164 +++++ cpp/runtime/state/sequenceSlots.h | 121 ++++ examples/llm/CMakeLists.txt | 11 + examples/llm/continuousBatchingProbe.cpp | 619 ++++++++++++++++++ examples/llm/continuousBatchingProbe.md | 52 ++ examples/llm/sequenceStepRuntime.md | 82 +++ .../cpp/runtime/state/sequenceSlotsTest.cpp | 159 +++++ 10 files changed, 1571 insertions(+) create mode 100644 cpp/runtime/sequenceStepRuntime.cpp create mode 100644 cpp/runtime/sequenceStepRuntime.h create mode 100644 cpp/runtime/state/sequenceSlots.cpp create mode 100644 cpp/runtime/state/sequenceSlots.h create mode 100644 examples/llm/continuousBatchingProbe.cpp create mode 100644 examples/llm/continuousBatchingProbe.md create mode 100644 examples/llm/sequenceStepRuntime.md create mode 100644 unittests/cpp/runtime/state/sequenceSlotsTest.cpp diff --git a/cpp/runtime/llmRankRuntime.h b/cpp/runtime/llmRankRuntime.h index 984caff57..195728e9a 100644 --- a/cpp/runtime/llmRankRuntime.h +++ b/cpp/runtime/llmRankRuntime.h @@ -240,6 +240,9 @@ class LLMRankRuntime } private: + friend class ContinuousBatchingProbe; + friend class SequenceStepRuntime; + void initializeFromEngineDir(std::string const& engineDir, std::string const& multimodalEngineDir, std::unordered_map const& loraWeightsMap, std::optional const& draftingConfig, cudaStream_t stream, diff --git a/cpp/runtime/sequenceStepRuntime.cpp b/cpp/runtime/sequenceStepRuntime.cpp new file mode 100644 index 000000000..4abac3b9f --- /dev/null +++ b/cpp/runtime/sequenceStepRuntime.cpp @@ -0,0 +1,296 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ +#include "runtime/sequenceStepRuntime.h" +#include "common/bindingNames.h" +#include "kernels/embeddingKernels/embeddingKernels.h" +#include "kernels/posEncoding/initializeCosSinCache.h" +#include "runtime/llmRankRuntime.h" + +#include +#include +#include + +namespace trt_edgellm +{ +namespace rt +{ +namespace +{ +void require(bool condition, char const* message) +{ + if (!condition) + { + throw std::logic_error(message); + } +} +} // namespace + +struct SequenceStepRuntime::Lease +{ + explicit Lease(std::atomic& gate) + : gate(gate) + { + bool expected = false; + require(gate.compare_exchange_strong(expected, true, std::memory_order_acquire), "Runtime already leased"); + } + ~Lease() + { + if (!poisoned) + { + gate.store(false, std::memory_order_release); + } + } + std::atomic& gate; + bool poisoned{}; +}; + +struct SequenceStepRuntime::View +{ + TensorMap bindings; + std::map rows; +}; + +SequenceStepRuntime::SequenceStepRuntime(LLMRankRuntime& runtime, cudaStream_t stream) + : mRuntime(runtime) + , mStream(stream) + , mLease(std::make_unique(runtime.mHandleRequestInProgress)) + , mSlots(2, runtime.mDeployment.base.maxSupportedInputLength, runtime.mDeployment.base.maxKVCacheCapacity) + , mHostStarts({2}, DeviceType::kCPU, nvinfer1::DataType::kINT32) + , mDeviceStarts({2}, DeviceType::kGPU, nvinfer1::DataType::kINT32) + , mHostSelect({2, 1}, DeviceType::kCPU, nvinfer1::DataType::kINT64) +{ + auto const& config = runtime.mDeployment.base; + require(stream != nullptr, "Step runtime requires an explicit stream"); + require(runtime.mMaxRuntimeBatchSize == 2 && !runtime.mMapping.isParallel() && config.modelType == "qwen3_5_text" + && !runtime.hasDraftModel() && !runtime.mContextCache && config.maxSupportedLoraRank == 0 + && !config.isDiffusionBackbone && !config.useVisionBidirectionalAttention && !config.useContextDependentRope + && !config.useDualRope && !runtime.mDeepstack && !runtime.mGemma4Ple && config.reducedVocabSize == 0, + "Step runtime supports two-slot vanilla text-only Qwen3.5 without context reuse or LoRA"); + require(runtime.mSharedResources->kvPageTables[0]->isIdentity(), "Step session requires identity page ownership"); + for (int32_t index = 0; index < 3; ++index) + { + auto view = std::make_unique(); + view->bindings = runtime.mBaseTensorMap; + int32_t const first = index == 1 ? 1 : 0; + int32_t const count = index == 2 ? 2 : 1; + auto bindRow = [&](std::string const& name) { + Tensor* original = runtime.mBaseTensorMap.get(name); + require(original != nullptr, "Missing selected-row binding"); + auto shape = original->getShape(); + require(shape[0] == 2, "Step session must be constructed before legacy execution reshapes state"); + size_t const bytes = original->getMemoryCapacity() / 2; + shape[0] = count; + auto* pointer = static_cast(original->rawPointer()) + first * bytes; + auto inserted = view->rows.emplace(name, Tensor(pointer, shape, DeviceType::kGPU, original->getDataType())); + view->bindings.set(name, inserted.first->second); + }; + for (int32_t layer = 0; layer < config.numLinearAttnLayers; ++layer) + { + bindRow(binding_names::formatRecurrentStateName(layer, true)); + bindRow(binding_names::formatRecurrentStateName(layer, false)); + bindRow(binding_names::formatConvStateName(layer, true)); + bindRow(binding_names::formatConvStateName(layer, false)); + } + bindRow(binding_names::kKVPageTable); + if (config.ropeConfig.type == RopeType::kMRope) + { + bindRow(binding_names::kRopeCosSin); + } + view->bindings.set(binding_names::kKVCacheStartIndex, mDeviceStarts); + mViews[index] = std::move(view); + } + CUDA_CHECK(cudaEventCreateWithFlags(&mComplete, cudaEventDisableTiming)); +} + +SequenceStepRuntime::~SequenceStepRuntime() +{ + // Drain borrowed buffers before releasing the lease, including partially enqueued failed steps. + if (cudaStreamSynchronize(mStream) != cudaSuccess) + { + mLease->poisoned = true; + } + if (mComplete != nullptr) + { + cudaEventDestroy(mComplete); + } +} + +bool SequenceStepRuntime::healthy() const noexcept +{ + return !mLease->poisoned; +} + +void SequenceStepRuntime::requireIdle() const +{ + require(healthy() && !mPending, "Step runtime is failed or has an unfinished forward"); +} + +SequenceState const& SequenceStepRuntime::state(SequenceHandle handle) const +{ + return mSlots.get(handle); +} + +SequenceHandle SequenceStepRuntime::acquire(uint64_t requestId, std::vector prompt, SequenceOptions options) +{ + requireIdle(); + int64_t const vocabulary = mRuntime.mEmbedding.table.getShape()[0]; + for (int32_t token : prompt) + { + require(token >= 0 && token < vocabulary, "Prompt token outside embedding vocabulary"); + } + auto handle = mSlots.acquire(requestId, std::move(prompt), std::move(options)); + try + { + mRuntime.zeroRecurrentStates(handle.slot, mStream); + } + catch (...) + { + mLease->poisoned = true; + throw; + } + return handle; +} + +Tensor const& SequenceStepRuntime::beginPrefill(SequenceHandle handle, int32_t count) +{ + requireIdle(); + auto const& sequence = state(handle); + require(sequence.phase() == SequencePhase::kPrefill && count > 0 + && count <= static_cast(sequence.prompt().size()) - sequence.promptCursor(), + "Invalid prefill span"); + return enqueue({handle, {}}, 1, count, false); +} + +Tensor const& SequenceStepRuntime::beginDecode(std::array const& handles, int32_t count) +{ + requireIdle(); + require(count == 1 || count == 2, "Invalid decode batch size"); + for (int32_t i = 0; i < count; ++i) + { + auto const& sequence = state(handles[i]); + require(sequence.phase() == SequencePhase::kDecode, "Decode requires a pending sampled token"); + require(i == 0 || handles[i].slot == handles[i - 1].slot + 1, "Slots must be distinct and physically ordered"); + } + return enqueue(handles, count, 1, true); +} + +Tensor const& SequenceStepRuntime::enqueue( + std::array const& handles, int32_t count, int32_t span, bool decode) +{ + auto& io = *mRuntime.mPipelineIO; + auto const& config = mRuntime.mDeployment.base; + // A nonempty-cache S=1 binding selects attention decode and requires absolute context lengths. + bool const executeDecode = decode || (span == 1 && state(handles[0]).committedTokens() > 0); + int32_t const physicalSpan = span; + require(mRuntime.mIdsInput.reshape({count, physicalSpan}) + && mRuntime.mHostPackedTokenIds.reshape({count, physicalSpan}) + && io.inputsEmbeds.reshape({count, physicalSpan, config.hiddenSize}) + && io.outputLogits.reshape({count, config.outputVocabSize}) && io.contextLengths.reshape({count}) + && io.hostContextLengths.reshape({count}) && io.selectTokenIndices.reshape({count, 1}), + "Step tensor shape exceeds allocated capacity"); + auto* packed = mRuntime.mHostPackedTokenIds.dataPointer(); + std::fill_n(packed, count * physicalSpan, 0); + for (int32_t row = 0; row < count; ++row) + { + auto const& sequence = state(handles[row]); + mHostStarts.dataPointer()[row] = sequence.committedTokens(); + io.hostContextLengths.dataPointer()[row] = executeDecode ? sequence.committedTokens() + 1 : span; + mHostSelect.dataPointer()[row] = executeDecode ? 0 : span - 1; + if (decode) + { + packed[row] = sequence.output().back(); + } + else + { + std::copy_n(sequence.prompt().data() + sequence.promptCursor(), span, packed + row * physicalSpan); + } + } + mPending = true; + mHandles = handles; + mCount = count; + mSpan = span; + mDecode = decode; + try + { + // Text positions are request-independent; initialize once before the first forward in this session. + if (!mRopeInitialized && config.ropeConfig.type == RopeType::kMRope) + { + kernel::initializeTextOnlyMRopeCosSin(io.mropeCosSin.dataPointer(), config.ropeConfig.rotaryTheta, + config.rotaryDim, config.maxKVCacheCapacity, 2, mStream); + mRopeInitialized = true; + } + CUDA_CHECK(cudaMemcpyAsync(mRuntime.mIdsInput.rawPointer(), packed, + static_cast(count) * physicalSpan * sizeof(int32_t), cudaMemcpyHostToDevice, mStream)); + CUDA_CHECK(cudaMemcpyAsync(mDeviceStarts.rawPointer(), mHostStarts.rawPointer(), count * sizeof(int32_t), + cudaMemcpyHostToDevice, mStream)); + CUDA_CHECK(cudaMemcpyAsync(io.contextLengths.rawPointer(), io.hostContextLengths.rawPointer(), + count * sizeof(int32_t), cudaMemcpyHostToDevice, mStream)); + CUDA_CHECK(cudaMemcpyAsync(io.selectTokenIndices.rawPointer(), mHostSelect.rawPointer(), + count * sizeof(int64_t), cudaMemcpyHostToDevice, mStream)); + kernel::embeddingLookup(mRuntime.mIdsInput, mRuntime.mEmbedding.table, mRuntime.mEmbedding.scalesAsOptional(), + io.inputsEmbeds, mStream); + auto const dims = executeDecode + ? config.decodeDims(count) + : config.prefillDims(count, physicalSpan, state(handles[0]).committedTokens() == 0); + auto const& view = *mViews[count == 2 ? 2 : handles[0].slot]; + require(mRuntime.mBaseExecutor->prepare(executeDecode ? 1 : 0, dims, view.bindings, mStream), + "Step engine preparation failed"); + require(mRuntime.mBaseExecutor->execute(mStream), "Step engine execution failed"); + CUDA_CHECK(cudaEventRecord(mComplete, mStream)); + } + catch (...) + { + mLease->poisoned = true; + throw; + } + return io.outputLogits; +} + +void SequenceStepRuntime::completeStep() +{ + require(healthy() && mPending, "No healthy pending step"); + try + { + CUDA_CHECK(cudaEventSynchronize(mComplete)); + for (int32_t row = 0; row < mCount; ++row) + { + if (mDecode) + { + mSlots.commitDecode(mHandles[row]); + } + else + { + mSlots.commitPrompt(mHandles[row], mSpan); + } + } + mPending = false; + } + catch (...) + { + mLease->poisoned = true; + throw; + } +} + +void SequenceStepRuntime::acceptToken(SequenceHandle handle, int32_t token, uint64_t randomDraws) +{ + requireIdle(); + require(token >= 0 && token < mRuntime.mDeployment.base.outputVocabSize, "Sample outside vocabulary"); + mSlots.acceptToken(handle, token, randomDraws); +} + +void SequenceStepRuntime::finish(SequenceHandle handle) +{ + requireIdle(); + mSlots.finish(handle); +} + +void SequenceStepRuntime::release(SequenceHandle handle) +{ + requireIdle(); + mSlots.release(handle); +} +} // namespace rt +} // namespace trt_edgellm diff --git a/cpp/runtime/sequenceStepRuntime.h b/cpp/runtime/sequenceStepRuntime.h new file mode 100644 index 000000000..30ac79686 --- /dev/null +++ b/cpp/runtime/sequenceStepRuntime.h @@ -0,0 +1,64 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ +#pragma once + +#include "common/tensor.h" +#include "runtime/state/sequenceSlots.h" +#include +#include + +namespace trt_edgellm +{ +namespace rt +{ +class LLMRankRuntime; + +//! Exclusive, single-threaded text-only step session borrowing one rank runtime and its stream. +//! Runtime and stream must outlive the session. Do not invoke other runtime APIs during this lease. +//! Engine/CUDA failure poisons the lease; recreate the parent runtime before further use. +class SequenceStepRuntime +{ +public: + SequenceStepRuntime(LLMRankRuntime& runtime, cudaStream_t stream); + ~SequenceStepRuntime(); + SequenceStepRuntime(SequenceStepRuntime const&) = delete; + SequenceStepRuntime& operator=(SequenceStepRuntime const&) = delete; + + SequenceHandle acquire(uint64_t requestId, std::vector prompt, SequenceOptions options); + SequenceState const& state(SequenceHandle handle) const; + //! Enqueue a prompt span or pending output tokens. No sampling, output publication or CPU state commit occurs. + //! Logits are borrowed until the next begin call; they become readable after completion on the supplied stream. + Tensor const& beginPrefill(SequenceHandle handle, int32_t count); + Tensor const& beginDecode(std::array const& handles, int32_t count); + //! Wait for forward completion before publishing endpoints and reusing pinned staging or physical slots. + void completeStep(); + void acceptToken(SequenceHandle handle, int32_t token, uint64_t randomDraws = 0); + void finish(SequenceHandle handle); + void release(SequenceHandle handle); + bool healthy() const noexcept; + +private: + struct Lease; + struct View; + void requireIdle() const; + Tensor const& enqueue(std::array const& handles, int32_t count, int32_t span, bool decode); + LLMRankRuntime& mRuntime; + cudaStream_t mStream; + std::unique_ptr mLease; + SequenceSlots mSlots; + std::array, 3> mViews; + Tensor mHostStarts; + Tensor mDeviceStarts; + Tensor mHostSelect; + cudaEvent_t mComplete{}; + bool mPending{}; + bool mRopeInitialized{}; + bool mDecode{}; + int32_t mCount{}; + int32_t mSpan{}; + std::array mHandles{}; +}; +} // namespace rt +} // namespace trt_edgellm diff --git a/cpp/runtime/state/sequenceSlots.cpp b/cpp/runtime/state/sequenceSlots.cpp new file mode 100644 index 000000000..aba454f76 --- /dev/null +++ b/cpp/runtime/state/sequenceSlots.cpp @@ -0,0 +1,164 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ +#include "runtime/state/sequenceSlots.h" + +#include +#include +#include +#include +#include + +namespace trt_edgellm +{ +namespace rt +{ +namespace +{ +uint64_t nextOwner() +{ + static std::atomic sNext{1}; + uint64_t candidate = sNext.load(std::memory_order_relaxed); + do + { + if (candidate == std::numeric_limits::max()) + { + throw std::overflow_error("Sequence owner IDs exhausted"); + } + } while (!sNext.compare_exchange_weak(candidate, candidate + 1, std::memory_order_relaxed)); + return candidate; +} +} // namespace + +SequenceSlots::SequenceSlots(int32_t slots, int32_t maxInputTokens, int32_t maxSequenceTokens) + : mOwner(nextOwner()) + , mMaxInputTokens(maxInputTokens) + , mMaxSequenceTokens(maxSequenceTokens) +{ + if (slots <= 0 || maxInputTokens <= 0 || maxSequenceTokens < maxInputTokens) + { + throw std::invalid_argument("Invalid sequence slot limits"); + } + mSlots.resize(slots); +} + +SequenceHandle SequenceSlots::acquire(uint64_t requestId, std::vector prompt, SequenceOptions options) +{ + if (requestId == 0 || prompt.empty() || prompt.size() > static_cast(mMaxInputTokens) + || options.maxOutputTokens <= 0 + || static_cast(prompt.size()) + options.maxOutputTokens > mMaxSequenceTokens + || !std::isfinite(options.temperature) || options.temperature < 0.0F || !std::isfinite(options.topP) + || options.topP <= 0.0F || options.topP > 1.0F || options.topK < 0 || options.numLogprobs < 0) + { + throw std::invalid_argument("Invalid sequence request or capacity"); + } + for (auto const& state : mSlots) + { + if (state.mPhase != SequencePhase::kFree && state.mRequestId == requestId) + { + throw std::invalid_argument("Duplicate live request ID"); + } + } + for (size_t slot = 0; slot < mSlots.size(); ++slot) + { + auto& state = mSlots[slot]; + if (state.mPhase != SequencePhase::kFree) + { + continue; + } + if (state.mGeneration == std::numeric_limits::max()) + { + continue; + } + SequenceState next; + next.mGeneration = state.mGeneration + 1; + next.mRequestId = requestId; + next.mPrompt = std::move(prompt); + next.mOptions = std::move(options); + next.mOutput.reserve(next.mOptions.maxOutputTokens); + next.mPhase = SequencePhase::kPrefill; + state = std::move(next); + return {mOwner, state.mGeneration, static_cast(slot)}; + } + throw std::runtime_error("No free sequence slot"); +} + +SequenceState const& SequenceSlots::get(SequenceHandle handle) const +{ + if (handle.owner != mOwner || handle.slot < 0 || static_cast(handle.slot) >= mSlots.size()) + { + throw std::invalid_argument("Foreign or invalid sequence handle"); + } + auto const& state = mSlots[handle.slot]; + if (state.mPhase == SequencePhase::kFree || state.mGeneration != handle.generation) + { + throw std::invalid_argument("Stale sequence handle"); + } + return state; +} + +SequenceState& SequenceSlots::checked(SequenceHandle handle) +{ + get(handle); + return mSlots[handle.slot]; +} + +void SequenceSlots::commitPrompt(SequenceHandle handle, int32_t count) +{ + auto& state = checked(handle); + if (state.mPhase != SequencePhase::kPrefill || count <= 0 + || count > static_cast(state.mPrompt.size()) - state.mPromptCursor) + { + throw std::logic_error("Invalid prompt commit"); + } + state.mPromptCursor += count; + state.mCommittedTokens += count; + if (state.mPromptCursor == static_cast(state.mPrompt.size())) + { + state.mPhase = SequencePhase::kAwaitingSample; + } +} + +void SequenceSlots::commitDecode(SequenceHandle handle) +{ + auto& state = checked(handle); + if (state.mPhase != SequencePhase::kDecode + || state.mCommittedTokens + 1 != static_cast(state.mPrompt.size() + state.mOutput.size()) + || state.mCommittedTokens >= mMaxSequenceTokens) + { + throw std::logic_error("Decode requires exactly one uncommitted output token"); + } + ++state.mCommittedTokens; + state.mPhase = SequencePhase::kAwaitingSample; +} + +void SequenceSlots::acceptToken(SequenceHandle handle, int32_t token, uint64_t randomDraws) +{ + auto& state = checked(handle); + if (state.mPhase != SequencePhase::kAwaitingSample || token < 0 + || randomDraws > std::numeric_limits::max() - state.mRandomCounter) + { + throw std::logic_error("Invalid sampled-token acceptance"); + } + state.mOutput.push_back(token); + state.mRandomCounter += randomDraws; + state.mPhase = static_cast(state.mOutput.size()) == state.mOptions.maxOutputTokens + ? SequencePhase::kFinished + : SequencePhase::kDecode; +} + +void SequenceSlots::finish(SequenceHandle handle) +{ + checked(handle).mPhase = SequencePhase::kFinished; +} + +void SequenceSlots::release(SequenceHandle handle) +{ + auto& state = checked(handle); + auto const generation = state.mGeneration; + state = SequenceState{}; + state.mGeneration = generation; +} +} // namespace rt +} // namespace trt_edgellm diff --git a/cpp/runtime/state/sequenceSlots.h b/cpp/runtime/state/sequenceSlots.h new file mode 100644 index 000000000..aeb250477 --- /dev/null +++ b/cpp/runtime/state/sequenceSlots.h @@ -0,0 +1,121 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ +#pragma once + +#include +#include +#include +#include + +namespace trt_edgellm +{ +namespace rt +{ +//! Handles are valid only for one allocation in one slot pool. +struct SequenceHandle +{ + uint64_t owner{}; + uint64_t generation{}; + int32_t slot{-1}; +}; + +//! Request-local controls; sampling and stopping policy are applied by the scheduler, not the step executor. +struct SequenceOptions +{ + int32_t maxOutputTokens{128}; + float temperature{1.0F}; + float topP{1.0F}; + int64_t topK{}; + uint64_t seed{42}; + int32_t numLogprobs{}; + bool enableThinking{}; + std::vector eosTokenIds; + std::vector stopStrings; + std::unordered_map logitBias; +}; + +enum class SequencePhase +{ + kFree, + kPrefill, + kAwaitingSample, + kDecode, + kFinished, +}; + +//! Persistent logical state; the physical slot owns its KV, recurrent and convolution rows. +class SequenceState +{ +public: + uint64_t requestId() const noexcept + { + return mRequestId; + } + SequencePhase phase() const noexcept + { + return mPhase; + } + int32_t promptCursor() const noexcept + { + return mPromptCursor; + } + int32_t committedTokens() const noexcept + { + return mCommittedTokens; + } + uint64_t randomCounter() const noexcept + { + return mRandomCounter; + } + std::vector const& prompt() const noexcept + { + return mPrompt; + } + std::vector const& output() const noexcept + { + return mOutput; + } + SequenceOptions const& options() const noexcept + { + return mOptions; + } + +private: + friend class SequenceSlots; + uint64_t mRequestId{}; + uint64_t mGeneration{}; + SequencePhase mPhase{SequencePhase::kFree}; + int32_t mPromptCursor{}; + int32_t mCommittedTokens{}; + uint64_t mRandomCounter{}; + std::vector mPrompt; + std::vector mOutput; + SequenceOptions mOptions; +}; + +//! Single-owner bookkeeping with no CUDA dependency; allocation occurs only at admission. +class SequenceSlots +{ +public: + SequenceSlots(int32_t slots, int32_t maxInputTokens, int32_t maxSequenceTokens); + SequenceSlots(SequenceSlots const&) = delete; + SequenceSlots& operator=(SequenceSlots const&) = delete; + SequenceHandle acquire(uint64_t requestId, std::vector prompt, SequenceOptions options); + SequenceState const& get(SequenceHandle handle) const; + void commitPrompt(SequenceHandle handle, int32_t count); + void commitDecode(SequenceHandle handle); + void acceptToken(SequenceHandle handle, int32_t token, uint64_t randomDraws = 0); + void finish(SequenceHandle handle); + void release(SequenceHandle handle); + +private: + SequenceState& checked(SequenceHandle handle); + uint64_t mOwner; + int32_t mMaxInputTokens; + int32_t mMaxSequenceTokens; + std::vector mSlots; +}; +} // namespace rt +} // namespace trt_edgellm diff --git a/examples/llm/CMakeLists.txt b/examples/llm/CMakeLists.txt index b4256f3d8..f9f213cbb 100644 --- a/examples/llm/CMakeLists.txt +++ b/examples/llm/CMakeLists.txt @@ -10,6 +10,17 @@ # prohibited. # Build executables for chat, benchmark and accuracy mode +option(BUILD_CONTINUOUS_BATCHING_PROBE + "Build the model-dependent eager slot/chunk feasibility probe" OFF) +if(BUILD_CONTINUOUS_BATCHING_PROBE) + add_executable(continuous_batching_probe continuousBatchingProbe.cpp) + target_link_libraries(continuous_batching_probe PRIVATE edgellmCore + commonLibraryExt) + target_include_directories(continuous_batching_probe + PRIVATE ${COMMON_INCLUDE_DIRS}) + add_cross_build_link_options(continuous_batching_probe) +endif() + add_executable(llm_build llm_build.cpp) target_link_libraries(llm_build PRIVATE ${NV_ONNX_PARSER_LIB} commonLibraryExt edgellmBuilder) diff --git a/examples/llm/continuousBatchingProbe.cpp b/examples/llm/continuousBatchingProbe.cpp new file mode 100644 index 000000000..29bb5322e --- /dev/null +++ b/examples/llm/continuousBatchingProbe.cpp @@ -0,0 +1,619 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#include "common/bindingNames.h" +#include "common/trtUtils.h" +#include "kernels/posEncoding/initializeCosSinCache.h" +#include "runtime/llmRankRuntime.h" +#include "runtime/sequenceStepRuntime.h" + +#include +#include +#include +#include +#include +#include +#include + +namespace trt_edgellm +{ +namespace rt +{ +namespace +{ +void require(bool condition, std::string const& message) +{ + if (!condition) + { + throw std::runtime_error(message); + } +} + +int32_t greedy(std::vector const& logits) +{ + return static_cast(std::max_element(logits.begin(), logits.end()) - logits.begin()); +} + +bool compare(std::string const& label, std::vector const& reference, std::vector const& candidate, + bool exact = false) +{ + require(reference.size() == candidate.size(), "Logit count mismatch"); + double squaredError = 0.0; + double squaredReference = 0.0; + float maxError = 0.0F; + bool finite = true; + for (size_t i = 0; i < reference.size(); ++i) + { + finite = finite && std::isfinite(reference[i]) && std::isfinite(candidate[i]); + double const error = static_cast(candidate[i]) - reference[i]; + squaredError += error * error; + squaredReference += static_cast(reference[i]) * reference[i]; + maxError = std::max(maxError, static_cast(std::abs(error))); + } + double const relativeL2 = std::sqrt(squaredError / std::max(squaredReference, 1.0e-30)); + bool const sameGreedy = greedy(reference) == greedy(candidate); + // Preliminary FP16 correctness screen, not a model-quality acceptance criterion. + bool const passed = finite && sameGreedy && (exact ? maxError == 0.0F : maxError <= 0.1F && relativeL2 <= 0.005); + std::cout << "COMPARE " << label << " max_abs=" << maxError << " relative_l2=" << relativeL2 + << " reference_greedy=" << greedy(reference) << " candidate_greedy=" << greedy(candidate) + << " passed=" << passed << std::endl; + return passed; +} +} // namespace + +//! Model-dependent, eager-only feasibility probe; deliberately not a serving API. +class ContinuousBatchingProbe +{ +public: + ContinuousBatchingProbe(LLMRankRuntime& runtime, cudaStream_t stream) + : mRuntime(runtime) + , mStream(stream) + , mOriginalMap(runtime.mBaseTensorMap) + , mHostLengths({2}, DeviceType::kCPU, nvinfer1::DataType::kINT32) + , mHostLogits({2, runtime.mDeployment.base.outputVocabSize}, DeviceType::kCPU, nvinfer1::DataType::kFLOAT) + , mScratch({kCOPY_BYTES}, DeviceType::kCPU, nvinfer1::DataType::kUINT8) + { + auto const& config = runtime.mDeployment.base; + require(runtime.mMaxRuntimeBatchSize == 2 && !runtime.hasDraftModel(), "Requires two-slot vanilla engine"); + require(config.maxSupportedLoraRank == 0 && !config.useVisionBidirectionalAttention + && config.reducedVocabSize == 0 && !config.isDiffusionBackbone, + "Unsupported probe configuration"); + require(runtime.mPipelineIO->outputLogits.getDataType() == nvinfer1::DataType::kFLOAT, + "Probe requires FP32 logits"); + require(runtime.mSharedResources->kvPageTables[0]->isIdentity(), "Probe requires identity KV pages"); + if (config.ropeConfig.type == RopeType::kMRope) + { + auto& rope = runtime.mPipelineIO->mropeCosSin; + require(rope.reshape({2, config.maxKVCacheCapacity, config.rotaryDim}), "RoPE shape"); + kernel::initializeTextOnlyMRopeCosSin(rope.dataPointer(), config.ropeConfig.rotaryTheta, + config.rotaryDim, config.maxKVCacheCapacity, 2, stream); + } + auto& mamba = cache().getMambaCacheManager(); + require(mamba.numLayers() > 0, "Probe requires hybrid recurrent state"); + for (int32_t layer = 0; layer < mamba.numLayers(); ++layer) + { + addBinding(binding_names::formatRecurrentStateName(layer, true)); + addBinding(binding_names::formatRecurrentStateName(layer, false)); + addBinding(binding_names::formatConvStateName(layer, true)); + addBinding(binding_names::formatConvStateName(layer, false)); + } + addBinding(binding_names::kKVPageTable); + reset(0); + reset(1); + memory("initialized"); + } + + ~ContinuousBatchingProbe() + { + mRuntime.mBaseTensorMap = mOriginalMap; + } + + //! Reset only the selected sequence's recurrent state and logical endpoint. + void reset(int32_t slot) + { + require(slot >= 0 && slot < 2, "Invalid slot"); + mRuntime.zeroRecurrentStates(slot, mStream); + mCommitted[slot] = 0; + } + + //! Execute one contiguous physical-slot view using existing native primitives. + std::vector> step( + int32_t firstSlot, std::vector> const& tokens, bool decode = false) + { + int32_t const count = static_cast(tokens.size()); + require(count > 0 && count <= 2 && firstSlot >= 0 && firstSlot + count <= 2, "Invalid selected slots"); + require(decode || count == 1, "Probe prefill selects one slot at a time"); + for (auto const& row : tokens) + { + require(!row.empty() && (!decode || row.size() == 1), "Invalid step token count"); + } + select(firstSlot, count); + DecodingInferenceContext context; + context.initialize(count, 8, std::nullopt, {}, "", mStream); + context.temperature = 0.0F; + context.topK = 1; + context.tokenIds = tokens; + for (int32_t i = 0; i < count; ++i) + { + require(mCommitted[firstSlot + i] + static_cast(tokens[i].size()) + <= mRuntime.mDeployment.base.maxKVCacheCapacity, + "KV endpoint exceeds engine capacity"); + context.effectivePrefillLengths[i] = static_cast(tokens[i].size()); + } + // Nonempty-cache S=1 is decoded by the attention plugin and needs absolute context lengths. + bool const promptTail = !decode && tokens[0].size() == 1 && mCommitted[firstSlot] > 0; + bool const executeDecode = decode || promptTail; + std::cout << "STEP slot=" << firstSlot << " batch=" << count << " profile=" << (executeDecode ? 1 : 0) + << " logical_decode=" << decode << " prompt_tail=" << promptTail << " start=" << mCommitted[firstSlot] + << " tokens=" << tokens[0].size() << std::endl; + bool const success = executeDecode ? mRuntime.mDecoderRegistry->cachePrimingStrategy().decodeStep(context) + : mRuntime.runBaseModelPrefill(context, nullptr, false); + require(success, "Native execution failed"); + int32_t const vocab = mRuntime.mDeployment.base.outputVocabSize; + CUDA_CHECK(cudaMemcpyAsync(mHostLogits.rawPointer(), mRuntime.mPipelineIO->outputLogits.rawPointer(), + static_cast(count) * vocab * sizeof(float), cudaMemcpyDeviceToHost, mStream)); + CUDA_CHECK(cudaMemcpyAsync(mHostLengths.rawPointer(), cache().getKVCacheLengths().rawPointer(), + count * sizeof(int32_t), cudaMemcpyDeviceToHost, mStream)); + CUDA_CHECK(cudaStreamSynchronize(mStream)); + std::vector> results; + for (int32_t i = 0; i < count; ++i) + { + mCommitted[firstSlot + i] += static_cast(tokens[i].size()); + require(mHostLengths.dataPointer()[i] == mCommitted[firstSlot + i], "Incorrect committed length"); + auto const* begin = mHostLogits.dataPointer() + i * vocab; + results.emplace_back(begin, begin + vocab); + } + return results; + } + + //! Snapshot all recurrent/conv bytes and the materialized attention-KV prefix. + void observeEndpoint(int32_t slot, int32_t length) + { + mCommitted.at(slot) = length; + } + + //! Capture sampler output scratch to detect accidental sampling by forward-only APIs. + std::vector samplingBytes() + { + auto& indices = mRuntime.mSamplingIndices; + std::vector result(indices.getMemoryCapacity()); + copyToHost(indices.rawPointer(), result.data(), result.size()); + return result; + } + + //! Snapshot all recurrent/conv bytes and the materialized attention-KV prefix. + void snapshot(int32_t slot) + { + mSnapshots.clear(); + mSnapshotSlot = slot; + mSnapshotLength = mCommitted[slot]; + auto capture = [&](void* pointer, size_t bytes) { + Snapshot entry{pointer, std::vector(bytes)}; + copyToHost(pointer, entry.bytes.data(), bytes); + mSnapshots.push_back(std::move(entry)); + }; + auto& mamba = cache().getMambaCacheManager(); + for (int32_t layer = 0; layer < mamba.numLayers(); ++layer) + { + for (Tensor* tensor : {&mamba.getRecurrentState(layer), &mamba.getConvState(layer)}) + { + size_t const bytes = tensor->getMemoryCapacity() / 2; + capture(static_cast(tensor->rawPointer()) + slot * bytes, bytes); + } + } + auto& kv = cache().getKVCacheManager(); + for (int32_t layer = 0; layer < kv.numLayers(); ++layer) + { + auto& tensor = kv.getCombinedKVCache(layer); + auto const shape = tensor.getShape(); + size_t const tokenBytes = shape[3] * shape[4] * sizeof(half); + require(tensor.getDataType() == nvinfer1::DataType::kHALF, "Probe requires FP16 KV"); + int64_t const pagesPerSlot = (mRuntime.mDeployment.base.maxKVCacheCapacity + shape[2] - 1) / shape[2]; + require(shape[1] == 2 * pagesPerSlot, "Probe requires a two-slot pool with no surplus pages"); + size_t const halfBytes = tensor.getMemoryCapacity() / 2; + size_t const slotBytes = halfBytes / 2; + require(slotBytes / tokenBytes >= static_cast(mCommitted[slot]), "KV snapshot range"); + for (int32_t plane = 0; plane < 2; ++plane) + { + capture(static_cast(tensor.rawPointer()) + plane * halfBytes + slot * slotBytes, + mCommitted[slot] * tokenBytes); + } + } + auto& table = mRuntime.mSharedResources->kvPageTables[0]->kernelView(); + size_t const tableBytes = table.getMemoryCapacity() / 2; + capture(static_cast(table.rawPointer()) + slot * tableBytes, tableBytes); + } + + //! Require byte-for-byte preservation of the inactive request's live state. + void verifySnapshot() + { + require(mSnapshotSlot >= 0 && mCommitted[mSnapshotSlot] == mSnapshotLength, "Inactive endpoint changed"); + size_t totalBytes = 0; + for (auto const& entry : mSnapshots) + { + for (size_t offset = 0; offset < entry.bytes.size(); offset += kCOPY_BYTES) + { + size_t const bytes = std::min(static_cast(kCOPY_BYTES), entry.bytes.size() - offset); + CUDA_CHECK(cudaMemcpyAsync(mScratch.rawPointer(), static_cast(entry.pointer) + offset, + bytes, cudaMemcpyDeviceToHost, mStream)); + CUDA_CHECK(cudaStreamSynchronize(mStream)); + require(std::memcmp(mScratch.rawPointer(), entry.bytes.data() + offset, bytes) == 0, + "Inactive state modified"); + totalBytes += bytes; + } + } + std::cout << "ISOLATION slot=" << mSnapshotSlot << " exact_bytes=" << totalBytes << " passed=1" << std::endl; + } + + //! Record device allocator availability, not process RSS or exclusive GPU use. + void memory(std::string const& label) + { + size_t freeBytes = 0; + size_t totalBytes = 0; + CUDA_CHECK(cudaMemGetInfo(&freeBytes, &totalBytes)); + std::cout << "MEMORY " << label << " cuda_free=" << freeBytes << " cuda_total=" << totalBytes << std::endl; + } + +private: + struct Binding + { + Tensor* original; + Coords shape; + }; + struct Snapshot + { + void* pointer; + std::vector bytes; + }; + static constexpr int64_t kCOPY_BYTES = 1024 * 1024; + + HybridCacheManager& cache() + { + return *mRuntime.mSharedResources->cacheManagers[0]; + } + + void addBinding(std::string const& name) + { + auto* tensor = mOriginalMap.get(name); + require(tensor != nullptr && tensor->getShape()[0] == 2, "Invalid physical binding: " + name); + mBindings.emplace(name, Binding{tensor, tensor->getShape()}); + } + + void select(int32_t firstSlot, int32_t count) + { + for (auto const& [name, binding] : mBindings) + { + auto shape = binding.shape; + size_t const rowBytes = binding.original->getMemoryCapacity() / 2; + shape[0] = count; + auto* pointer = static_cast(binding.original->rawPointer()) + firstSlot * rowBytes; + mViews[name] = Tensor(pointer, shape, DeviceType::kGPU, binding.original->getDataType()); + mRuntime.mBaseTensorMap.set(name, mViews.at(name)); + } + require(mHostLengths.reshape({count}), "Selected lengths shape"); + for (int32_t i = 0; i < count; ++i) + { + mHostLengths.dataPointer()[i] = mCommitted[firstSlot + i]; + } + // Cache-manager lengths are execution scratch here; canonical endpoints stay slot-owned. + Tensor lengths(mHostLengths.rawPointer(), {count}, DeviceType::kCPU, nvinfer1::DataType::kINT32); + cache().resetForNewSequences(lengths, mStream); + } + + void copyToHost(void* source, std::byte* destination, size_t total) + { + for (size_t offset = 0; offset < total; offset += kCOPY_BYTES) + { + size_t const bytes = std::min(static_cast(kCOPY_BYTES), total - offset); + CUDA_CHECK(cudaMemcpyAsync(mScratch.rawPointer(), static_cast(source) + offset, bytes, + cudaMemcpyDeviceToHost, mStream)); + CUDA_CHECK(cudaStreamSynchronize(mStream)); + std::memcpy(destination + offset, mScratch.rawPointer(), bytes); + } + } + + LLMRankRuntime& mRuntime; + cudaStream_t mStream; + TensorMap mOriginalMap; + Tensor mHostLengths; + Tensor mHostLogits; + Tensor mScratch; + std::array mCommitted{}; + std::map mBindings; + std::map mViews; + std::vector mSnapshots; + int32_t mSnapshotSlot{-1}; + int32_t mSnapshotLength{}; +}; + +//! Run bounded synthetic continuation and staggered slot-isolation fixtures. +bool runProbe(LLMRankRuntime& runtime, tokenizer::Tokenizer& tokenizer, cudaStream_t stream) +{ + ContinuousBatchingProbe probe(runtime, stream); + auto const vocabulary = tokenizer.encode("The red fox walks beside the blue river. Count one two three four. "); + require(!vocabulary.empty(), "Empty synthetic fixture"); + std::vector prompt; + for (int32_t i = 0; i < 129; ++i) + { + prompt.push_back(static_cast(vocabulary[i % vocabulary.size()])); + } + bool passed = true; + for (int32_t length : {1, 3, 4, 63, 64, 65, 127, 128, 129}) + { + std::vector const input(prompt.begin(), prompt.begin() + length); + probe.reset(0); + auto const reference = probe.step(0, {input})[0]; + int32_t const forcedToken = greedy(reference); + auto const referenceDecode = probe.step(0, {{forcedToken}}, true)[0]; + probe.reset(1); + std::vector candidate; + for (int32_t offset = 0; offset < length;) + { + int32_t const end = std::min(offset + 64, length); + candidate = probe.step(1, {{input.begin() + offset, input.begin() + end}})[0]; + offset = end; + } + passed = compare("chunk_prefill_" + std::to_string(length), reference, candidate, length <= 64) && passed; + auto const candidateDecode = probe.step(1, {{forcedToken}}, true)[0]; + passed = compare("chunk_decode_" + std::to_string(length), referenceDecode, candidateDecode, length <= 64) + && passed; + } + + probe.reset(0); + auto const firstA = probe.step(0, {prompt})[0]; + auto const tokenA = greedy(firstA); + auto const nextA = probe.step(0, {{tokenA}}, true)[0]; + probe.reset(1); + std::vector promptB(prompt.rbegin(), prompt.rend()); + auto const firstB = probe.step(1, {promptB})[0]; + auto const tokenB = greedy(firstB); + auto const nextB = probe.step(1, {{tokenB}}, true)[0]; + + probe.reset(0); + passed = compare("A_reuse", firstA, probe.step(0, {prompt})[0], true) && passed; + probe.snapshot(0); + probe.reset(1); + probe.step(1, {{promptB.begin(), promptB.begin() + 64}}); + probe.verifySnapshot(); + probe.snapshot(1); + passed = compare("A_decode_during_B_prefill", nextA, probe.step(0, {{tokenA}}, true)[0], true) && passed; + probe.verifySnapshot(); + probe.snapshot(0); + probe.step(1, {{promptB.begin() + 64, promptB.begin() + 128}}); + probe.verifySnapshot(); + auto const chunkB = probe.step(1, {{promptB.back()}})[0]; + probe.verifySnapshot(); + passed = compare("B_one_token_tail", firstB, chunkB) && passed; + probe.snapshot(0); + passed = compare("B_decode_after_A", nextB, probe.step(1, {{tokenB}}, true)[0]) && passed; + probe.verifySnapshot(); + + auto const secondTokenA = greedy(nextA); + auto const secondTokenB = greedy(nextB); + probe.reset(0); + probe.reset(1); + probe.step(0, {prompt}); + probe.step(1, {promptB}); + probe.step(0, {{tokenA}}, true); + probe.step(1, {{tokenB}}, true); + auto const secondA = probe.step(0, {{secondTokenA}}, true)[0]; + auto const secondB = probe.step(1, {{secondTokenB}}, true)[0]; + probe.reset(0); + probe.reset(1); + probe.step(0, {prompt}); + probe.step(1, {promptB}); + probe.step(0, {{tokenA}}, true); + probe.step(1, {{tokenB}}, true); + auto const batched = probe.step(0, {{secondTokenA}, {secondTokenB}}, true); + passed = compare("two_slot_decode_A", secondA, batched[0]) && passed; + passed = compare("two_slot_decode_B", secondB, batched[1]) && passed; + probe.memory("completed"); + return passed; +} +} // namespace rt +} // namespace trt_edgellm + +namespace trt_edgellm +{ +namespace rt +{ +//! Validate the production step interfaces against the independent P1/legacy path. +bool runStepTests(LLMRankRuntime& runtime, tokenizer::Tokenizer& tokenizer, cudaStream_t stream) +{ + auto const words = tokenizer.encode("A red fox crosses a blue river. One two three four. "); + require(!words.empty(), "Missing fixture tokens"); + std::vector promptA; + std::vector promptB; + for (int32_t i = 0; i < 129; ++i) + { + promptB.push_back(words[i % words.size()]); + if (i < 65) + { + promptA.push_back(words[i % words.size()]); + } + } + std::vector firstA, firstB, nextA, nextB, chunkFirstB, chunkNextB; + { + ContinuousBatchingProbe reference(runtime, stream); + firstA = reference.step(0, {promptA})[0]; + firstB = reference.step(1, {promptB})[0]; + nextA = reference.step(0, {{greedy(firstA)}}, true)[0]; + nextB = reference.step(1, {{greedy(firstB)}}, true)[0]; + reference.reset(1); + reference.step(1, {{promptB.begin(), promptB.begin() + 64}}); + reference.step(1, {{promptB.begin() + 64, promptB.begin() + 128}}); + chunkFirstB = reference.step(1, {{promptB.back()}})[0]; + chunkNextB = reference.step(1, {{greedy(firstB)}}, true)[0]; + bool const baselineQuality = compare("P1_decode_tail_quality_diagnostic", nextB, chunkNextB); + std::cout << "P1_TAIL_BASELINE quality_passed=" << baselineQuality << std::endl; + } + bool passed = true; + Tensor hostLogits({2, static_cast(firstA.size())}, DeviceType::kCPU, nvinfer1::DataType::kFLOAT); + { + ContinuousBatchingProbe observer(runtime, stream); + auto const samplerBefore = observer.samplingBytes(); + SequenceStepRuntime steps(runtime, stream); + auto rejects = [&](auto operation, char const* label) { + bool rejected = false; + try + { + operation(); + } + catch (std::logic_error const&) + { + rejected = true; + } + require(rejected && steps.healthy(), label); + std::cout << "REJECT " << label << " passed=1" << std::endl; + }; + rejects([&] { SequenceStepRuntime duplicate(runtime, stream); }, "second_session"); + LLMGenerationRequest request{}; + LLMGenerationResponse response; + require(!runtime.handleRequest(request, response, stream), "Legacy call entered a leased runtime"); + auto finish = [&](Tensor const& logits) { + steps.completeStep(); + size_t const elements = logits.getShape().volume(); + CUDA_CHECK(cudaMemcpyAsync(hostLogits.rawPointer(), logits.rawPointer(), elements * sizeof(float), + cudaMemcpyDeviceToHost, stream)); + CUDA_CHECK(cudaStreamSynchronize(stream)); + std::vector> values; + for (int64_t row = 0; row < logits.getShape()[0]; ++row) + { + float const* begin = hostLogits.dataPointer() + row * firstA.size(); + values.emplace_back(begin, begin + firstA.size()); + } + return values; + }; + SequenceOptions optionsA; + optionsA.maxOutputTokens = 3; + optionsA.temperature = 0.0F; + SequenceOptions optionsB; + optionsB.maxOutputTokens = 7; + optionsB.seed = 123; + auto a = steps.acquire(1, promptA, optionsA); + auto b = steps.acquire(2, promptB, optionsB); + auto const& logitsA = steps.beginPrefill(a, 65); + require(steps.state(a).committedTokens() == 0, "Forward published state before completion"); + rejects([&] { steps.release(a); }, "release_in_flight"); + rejects([&] { steps.beginPrefill(b, 64); }, "overlapping_forward"); + passed = compare("native_prefill_A", firstA, finish(logitsA)[0], true) && passed; + require(steps.state(a).output().empty(), "Prefill sampled unexpectedly"); + steps.acceptToken(a, greedy(firstA)); + require(steps.state(a).committedTokens() == 65, "Sample committed to KV early"); + observer.observeEndpoint(a.slot, 65); + observer.snapshot(a.slot); + finish(steps.beginPrefill(b, 64)); + observer.verifySnapshot(); + observer.observeEndpoint(b.slot, 64); + observer.snapshot(b.slot); + passed = compare("native_decode_A_while_B_prefills", nextA, finish(steps.beginDecode({a, {}}, 1))[0], true) + && passed; + observer.verifySnapshot(); + observer.observeEndpoint(a.slot, 66); + observer.snapshot(a.slot); + finish(steps.beginPrefill(b, 64)); + auto const tail = finish(steps.beginPrefill(b, 1))[0]; + passed = compare("native_sampling_free_tail_matches_P1", chunkFirstB, tail, true) && passed; + require(steps.state(b).output().empty() && steps.state(b).committedTokens() == 129, + "Prompt tail generated a token or committed incorrectly"); + observer.verifySnapshot(); + steps.acceptToken(b, greedy(firstB), 3); + passed = compare("native_decode_B_after_tail_matches_P1", chunkNextB, finish(steps.beginDecode({b, {}}, 1))[0], + true) + && passed; + require(steps.state(b).randomCounter() == 3 && steps.state(a).randomCounter() == 0, + "Random counters crossed requests"); + + observer.observeEndpoint(b.slot, 130); + observer.snapshot(b.slot); + steps.finish(a); + steps.release(a); + auto c = steps.acquire(3, promptA, optionsA); + require(c.slot == a.slot && c.generation != a.generation, "Slot was not safely reused"); + rejects([&] { steps.release(a); }, "stale_release"); + passed = compare("native_reused_slot_C", firstA, finish(steps.beginPrefill(c, 65))[0], true) && passed; + observer.verifySnapshot(); + require(steps.state(b).options().maxOutputTokens == 7 && steps.state(b).output().size() == 1, + "Partner request lost independent state"); + steps.release(c); + steps.release(b); + + a = steps.acquire(4, promptA, optionsA); + b = steps.acquire(5, promptB, optionsB); + finish(steps.beginPrefill(a, 65)); + finish(steps.beginPrefill(b, 129)); + steps.acceptToken(a, greedy(firstA)); + steps.acceptToken(b, greedy(firstB)); + rejects([&] { steps.beginDecode({a, a}, 2); }, "duplicate_decode_slot"); + rejects([&] { steps.beginDecode({b, a}, 2); }, "unordered_decode_slots"); + auto const pair = finish(steps.beginDecode({a, b}, 2)); + passed = compare("native_unequal_batch_A", nextA, pair[0], true) && passed; + passed = compare("native_unequal_batch_B", nextB, pair[1], true) && passed; + require(steps.state(a).committedTokens() == 66 && steps.state(b).committedTokens() == 130, + "Per-row endpoints did not commit independently"); + steps.release(a); + steps.release(b); + require(observer.samplingBytes() == samplerBefore, "Forward-only steps touched sampler output scratch"); + std::cout << "FORWARD_ONLY sampler_unchanged=1 passed=1" << std::endl; + } + + LLMGenerationRequest legacy{}; + legacy.requests.resize(1); + legacy.formattedRequests.resize(1); + legacy.requests[0].messages.push_back({"user", {{"text", "Synthetic pretokenized regression fixture"}}}); + legacy.preTokenizedInputIds = {promptA}; + legacy.temperature = 0.0F; + legacy.topK = 1; + legacy.topP = 1.0F; + legacy.maxGenerateLength = 2; + legacy.applyChatTemplate = false; + LLMGenerationResponse response; + require(runtime.handleRequest(legacy, response, stream), "Legacy API failed after step session destruction"); + require( + response.outputIds.size() == 1 && response.outputIds[0] == std::vector{greedy(firstA), greedy(nextA)}, + "Legacy output changed after releasing the step lease"); + std::cout << "LEGACY_RESTORED passed=1" << std::endl; + std::cout << "P2_STATE_GATE passed=" << passed << " full_chunk_quality=separate_P1_TAIL_BASELINE" << std::endl; + return passed; +} +} // namespace rt +} // namespace trt_edgellm + +int main(int argc, char** argv) +{ + if (argc != 3 && (argc != 4 || std::string(argv[3]) != "--steps")) + { + std::cerr << "Usage: continuous_batching_probe ENGINE_DIR CHECKPOINT_DIR [--steps]\n"; + return 2; + } + cudaStream_t stream{}; + try + { + auto plugin = trt_edgellm::loadEdgellmPluginLib(); + trt_edgellm::rt::require(plugin != nullptr, "Plugin loading failed"); + CUDA_CHECK(cudaStreamCreateWithFlags(&stream, cudaStreamNonBlocking)); + bool passed = false; + { + trt_edgellm::tokenizer::Tokenizer tokenizer; + trt_edgellm::rt::require(tokenizer.loadFromHF(argv[1], false), "Tokenizer loading failed"); + trt_edgellm::rt::LLMRankRuntime runtime(argv[1], "", {}, std::nullopt, stream, + trt_edgellm::rt::ParallelMapping{}, tokenizer, trt_edgellm::rt::ContextCacheConfig{}, argv[2], ""); + passed = argc == 4 ? trt_edgellm::rt::runStepTests(runtime, tokenizer, stream) + : trt_edgellm::rt::runProbe(runtime, tokenizer, stream); + } + CUDA_CHECK(cudaStreamDestroy(stream)); + std::cout << "PROBE_RESULT passed=" << passed << std::endl; + return passed ? 0 : 1; + } + catch (std::exception const& error) + { + std::cerr << "PROBE_ERROR " << error.what() << std::endl; + if (stream != nullptr) + { + cudaStreamDestroy(stream); + } + return 1; + } +} diff --git a/examples/llm/continuousBatchingProbe.md b/examples/llm/continuousBatchingProbe.md new file mode 100644 index 000000000..261b3e7af --- /dev/null +++ b/examples/llm/continuousBatchingProbe.md @@ -0,0 +1,52 @@ +# Continuous batching: P1 native feasibility probe + +This optional diagnostic does **not** enable a scheduler or alter HTTP admission. +It uses one vanilla, two-slot hybrid runtime with eager execution and the existing +serialized engine. The only runtime-header hook is a friend declaration; no +production method, data layout or thread-safety contract changes. + +Build with `-DBUILD_CONTINUOUS_BATCHING_PROBE=ON`, with `TRT_PACKAGE_DIR` set and +submodules initialized. Run `continuous_batching_probe ENGINE_DIR CHECKPOINT_DIR` +with the normal `LD_LIBRARY_PATH` and `EDGELLM_PLUGIN_PATH`. Do not run it alongside +a resident model on memory-constrained devices. Use an external 300-second total +deadline and preserve/restore service and watchdog state during maintenance. + +The operator script in the OpenClaw project can link this source against the +exact pinned `e8b29522938901f6df19ebeedd4b69bc8edbcd97` archive using a header +overlay and the reference CUDA device-link object. That saves a complete rebuild; +it must not be used against another revision or after layout/method changes. + +## Checks and numerical policy + +- Synthetic token fixtures have lengths 1, 3, 4, 63, 64, 65, 127, 128 and 129. + Slot 0 executes full prefill; slot 1 uses chunks of at most 64 tokens. A + teacher-forced decode step follows each path. +- Identical singleton execution shapes must produce exactly identical finite + logits. Chunked and batch-two comparisons require identical greedy IDs, + maximum absolute logit error at most 0.1 and relative L2 error at most 0.005. + These are declared-before-run preliminary FP16 screens, not evidence of task + quality or sufficient production acceptance. Every measured error is printed; + failures are retained and investigated, not fixed by loosening tolerances. +- Staggered execution checks A decoding while B prefills, profile switches, + B's one-token tail, both physical slots independently and a two-row decode. +- Inactive state comparison is byte-for-byte: all recurrent and convolution + state, materialized attention KV and the physical page-table row. Unwritten + KV capacity is not semantically live state and is not compared. Snapshots are + host-only diagnostic allocations, copied through 1 MiB pinned staging; they + are not proposed serving-time state copies. +- Execution lengths are packed into the existing cache-manager scratch tensor; + canonical endpoints stay in a slot-indexed host array. Both input and output + recurrent/convolution bindings alias the selected physical rows. KV pools stay + fixed while only page-table rows are selected. These test-only bindings do not + implement the production ownership/lifetime API planned for P2. +- A resumed one-token prompt chunk uses the decode profile and absolute + `context_lengths = committed + 1`. The attention plugin chooses vanilla decode + from nonempty start indices and sequence length one, even under the prefill + profile. Supplying chunk length one instead corrupts addressing. The probe + discards the decoder's sample and advances prompt accounting by one; P3 must + expose a sampling-free forward boundary with this same execution contract. + +Exit 0 means these fixtures passed. Any error, nonfinite output, state mutation, +greedy mismatch or threshold failure is nonzero. Full active-state numerical +comparisons, larger prompts, diverse natural-language fixtures, long-term memory +stability, graph execution and server concurrency still require later coverage. diff --git a/examples/llm/sequenceStepRuntime.md b/examples/llm/sequenceStepRuntime.md new file mode 100644 index 000000000..b212b596c --- /dev/null +++ b/examples/llm/sequenceStepRuntime.md @@ -0,0 +1,82 @@ +# P2: persistent sequence ownership and forward-only execution + +`runtime/state/sequenceSlots.{h,cpp}` holds request-local state without CUDA. +`runtime/sequenceStepRuntime.{h,cpp}` binds that state to one existing rank runtime. +This is an internal execution mechanism, not an HTTP scheduler. + +## Lifecycle + +Construct `SequenceStepRuntime` before legacy inference reshapes state, and keep +the parent `LLMRankRuntime` and the explicit CUDA stream alive until the session +is destroyed. One externally serialized owner uses the session. It leases the +legacy overlap gate for its entire lifetime; a second session or legacy +`handleRequest` cannot enter. Other parent APIs must not be invoked under the +lease. A successful teardown drains the stream and releases the gate. Any engine +or CUDA failure poisons the session and keeps the parent gate closed: recreate +the parent runtime before continuing. Full injected-fault coverage belongs to P5. + +1. `acquire(id, preparedTokens, options)` assigns a free physical slot and resets + only its recurrent/convolution rows. Tokenization/templating is the caller's + responsibility. Admission validates prompt plus requested output capacity per + sequence; it does not clamp all rows to a shared shortest headroom. +2. `beginPrefill(handle, count)` enqueues a valid prompt span. `beginDecode` takes + one slot or two distinct slots in physical order, each with exactly one + pending sampled output token. Both return borrowed GPU logits without sampling. +3. `completeStep()` waits on the forward-completion event and then publishes + committed endpoints. Another forward, release or acceptance is prohibited + while a step is pending. The event covers forward execution, not any additional + sampling/D2H operations a caller subsequently enqueues on the stream. +4. `acceptToken` records an externally selected token and its random-draw count. + A sampled token is not committed to KV until the next decode forward. Prompt + chunks never create output tokens. The configured output cap ends the logical + sequence; EOS, stop strings and other generation policy are left to P5. +5. `finish` ends logical work early; `release` invalidates the handle. Another + request may reuse that slot while its partner stays live. Handles include + pool ownership and allocation generation, preventing stale or foreign use. + +State references are borrowed owner-thread views, not objects to publish to +other threads. Logits remain valid only until the next forward; consume/sample +them before reusing the shared output buffer. + +## Physical storage and supported scope + +The three execution views are `{0}`, `{1}`, and `{0,1}`. They are prebuilt and +bind aliased recurrent/convolution rows, selected page-table and text-RoPE rows, +and the unchanged attention pools. Logical endpoints are slot-owned; one +preallocated eight-byte device array stages selected start indices. Existing +embedding and pipeline buffers are reused. There are no serving-time state +snapshots or cache compaction copies. The extra pinned metadata payload is +24 bytes, excluding allocator granularity and event/view bookkeeping. + +The current guard accepts only single-rank, two-slot vanilla text-only +`qwen3_5_text`, without context reuse, LoRA, speculative decode, context-dependent +or dual RoPE, PLE or deepstack. It intentionally does not advertise support for +other engines, graphs, multimodal inputs or tensor parallelism. + +A resumed physical sequence length of one selects the attention decode kernel; +the forward adapter therefore uses absolute context lengths and the decode +profile while retaining **logical prompt** accounting. This exactly preserves +the P1 execution recipe without invoking its sampler. A new fixture exposed +full-versus-chunked numerical divergence in that existing recipe. P2's exact +execution-preservation gate must not be confused with passing that separate +numerical/quality gate. An experimental padded two-position tail also missed +the unchanged threshold and was not adopted. + +## Verification + +`unittests/cpp/runtime/state/sequenceSlotsTest.cpp` is picked up by the existing +runtime-state unit-test target. It can also be compiled with the state source +and bundled GoogleTest without CUDA. Nine CPU tests passed normally and with +AddressSanitizer/UndefinedBehaviorSanitizer, including 100 reuse cycles. + +The optional `continuous_batching_probe ENGINE_DIR CHECKPOINT_DIR --steps` +tests this API against independent P1 and legacy execution. The OpenClaw +`run-continuous-batching-p2.sh` operator wrapper preserves the production runtime, +builds the added sources beside the pinned archive, bounds inference and restores +the model/watchdog. P1's earlier standalone source/binary remain preserved in +their separate evidence directory. + +The `P2_STATE_GATE` line reports mechanism correctness. The retained +`P1_TAIL_BASELINE quality_passed=0` is a **failing** numerical diagnostic, even +when P2 succeeds. P3 must resolve or rigorously qualify it before production +chunking can be considered complete. No threshold was relaxed. diff --git a/unittests/cpp/runtime/state/sequenceSlotsTest.cpp b/unittests/cpp/runtime/state/sequenceSlotsTest.cpp new file mode 100644 index 000000000..400988300 --- /dev/null +++ b/unittests/cpp/runtime/state/sequenceSlotsTest.cpp @@ -0,0 +1,159 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ +#include "runtime/state/sequenceSlots.h" +#include +#include +#include + +namespace trt_edgellm +{ +namespace rt +{ +TEST(SequenceSlots, RejectsInvalidCapacity) +{ + EXPECT_THROW(SequenceSlots(0, 8, 16), std::invalid_argument); + EXPECT_THROW(SequenceSlots(2, 16, 8), std::invalid_argument); + SequenceSlots slots(2, 8, 16); + EXPECT_THROW(slots.acquire(1, {}, {}), std::invalid_argument); + EXPECT_THROW(slots.acquire(1, {1}, {}), std::invalid_argument); +} + +TEST(SequenceSlots, AdmissionIsBoundedAndIdsAreUnique) +{ + SequenceSlots slots(2, 8, 16); + SequenceOptions options; + options.maxOutputTokens = 4; + auto a = slots.acquire(1, {1, 2}, options); + EXPECT_THROW(slots.acquire(1, {3}, options), std::invalid_argument); + auto b = slots.acquire(2, {3}, options); + EXPECT_EQ(a.slot, 0); + EXPECT_EQ(b.slot, 1); + EXPECT_THROW(slots.acquire(3, {3}, options), std::runtime_error); +} + +TEST(SequenceSlots, RejectsStaleAndForeignHandles) +{ + SequenceSlots slots(2, 8, 16); + SequenceSlots other(2, 8, 16); + SequenceOptions options; + options.maxOutputTokens = 4; + auto a = slots.acquire(1, {1}, options); + EXPECT_THROW(other.get(a), std::invalid_argument); + slots.release(a); + EXPECT_THROW(slots.get(a), std::invalid_argument); + auto b = slots.acquire(2, {2}, options); + EXPECT_EQ(a.slot, b.slot); + EXPECT_NE(a.generation, b.generation); + EXPECT_THROW(slots.finish(a), std::invalid_argument); + EXPECT_EQ(slots.get(b).requestId(), 2U); +} + +TEST(SequenceSlots, SeparatesSampledAndCommittedTokens) +{ + SequenceSlots slots(2, 8, 16); + SequenceOptions options; + options.maxOutputTokens = 2; + auto a = slots.acquire(1, {1, 2, 3}, options); + slots.commitPrompt(a, 2); + EXPECT_EQ(slots.get(a).promptCursor(), 2); + EXPECT_THROW(slots.acceptToken(a, 4), std::logic_error); + slots.commitPrompt(a, 1); + EXPECT_EQ(slots.get(a).phase(), SequencePhase::kAwaitingSample); + slots.acceptToken(a, 4, 7); + EXPECT_EQ(slots.get(a).committedTokens(), 3); + EXPECT_EQ(slots.get(a).randomCounter(), 7U); + slots.commitDecode(a); + EXPECT_EQ(slots.get(a).committedTokens(), 4); + EXPECT_THROW(slots.commitDecode(a), std::logic_error); + slots.acceptToken(a, 5); + EXPECT_EQ(slots.get(a).phase(), SequencePhase::kFinished); + EXPECT_EQ(slots.get(a).output().size(), 2U); + EXPECT_EQ(slots.get(a).committedTokens(), 4); + EXPECT_THROW(slots.acceptToken(a, 6), std::logic_error); +} + +TEST(SequenceSlots, OneTokenOutputNeedsNoDecodeForward) +{ + SequenceSlots slots(2, 8, 16); + SequenceOptions options; + options.maxOutputTokens = 1; + auto a = slots.acquire(1, {1, 2}, options); + slots.commitPrompt(a, 2); + slots.acceptToken(a, 3); + EXPECT_EQ(slots.get(a).phase(), SequencePhase::kFinished); + EXPECT_EQ(slots.get(a).committedTokens(), 2); + EXPECT_THROW(slots.commitDecode(a), std::logic_error); +} + +TEST(SequenceSlots, RequestsKeepIndependentOptionsAndHeadroom) +{ + SequenceSlots slots(2, 12, 16); + SequenceOptions first; + first.maxOutputTokens = 2; + first.temperature = 0.0F; + first.stopStrings = {"END"}; + first.logitBias[7] = -2.0F; + SequenceOptions second; + second.maxOutputTokens = 8; + second.seed = 123; + auto a = slots.acquire(1, std::vector(12, 1), first); + auto b = slots.acquire(2, {1}, second); + slots.commitPrompt(a, 12); + slots.acceptToken(a, 2); + slots.finish(a); + slots.release(a); + EXPECT_EQ(slots.get(b).options().maxOutputTokens, 8); + EXPECT_EQ(slots.get(b).options().seed, 123U); + EXPECT_TRUE(slots.get(b).options().stopStrings.empty()); + EXPECT_EQ(slots.get(b).committedTokens(), 0); +} + +TEST(SequenceSlots, InvalidTransitionsDoNotMutateState) +{ + SequenceSlots slots(2, 8, 16); + SequenceOptions options; + options.maxOutputTokens = 2; + auto a = slots.acquire(1, {1, 2}, options); + EXPECT_THROW(slots.commitPrompt(a, 3), std::logic_error); + EXPECT_THROW(slots.commitPrompt(a, 0), std::logic_error); + EXPECT_THROW(slots.commitDecode(a), std::logic_error); + EXPECT_EQ(slots.get(a).promptCursor(), 0); + slots.commitPrompt(a, 2); + EXPECT_THROW(slots.commitPrompt(a, 1), std::logic_error); + EXPECT_THROW(slots.acceptToken(a, -1), std::logic_error); + EXPECT_TRUE(slots.get(a).output().empty()); +} + +TEST(SequenceSlots, DetectsRandomCounterOverflow) +{ + SequenceSlots slots(2, 8, 16); + SequenceOptions options; + options.maxOutputTokens = 3; + auto a = slots.acquire(1, {1}, options); + slots.commitPrompt(a, 1); + slots.acceptToken(a, 2, std::numeric_limits::max()); + slots.commitDecode(a); + EXPECT_THROW(slots.acceptToken(a, 3, 1), std::logic_error); + EXPECT_EQ(slots.get(a).output().size(), 1U); +} + +TEST(SequenceSlots, ReuseDoesNotCarryPromptOutputOrOptions) +{ + SequenceSlots slots(1, 8, 16); + SequenceOptions options; + options.maxOutputTokens = 1; + for (uint64_t request = 1; request <= 100; ++request) + { + auto handle = slots.acquire(request, {static_cast(request)}, options); + EXPECT_EQ(slots.get(handle).committedTokens(), 0); + EXPECT_EQ(slots.get(handle).randomCounter(), 0U); + EXPECT_TRUE(slots.get(handle).output().empty()); + slots.commitPrompt(handle, 1); + slots.acceptToken(handle, 1); + slots.release(handle); + } +} +} // namespace rt +} // namespace trt_edgellm From f416525af89eb7270129e5234089f312becd4fa4 Mon Sep 17 00:00:00 2001 From: ajmalrasi Date: Tue, 15 Sep 2026 09:04:25 +0530 Subject: [PATCH 2/7] feat: qualify bounded chunked prefill for Qwen3.5 Signed-off-by: ajmalrasi --- cpp/runtime/sequenceStepRuntime.cpp | 11 + cpp/runtime/sequenceStepRuntime.h | 3 + cpp/runtime/state/prefillChunk.h | 30 ++ examples/llm/chunkedPrefill.md | 79 +++++ examples/llm/continuousBatchingProbe.cpp | 318 +++++++++++++++++- examples/llm/sequenceStepRuntime.md | 7 + .../cpp/runtime/state/sequenceSlotsTest.cpp | 56 +++ 7 files changed, 501 insertions(+), 3 deletions(-) create mode 100644 cpp/runtime/state/prefillChunk.h create mode 100644 examples/llm/chunkedPrefill.md diff --git a/cpp/runtime/sequenceStepRuntime.cpp b/cpp/runtime/sequenceStepRuntime.cpp index 4abac3b9f..77a0f5d98 100644 --- a/cpp/runtime/sequenceStepRuntime.cpp +++ b/cpp/runtime/sequenceStepRuntime.cpp @@ -7,6 +7,7 @@ #include "kernels/embeddingKernels/embeddingKernels.h" #include "kernels/posEncoding/initializeCosSinCache.h" #include "runtime/llmRankRuntime.h" +#include "runtime/state/prefillChunk.h" #include #include @@ -163,6 +164,16 @@ Tensor const& SequenceStepRuntime::beginPrefill(SequenceHandle handle, int32_t c return enqueue({handle, {}}, 1, count, false); } +Tensor const& SequenceStepRuntime::beginPrefillChunk(SequenceHandle handle) +{ + requireIdle(); + auto const& sequence = state(handle); + int32_t const remaining = static_cast(sequence.prompt().size()) - sequence.promptCursor(); + require(sequence.promptCursor() % 64 == 0 && (sequence.promptCursor() == 0 || remaining >= 64), + "Chunk policy cannot continue an incompatible manually partitioned prompt"); + return beginPrefill(handle, nextPrefillChunkSize(remaining)); +} + Tensor const& SequenceStepRuntime::beginDecode(std::array const& handles, int32_t count) { requireIdle(); diff --git a/cpp/runtime/sequenceStepRuntime.h b/cpp/runtime/sequenceStepRuntime.h index 30ac79686..ca51d7ed6 100644 --- a/cpp/runtime/sequenceStepRuntime.h +++ b/cpp/runtime/sequenceStepRuntime.h @@ -30,7 +30,10 @@ class SequenceStepRuntime SequenceState const& state(SequenceHandle handle) const; //! Enqueue a prompt span or pending output tokens. No sampling, output publication or CPU state commit occurs. //! Logits are borrowed until the next begin call; they become readable after completion on the supplied stream. + //! Low-level span control for diagnostics; arbitrary partitions are not numerically qualified. Tensor const& beginPrefill(SequenceHandle handle, int32_t count); + //! Enqueue at most 128 true prompt tokens; completeStep publishes progress. Sample only at kAwaitingSample. + Tensor const& beginPrefillChunk(SequenceHandle handle); Tensor const& beginDecode(std::array const& handles, int32_t count); //! Wait for forward completion before publishing endpoints and reusing pinned staging or physical slots. void completeStep(); diff --git a/cpp/runtime/state/prefillChunk.h b/cpp/runtime/state/prefillChunk.h new file mode 100644 index 000000000..079446d4f --- /dev/null +++ b/cpp/runtime/state/prefillChunk.h @@ -0,0 +1,30 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ +#pragma once +#include +#include + +namespace trt_edgellm +{ +namespace rt +{ +//! Fixed prompt-only partition. Input is the number of unconsumed, already-tokenized prompt tokens. +//! Keep intermediate chunks on 64-token boundaries and keep resumed final chunks between 64 and 128 tokens. +inline int32_t nextPrefillChunkSize(int32_t remaining) +{ + if (remaining <= 0) + { + throw std::logic_error("No remaining prompt tokens"); + } + constexpr int32_t kCHUNK_CAP = 128; + constexpr int32_t kALIGNMENT = 64; + if (remaining <= kCHUNK_CAP) + { + return remaining; + } + return remaining < kCHUNK_CAP + kALIGNMENT ? kALIGNMENT : kCHUNK_CAP; +} +} // namespace rt +} // namespace trt_edgellm diff --git a/examples/llm/chunkedPrefill.md b/examples/llm/chunkedPrefill.md new file mode 100644 index 000000000..6c0c4c38e --- /dev/null +++ b/examples/llm/chunkedPrefill.md @@ -0,0 +1,79 @@ +# P3: bounded prompt continuation + +`SequenceStepRuntime::beginPrefillChunk(handle)` consumes the next bounded part +of the immutable token vector supplied at admission. The caller formats and +tokenizes the complete prompt once. It must call `completeStep()` before reading +progress, sampling, releasing a slot or starting another forward. + +The fixed policy in `runtime/state/prefillChunk.h` uses at most 128 true prompt +tokens per call. Intermediate endpoints remain multiples of 64. When 129–191 +tokens remain, it consumes 64 so the final chunk contains 65–127 tokens. Otherwise +it consumes up to 128. Thus a resumed final chunk has 64–128 tokens; a complete +short prompt, including a cold one-token prompt, can be shorter. + +Examples: + +| Prompt tokens | Chunk spans | +| --- | --- | +| 1 | 1 | +| 65 | 65 | +| 129 | 64, 65 | +| 130 | 64, 66 | +| 191 | 64, 127 | +| 192 | 128, 64 | +| 257 | 128, 64, 65 | +| 6144 | 48 chunks of 128 | + +No padding tokens, prompt replay, state-sized copies, new GPU allocation or +sampling occur in this helper. Attention KV, recurrent and convolution storage +continue through the existing selected-row forward mechanism. Each non-final +completion remains in `kPrefill`; only final completion reaches +`kAwaitingSample`. `acceptToken` rejects partial prefill. The first accepted +output is pending input for a later decode; accepting it does not increase the +committed cache length. Sampling/output policy itself remains P5 work. + +The fixed policy must be used from the start of prompt processing. It rejects +incompatible manual history (unaligned cursor or a resumed remainder below 64) +before enqueuing GPU work. `beginPrefill(handle, count)` is retained as a low-level +mechanism/diagnostic interface; arbitrary partitions are not numerically +qualified. HTTP integration, automatic scheduling and performance tuning remain +later phases. + +## Numerical scope and diagnostics + +The original 64+64+1 recipe remains a failing numerical control: its teacher +continuation exceeds the existing maximum absolute 0.1 / relative L2 0.005 logit +screens. The first candidate avoided only singleton tails; expanded tests also +found drift with 2–4-token resumed tails. No threshold was relaxed. The selected +policy avoids these execution shapes while preserving the exact prompt. + +On SM87, GDN selects different implementations for physical sequence length one +and greater than one. The prefill implementation is a sequential recurrence; +64-token alignment here is an empirical shape restriction, not a claim that GDN +requires blocks of 64. Kernel/tactic-level attribution of small-shape drift is +not established. Do not extrapolate this policy's qualification to other engines, +precisions, models, devices or chunk caps. + +The optional probe has two P3 modes: + +- `--chunks`: original boundary controls and longer synthetic prompts through + 6144 tokens, four teacher-forced steps, and all active state before/after decode. +- `--chunks-extra`: every length 129–193, additional multi-chunk tails, four + once-formatted chat prompts (including code, Unicode/JSON and numbered records), + eight teacher steps on chats, the retained raw-tail negative control, and + interleaving in both physical-slot orders. + +Logit qualification requires finite values, equal greedy IDs, max absolute error +at most 0.1 and relative L2 at most 0.005. Same-shape short controls and interleaved +references require exact equality. Active-state diagnostics apply finite, +0.1/0.005 screens to each recurrent/conv tensor and each materialized K/V plane; +inactive snapshots require every byte unchanged, including page-table rows. +Unused KV capacity is excluded. Host snapshot/staging copies belong only to the +probe, never to serving execution. + +`RAW_TAIL_DIAGNOSTIC quality_passed=0` is expected retained evidence for the +unsupported manual partition; it is separate from `P3_QUALITY_GATE`. Consult the +OpenClaw P3 evidence report for exact tested source revisions, results, failed +attempts, operational restoration, and remaining limitations. A passing short +numerical screen does not establish broad task quality, throughput or long-term +memory stability. diff --git a/examples/llm/continuousBatchingProbe.cpp b/examples/llm/continuousBatchingProbe.cpp index 29bb5322e..e2d0ce3b8 100644 --- a/examples/llm/continuousBatchingProbe.cpp +++ b/examples/llm/continuousBatchingProbe.cpp @@ -247,6 +247,73 @@ class ContinuousBatchingProbe std::cout << "ISOLATION slot=" << mSnapshotSlot << " exact_bytes=" << totalBytes << " passed=1" << std::endl; } + int32_t vocabularySize() const + { + return mRuntime.mDeployment.base.outputVocabSize; + } + + //! Compare live physical rows without allocating another state-sized snapshot. + bool compareActive(std::string const& label, int32_t length) + { + Tensor candidate({kCOPY_BYTES}, DeviceType::kCPU, nvinfer1::DataType::kUINT8); + bool passed = true; + auto check = [&](Tensor& tensor, size_t rowBytes, size_t planeOffset, size_t bytes, std::string const& name) { + double squaredError = 0.0, squaredReference = 0.0; + double maxError = 0.0; + bool finite = true; + size_t const elementBytes + = tensor.getDataType() == nvinfer1::DataType::kFLOAT ? sizeof(float) : sizeof(half); + for (size_t offset = 0; offset < bytes; offset += kCOPY_BYTES) + { + size_t const count = std::min(static_cast(kCOPY_BYTES), bytes - offset); + auto* source = static_cast(tensor.rawPointer()) + planeOffset + offset; + CUDA_CHECK(cudaMemcpyAsync(mScratch.rawPointer(), source, count, cudaMemcpyDeviceToHost, mStream)); + CUDA_CHECK( + cudaMemcpyAsync(candidate.rawPointer(), source + rowBytes, count, cudaMemcpyDeviceToHost, mStream)); + CUDA_CHECK(cudaStreamSynchronize(mStream)); + for (size_t i = 0; i < count / elementBytes; ++i) + { + double const a = elementBytes == sizeof(float) ? mScratch.dataPointer()[i] + : __half2float(mScratch.dataPointer()[i]); + double const b = elementBytes == sizeof(float) ? candidate.dataPointer()[i] + : __half2float(candidate.dataPointer()[i]); + finite = finite && std::isfinite(a) && std::isfinite(b); + maxError = std::max(maxError, std::abs(a - b)); + squaredError += (a - b) * (a - b); + squaredReference += a * a; + } + } + double const relative = std::sqrt(squaredError / std::max(squaredReference, 1.0e-30)); + bool const ok = finite && maxError <= 0.1 && relative <= 0.005; + passed = ok && passed; + std::cout << "ACTIVE_STATE " << label << " " << name << " max_abs=" << maxError + << " relative_l2=" << relative << " finite=" << finite << " passed=" << ok << std::endl; + }; + auto& mamba = cache().getMambaCacheManager(); + for (int32_t layer = 0; layer < mamba.numLayers(); ++layer) + { + auto& recurrent = mamba.getRecurrentState(layer); + auto& conv = mamba.getConvState(layer); + check(recurrent, recurrent.getMemoryCapacity() / 2, 0, recurrent.getMemoryCapacity() / 2, + "recurrent_" + std::to_string(layer)); + check(conv, conv.getMemoryCapacity() / 2, 0, conv.getMemoryCapacity() / 2, "conv_" + std::to_string(layer)); + } + auto& kv = cache().getKVCacheManager(); + for (int32_t layer = 0; layer < kv.numLayers(); ++layer) + { + auto& tensor = kv.getCombinedKVCache(layer); + auto const shape = tensor.getShape(); + size_t const tokenBytes = shape[3] * shape[4] * sizeof(half); + size_t const planeBytes = tensor.getMemoryCapacity() / 2; + for (int32_t plane = 0; plane < 2; ++plane) + { + check(tensor, planeBytes / 2, plane * planeBytes, length * tokenBytes, + "kv_" + std::to_string(layer) + "_" + std::to_string(plane)); + } + } + return passed; + } + //! Record device allocator availability, not process RSS or exclusive GPU use. void memory(std::string const& label) { @@ -578,14 +645,256 @@ bool runStepTests(LLMRankRuntime& runtime, tokenizer::Tokenizer& tokenizer, cuda std::cout << "P2_STATE_GATE passed=" << passed << " full_chunk_quality=separate_P1_TAIL_BASELINE" << std::endl; return passed; } + +//! Full-prompt numerical screen, separate from the matching-recipe P2 mechanism gate. +bool runChunkTests(LLMRankRuntime& runtime, tokenizer::Tokenizer& tokenizer, cudaStream_t stream, bool extended) +{ + auto const words = tokenizer.encode("A red fox crosses a blue river. One two three four. "); + require(!words.empty(), "Missing fixture tokens"); + ContinuousBatchingProbe observer(runtime, stream); + SequenceStepRuntime steps(runtime, stream); + Tensor host({2, observer.vocabularySize()}, DeviceType::kCPU, nvinfer1::DataType::kFLOAT); + auto finish = [&](Tensor const& logits) { + steps.completeStep(); + size_t const count = logits.getShape().volume(); + require(static_cast(host.getMemoryCapacity()) >= count * sizeof(float), "Host logit capacity"); + CUDA_CHECK(cudaMemcpyAsync( + host.rawPointer(), logits.rawPointer(), count * sizeof(float), cudaMemcpyDeviceToHost, stream)); + CUDA_CHECK(cudaStreamSynchronize(stream)); + return std::vector(host.dataPointer(), host.dataPointer() + logits.getShape()[1]); + }; + SequenceOptions options; + options.maxOutputTokens = 12; + bool passed = true; + uint64_t id = 0; + std::vector>> fixtures; + auto lengths = extended + ? std::vector{130, 131, 132, 191, 192, 193, 258, 259, 260, 319, 320, 321} + : std::vector{1, 3, 4, 63, 64, 65, 127, 128, 129, 255, 256, 257, 513, 1025, 2049, 6144}; + if (extended) + { + for (int32_t length = 129; length <= 193; ++length) + { + if (std::find(lengths.begin(), lengths.end(), length) == lengths.end()) + { + lengths.push_back(length); + } + } + } + for (int32_t length : lengths) + { + std::vector prompt; + for (int32_t i = 0; i < length; ++i) + { + prompt.push_back(words[i % words.size()]); + } + fixtures.emplace_back("policy_" + std::to_string(length), std::move(prompt)); + } + if (extended) + { + std::string records; + for (int32_t i = 0; i < 100; ++i) + { + records += "Record " + std::to_string(i) + ": warehouse " + std::to_string(i % 7) + " has " + + std::to_string(17 * i + 3) + " blue items and " + std::to_string(11 * i + 5) + " red items.\n"; + } + std::vector texts{ + "Explain why the Moon changes shape during a month. Use plain language and distinguish phases from " + "eclipses.", + "Write a Python function that merges two sorted lists. Explain empty inputs and duplicate values.\n" + "def merge(a, b):\n # Preserve ordering and duplicates.\n pass\n", + "Return JSON with fields city and greeting for Chennai, 東京, and Zürich. Preserve Unicode text. " + "The greeting in Tamil is வணக்கம். Do not invent population values.", + records + "Which warehouse appears in record 73, and how many blue items does that record contain?"}; + for (size_t i = 0; i < texts.size(); ++i) + { + if (i < 3) + { + std::string const topic = texts[i]; + for (int32_t repeat = 0; repeat < 4; ++repeat) + { + texts[i] += "\nAdditional requested detail " + std::to_string(repeat) + ": " + topic; + } + } + LLMGenerationRequest::Request request; + request.messages.push_back({"system", {{"text", "You are a helpful assistant. Answer directly."}}}); + request.messages.push_back({"user", {{"text", texts[i]}}}); + LLMGenerationRequest::FormattedRequest formatted; + require(tokenizer.applyChatTemplate(request, formatted, true, true, false), "Chat formatting failed"); + auto tokens = tokenizer.encode(formatted.formattedCompleteRequest); + require(!tokens.empty() && tokens.size() <= 6144, "Formatted fixture length"); + fixtures.emplace_back("chat_" + std::to_string(i) + "_" + std::to_string(tokens.size()), std::move(tokens)); + } + } + for (auto const& [label, prompt] : fixtures) + { + int32_t const length = static_cast(prompt.size()); + auto a = steps.acquire(++id, prompt, options); + auto b = steps.acquire(++id, prompt, options); + auto reference = finish(steps.beginPrefill(a, length)); + std::vector actual; + int32_t chunks = 0; + while (steps.state(b).phase() == SequencePhase::kPrefill) + { + int32_t const before = steps.state(b).promptCursor(); + actual = finish(steps.beginPrefillChunk(b)); + int32_t const span = steps.state(b).promptCursor() - before; + require(span > 0 && span <= 128 && (before == 0 || span >= 64), "Invalid bounded partition"); + require(steps.state(b).output().empty(), "Prefill emitted a completion"); + ++chunks; + } + std::cout << "PARTITION " << label << " chunks=" << chunks << std::endl; + passed = compare(label, reference, actual, length <= 128) && passed; + passed = observer.compareActive(label, length) && passed; + int32_t const continuation = label.find("chat_") == 0 ? 8 : 4; + for (int32_t token = 0; token < continuation; ++token) + { + int32_t const teacher = greedy(reference); + steps.acceptToken(a, teacher); + steps.acceptToken(b, teacher); + reference = finish(steps.beginDecode({a, {}}, 1)); + actual = finish(steps.beginDecode({b, {}}, 1)); + passed = compare(label + "_teacher_" + std::to_string(token), reference, actual, length <= 128) && passed; + } + passed = observer.compareActive(label + "_after_teacher", length + continuation) && passed; + steps.release(a); + steps.release(b); + } + if (extended) + { + std::vector prompt(129); + for (size_t i = 0; i < prompt.size(); ++i) + { + prompt[i] = words[i % words.size()]; + } + auto a = steps.acquire(++id, prompt, options); + auto b = steps.acquire(++id, prompt, options); + auto reference = finish(steps.beginPrefill(a, 129)); + finish(steps.beginPrefill(b, 64)); + finish(steps.beginPrefill(b, 64)); + auto raw = finish(steps.beginPrefill(b, 1)); + bool rawQuality = compare("raw_singleton_prefill", reference, raw); + rawQuality = observer.compareActive("raw_singleton_prefill", 129) && rawQuality; + steps.acceptToken(a, greedy(reference)); + steps.acceptToken(b, greedy(reference)); + reference = finish(steps.beginDecode({a, {}}, 1)); + raw = finish(steps.beginDecode({b, {}}, 1)); + rawQuality = compare("raw_singleton_teacher", reference, raw) && rawQuality; + rawQuality = observer.compareActive("raw_singleton_teacher", 130) && rawQuality; + std::cout << "RAW_TAIL_DIAGNOSTIC quality_passed=" << rawQuality << std::endl; + steps.release(a); + steps.release(b); + + std::vector shortPrompt(prompt.begin(), prompt.begin() + 65); + std::vector longPrompt(513); + for (size_t i = 0; i < longPrompt.size(); ++i) + { + longPrompt[i] = words[(i + 3) % words.size()]; + } + auto baseline = steps.acquire(++id, shortPrompt, options); + std::vector> referenceA{finish(steps.beginPrefill(baseline, 65))}; + for (int32_t i = 0; i < 5; ++i) + { + steps.acceptToken(baseline, greedy(referenceA.back())); + referenceA.push_back(finish(steps.beginDecode({baseline, {}}, 1))); + } + steps.release(baseline); + baseline = steps.acquire(++id, longPrompt, options); + auto const referenceB = finish(steps.beginPrefill(baseline, 513)); + steps.acceptToken(baseline, greedy(referenceB)); + auto const nextB = finish(steps.beginDecode({baseline, {}}, 1)); + steps.release(baseline); + + auto incompatible = steps.acquire(++id, shortPrompt, options); + finish(steps.beginPrefill(incompatible, 64)); + bool rejectedPolicy = false; + try + { + steps.beginPrefillChunk(incompatible); + } + catch (std::logic_error const&) + { + rejectedPolicy = true; + } + require(rejectedPolicy && steps.healthy(), "Policy accepted manually created singleton tail"); + steps.release(incompatible); + std::cout << "REJECT mixed_manual_partition passed=1" << std::endl; + for (bool reversed : {false, true}) + { + if (reversed) + { + b = steps.acquire(++id, longPrompt, options); + a = steps.acquire(++id, shortPrompt, options); + } + else + { + a = steps.acquire(++id, shortPrompt, options); + b = steps.acquire(++id, longPrompt, options); + } + std::cout << "INTERLEAVE decoding_slot=" << a.slot << " prefill_slot=" << b.slot << std::endl; + auto const scratch = observer.samplingBytes(); + passed + = compare("interleaved_A_prefill", referenceA[0], finish(steps.beginPrefillChunk(a)), true) && passed; + steps.acceptToken(a, greedy(referenceA[0])); + int32_t index = 0; + while (steps.state(b).phase() == SequencePhase::kPrefill) + { + observer.observeEndpoint(a.slot, steps.state(a).committedTokens()); + observer.snapshot(a.slot); + auto const chunk = finish(steps.beginPrefillChunk(b)); + observer.verifySnapshot(); + observer.observeEndpoint(b.slot, steps.state(b).committedTokens()); + observer.snapshot(b.slot); + passed = compare("interleaved_A_decode_" + std::to_string(index), referenceA[index + 1], + finish(steps.beginDecode({a, {}}, 1)), true) + && passed; + observer.verifySnapshot(); + ++index; + if (steps.state(b).phase() == SequencePhase::kPrefill) + { + steps.acceptToken(a, greedy(referenceA[index])); + bool rejected = false; + try + { + steps.acceptToken(b, 0); + } + catch (std::logic_error const&) + { + rejected = true; + } + require(rejected, "Sample accepted during partial prefill"); + } + else + { + passed = compare("interleaved_B_final", referenceB, chunk, true) && passed; + } + } + require(index == 5 && steps.state(b).output().empty(), "Interleaved prompt accounting"); + steps.acceptToken(b, greedy(referenceB)); + passed = compare("interleaved_B_teacher", nextB, finish(steps.beginDecode({b, {}}, 1)), true) && passed; + require(observer.samplingBytes() == scratch, "Chunk forward touched sampler scratch"); + require(steps.state(b).committedTokens() == 514 && steps.state(b).output().size() == 1, + "Final sample/cache accounting"); + steps.release(a); + steps.release(b); + } + std::cout << "P3_INTERLEAVING_GATE passed=" << passed << std::endl; + } + observer.memory("p3_complete"); + std::cout << "P3_QUALITY_GATE passed=" << passed << std::endl; + return passed; +} } // namespace rt } // namespace trt_edgellm int main(int argc, char** argv) { - if (argc != 3 && (argc != 4 || std::string(argv[3]) != "--steps")) + if (argc != 3 + && (argc != 4 + || (std::string(argv[3]) != "--steps" + && (std::string(argv[3]) != "--chunks" && std::string(argv[3]) != "--chunks-extra")))) { - std::cerr << "Usage: continuous_batching_probe ENGINE_DIR CHECKPOINT_DIR [--steps]\n"; + std::cerr << "Usage: continuous_batching_probe ENGINE_DIR CHECKPOINT_DIR [--steps|--chunks|--chunks-extra]\n"; return 2; } cudaStream_t stream{}; @@ -600,7 +909,10 @@ int main(int argc, char** argv) trt_edgellm::rt::require(tokenizer.loadFromHF(argv[1], false), "Tokenizer loading failed"); trt_edgellm::rt::LLMRankRuntime runtime(argv[1], "", {}, std::nullopt, stream, trt_edgellm::rt::ParallelMapping{}, tokenizer, trt_edgellm::rt::ContextCacheConfig{}, argv[2], ""); - passed = argc == 4 ? trt_edgellm::rt::runStepTests(runtime, tokenizer, stream) + passed = argc == 4 ? (std::string(argv[3]).find("--chunks") == 0 + ? trt_edgellm::rt::runChunkTests( + runtime, tokenizer, stream, std::string(argv[3]) == "--chunks-extra") + : trt_edgellm::rt::runStepTests(runtime, tokenizer, stream)) : trt_edgellm::rt::runProbe(runtime, tokenizer, stream); } CUDA_CHECK(cudaStreamDestroy(stream)); diff --git a/examples/llm/sequenceStepRuntime.md b/examples/llm/sequenceStepRuntime.md index b212b596c..9c25b7ec4 100644 --- a/examples/llm/sequenceStepRuntime.md +++ b/examples/llm/sequenceStepRuntime.md @@ -80,3 +80,10 @@ The `P2_STATE_GATE` line reports mechanism correctness. The retained `P1_TAIL_BASELINE quality_passed=0` is a **failing** numerical diagnostic, even when P2 succeeds. P3 must resolve or rigorously qualify it before production chunking can be considered complete. No threshold was relaxed. + +## Bounded prompt continuation + +P3 adds `beginPrefillChunk(handle)` for a fixed 128-token cap with numerically +qualified tail partitioning. Use it from the first prompt step; low-level manual +partitions remain diagnostic. See [chunked prefill](chunkedPrefill.md) for the +shape restriction, final-sample accounting and retained failing controls. diff --git a/unittests/cpp/runtime/state/sequenceSlotsTest.cpp b/unittests/cpp/runtime/state/sequenceSlotsTest.cpp index 400988300..b00184d51 100644 --- a/unittests/cpp/runtime/state/sequenceSlotsTest.cpp +++ b/unittests/cpp/runtime/state/sequenceSlotsTest.cpp @@ -3,6 +3,7 @@ * SPDX-License-Identifier: Apache-2.0 */ #include "runtime/state/sequenceSlots.h" +#include "runtime/state/prefillChunk.h" #include #include #include @@ -11,6 +12,61 @@ namespace trt_edgellm { namespace rt { +TEST(PrefillChunks, EverySupportedPromptPreservesTokensAndBounds) +{ + for (int32_t length = 1; length <= 6144; ++length) + { + SequenceSlots slots(2, 6144, 8192); + SequenceOptions options; + options.maxOutputTokens = 1; + std::vector prompt(length); + for (int32_t i = 0; i < length; ++i) + { + prompt[i] = i; + } + auto handle = slots.acquire(1, prompt, options); + int32_t consumed = 0; + while (consumed < length) + { + int32_t const count = nextPrefillChunkSize(length - consumed); + ASSERT_GT(count, 0); + ASSERT_LE(count, 128); + ASSERT_LE(count, length - consumed); + if (consumed > 0) + { + ASSERT_GE(count, 64); + ASSERT_EQ(consumed % 64, 0); + } + EXPECT_THROW(slots.acceptToken(handle, 0), std::logic_error); + for (int32_t i = 0; i < count; ++i) + { + ASSERT_EQ(slots.get(handle).prompt()[consumed + i], consumed + i); + } + slots.commitPrompt(handle, count); + consumed += count; + ASSERT_EQ(slots.get(handle).promptCursor(), consumed); + ASSERT_EQ(slots.get(handle).committedTokens(), consumed); + ASSERT_TRUE(slots.get(handle).output().empty()); + } + ASSERT_EQ(slots.get(handle).phase(), SequencePhase::kAwaitingSample); + slots.acceptToken(handle, 7); + ASSERT_EQ(slots.get(handle).output().size(), 1U); + ASSERT_EQ(slots.get(handle).committedTokens(), length); + EXPECT_THROW(slots.acceptToken(handle, 7), std::logic_error); + } +} + +TEST(PrefillChunks, InvalidRemaindersAndShortTails) +{ + EXPECT_THROW(nextPrefillChunkSize(0), std::logic_error); + EXPECT_THROW(nextPrefillChunkSize(-1), std::logic_error); + EXPECT_EQ(nextPrefillChunkSize(1), 1); + EXPECT_EQ(nextPrefillChunkSize(129), 64); + EXPECT_EQ(nextPrefillChunkSize(191), 64); + EXPECT_EQ(nextPrefillChunkSize(192), 128); + EXPECT_EQ(nextPrefillChunkSize(257), 128); +} + TEST(SequenceSlots, RejectsInvalidCapacity) { EXPECT_THROW(SequenceSlots(0, 8, 16), std::invalid_argument); From e37d897ad3c1aa8fc325a5b4e90c491689e0f031 Mon Sep 17 00:00:00 2001 From: ajmalrasi Date: Tue, 15 Sep 2026 09:28:19 +0530 Subject: [PATCH 3/7] feat: Add native continuous scheduler with bounded slot reuse Signed-off-by: ajmalrasi --- cpp/CMakeLists.txt | 3 +- cpp/runtime/continuousScheduler.cpp | 258 ++++++++++++++++ cpp/runtime/continuousScheduler.h | 146 +++++++++ cpp/runtime/greedySchedulerBackend.cpp | 94 ++++++ cpp/runtime/greedySchedulerBackend.h | 42 +++ cpp/runtime/sequenceStepRuntime.cpp | 5 + cpp/runtime/sequenceStepRuntime.h | 2 + examples/llm/continuousBatchingProbe.cpp | 110 ++++++- examples/llm/continuousScheduler.md | 82 +++++ .../cpp/runtime/continuousSchedulerTest.cpp | 292 ++++++++++++++++++ 10 files changed, 1026 insertions(+), 8 deletions(-) create mode 100644 cpp/runtime/continuousScheduler.cpp create mode 100644 cpp/runtime/continuousScheduler.h create mode 100644 cpp/runtime/greedySchedulerBackend.cpp create mode 100644 cpp/runtime/greedySchedulerBackend.h create mode 100644 examples/llm/continuousScheduler.md create mode 100644 unittests/cpp/runtime/continuousSchedulerTest.cpp diff --git a/cpp/CMakeLists.txt b/cpp/CMakeLists.txt index 1b4840de9..30a30f21b 100644 --- a/cpp/CMakeLists.txt +++ b/cpp/CMakeLists.txt @@ -220,7 +220,8 @@ if(ENABLE_CUTEDSL_MODULE_TEST_HOOK) target_compile_definitions(edgellmCore PRIVATE ENABLE_CUTEDSL_MODULE_TEST_HOOK) endif() -target_link_libraries(edgellmCore PRIVATE ${CMAKE_DL_LIBS}) +find_package(Threads REQUIRED) +target_link_libraries(edgellmCore PRIVATE ${CMAKE_DL_LIBS} Threads::Threads) # Apply FMHA SM exclusion definitions target_compile_definitions(edgellmCore PRIVATE ${FMHA_EXCLUDE_DEFINITIONS}) edgellm_apply_qnx_warning_suppressions(edgellmCore) diff --git a/cpp/runtime/continuousScheduler.cpp b/cpp/runtime/continuousScheduler.cpp new file mode 100644 index 000000000..042d2f4e7 --- /dev/null +++ b/cpp/runtime/continuousScheduler.cpp @@ -0,0 +1,258 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ +#include "runtime/continuousScheduler.h" +#include +#include +#include + +namespace trt_edgellm +{ +namespace rt +{ +ContinuousScheduler::ContinuousScheduler( + std::unique_ptr backend, size_t maxQueued, size_t maxQueuedBytes, Observer observer) + : mBackend(std::move(backend)) + , mMaxQueued(maxQueued) + , mMaxQueuedBytes(maxQueuedBytes) + , mObserver(std::move(observer)) +{ + if (!mBackend || !maxQueued || !maxQueuedBytes) + { + throw std::invalid_argument("Scheduler requires a backend and positive queue bounds"); + } + mWorker = std::thread(&ContinuousScheduler::run, this); +} +ContinuousScheduler::~ContinuousScheduler() +{ + close(); +} + +SchedulerTicket ContinuousScheduler::submit(std::vector const& prompt, int32_t maxOutput) +{ + if (prompt.empty() || prompt.size() > 6144 || maxOutput <= 0 + || maxOutput > 8192 - static_cast(prompt.size())) + { + throw std::invalid_argument("P4 request exceeds qualified token bounds"); + } + for (auto token : prompt) + { + if (!mBackend->validToken(token)) + { + throw std::invalid_argument("Token outside supported vocabulary"); + } + } + std::lock_guard lock(mMutex); + if (mClosing || !mHealthy.load()) + { + throw std::runtime_error("Scheduler closed or failed"); + } + size_t const bytes = prompt.size() * sizeof(int32_t); + if (mQueue.size() >= mMaxQueued || bytes > mMaxQueuedBytes - mQueuedBytes) + { + throw std::runtime_error("Scheduler queue full"); + } + if (mNextId == std::numeric_limits::max()) + { + throw std::overflow_error("Scheduler ticket IDs exhausted"); + } + auto request = std::make_unique(); + request->id = ++mNextId; + request->prompt = prompt; + request->maxOutput = maxOutput; + request->cancelled = std::make_shared>(false); + SchedulerTicket ticket; + ticket.mId = request->id; + ticket.mCancelled = request->cancelled; + ticket.mResult = request->promise.get_future().share(); + mQueue.push_back(std::move(request)); + mQueuedBytes += bytes; + mWake.notify_one(); + return ticket; +} +void ContinuousScheduler::close() +{ + std::lock_guard joinLock(mCloseMutex); + { + std::lock_guard lock(mMutex); + mClosing = true; + } + mWake.notify_one(); + if (mWorker.joinable()) + { + mWorker.join(); + } +} +void ContinuousScheduler::emit(SchedulerEvent::Kind kind, Request const* request, uint64_t partner) +{ + if (!mObserver) + { + return; + } + SchedulerEvent event; + event.kind = kind; + event.partner = partner; + event.microseconds + = std::chrono::duration_cast(std::chrono::steady_clock::now().time_since_epoch()) + .count(); + if (request) + { + event.request = request->id; + event.handle = request->handle; + event.cursor = mBackend->state(request->handle).promptCursor(); + } + mObserver(event); +} +void ContinuousScheduler::terminal( + std::unique_ptr& request, SchedulerStatus status, bool release, std::exception_ptr error) +{ + SchedulerResult result; + result.status = status; + result.error = error; + if (release) + { + result.tokens = mBackend->state(request->handle).output(); + emit(SchedulerEvent::Kind::kRelease, request.get()); + mBackend->release(request->handle); + } + request->promise.set_value(std::move(result)); + request.reset(); +} +void ContinuousScheduler::boundary() +{ + // No GPU work is outstanding here. Reclaim before admission, including between decode and prefill. + for (auto& request : mActive) + { + if (!request) + { + continue; + } + if (request->cancelled->load()) + { + terminal(request, SchedulerStatus::kCancelled, true); + } + else if (mBackend->state(request->handle).phase() == SequencePhase::kFinished) + { + terminal(request, SchedulerStatus::kCompleted, true); + } + } + for (auto& request : mActive) + { + size_t examined = 0; + while (!request && examined++ < mMaxQueued) + { + { + std::lock_guard lock(mMutex); + if (mClosing || mQueue.empty()) + { + break; + } + request = std::move(mQueue.front()); + mQueue.pop_front(); + mQueuedBytes -= request->prompt.size() * sizeof(int32_t); + } + if (request->cancelled->load()) + { + terminal(request, SchedulerStatus::kCancelled, false); + continue; + } + request->handle = mBackend->acquire(request->id, std::move(request->prompt), request->maxOutput); + emit(SchedulerEvent::Kind::kAdmit, request.get()); + } + } +} +void ContinuousScheduler::run() noexcept +{ + try + { + mBackend->start(); + while (true) + { + { + std::unique_lock lock(mMutex); + if (!mClosing && mQueue.empty() && !mActive[0] && !mActive[1]) + { + lock.unlock(); + emit(SchedulerEvent::Kind::kIdle); + lock.lock(); + mWake.wait(lock, [&] { return mClosing || !mQueue.empty(); }); + } + if (mClosing) + { + break; + } + } + boundary(); + std::array handles{}; + std::array decoding{}; + int32_t count = 0; + for (auto& request : mActive) + { + if (request && mBackend->state(request->handle).phase() == SequencePhase::kDecode) + { + handles[count] = request->handle; + decoding[count++] = request.get(); + } + } + if (count == 2 && handles[0].slot > handles[1].slot) + { + std::swap(handles[0], handles[1]); + std::swap(decoding[0], decoding[1]); + } + if (count) + { + mBackend->decode(handles, count); + for (int32_t i = 0; i < count; ++i) + { + emit(SchedulerEvent::Kind::kDecode, decoding[i], count == 2 ? decoding[1 - i]->id : 0); + } + } + boundary(); + for (int32_t offset = 0; offset < 2; ++offset) + { + int32_t const slot = (mNextPrefill + offset) % 2; + auto& request = mActive[slot]; + if (request && mBackend->state(request->handle).phase() == SequencePhase::kPrefill) + { + mBackend->prefill(request->handle); + emit(SchedulerEvent::Kind::kPrefill, request.get()); + mNextPrefill = (slot + 1) % 2; + break; + } + } + boundary(); + } + for (auto& request : mActive) + { + if (request) + { + terminal(request, SchedulerStatus::kCancelled, true); + } + } + } + catch (...) + { + mHealthy.store(false); + mBackend->invalidate(); + auto const error = std::current_exception(); + for (auto& request : mActive) + { + if (request) + { + terminal(request, SchedulerStatus::kFailed, false, error); + } + } + } + std::lock_guard lock(mMutex); + mClosing = true; + for (auto& request : mQueue) + { + terminal(request, mHealthy.load() ? SchedulerStatus::kCancelled : SchedulerStatus::kFailed, false, + mHealthy.load() ? std::exception_ptr{} : std::make_exception_ptr(std::runtime_error("Worker failed"))); + } + mQueue.clear(); + mQueuedBytes = 0; +} +} // namespace rt +} // namespace trt_edgellm diff --git a/cpp/runtime/continuousScheduler.h b/cpp/runtime/continuousScheduler.h new file mode 100644 index 000000000..3fdac24dc --- /dev/null +++ b/cpp/runtime/continuousScheduler.h @@ -0,0 +1,146 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ +#pragma once +#include "runtime/state/sequenceSlots.h" +#include +#include +#include +#include +#include +#include +#include +#include + +namespace trt_edgellm +{ +namespace rt +{ +//! Forward methods are worker-only; validToken must read immutable metadata and be thread-safe. +//! Worker-only forward boundary. Implementations complete GPU work and accept one greedy token per ready row. +class SchedulerBackend +{ +public: + virtual ~SchedulerBackend() = default; + virtual void start() {} + virtual void invalidate() noexcept {} + virtual bool validToken(int32_t token) const + { + return token >= 0; + } + virtual SequenceHandle acquire(uint64_t id, std::vector prompt, int32_t maxOutput) = 0; + virtual SequenceState const& state(SequenceHandle handle) const = 0; + virtual void prefill(SequenceHandle handle) = 0; + virtual void decode(std::array const& handles, int32_t count) = 0; + virtual void release(SequenceHandle handle) = 0; +}; + +enum class SchedulerStatus +{ + kCompleted, + kCancelled, + kFailed +}; +struct SchedulerResult +{ + SchedulerStatus status{SchedulerStatus::kFailed}; + std::vector tokens; + std::exception_ptr error; +}; + +//! Ticket cancellation targets its submission, never a recycled physical slot. +class SchedulerTicket +{ +public: + uint64_t id() const + { + return mId; + } + void cancel() const + { + if (mCancelled) + { + mCancelled->store(true); + } + } + std::shared_future result() const + { + return mResult; + } + +private: + friend class ContinuousScheduler; + uint64_t mId{}; + std::shared_ptr> mCancelled; + std::shared_future mResult; +}; + +//! Fixed-size diagnostic event; callbacks run on the worker and must not block or call close(). +struct SchedulerEvent +{ + enum class Kind + { + kAdmit, + kPrefill, + kDecode, + kRelease, + kIdle + }; + Kind kind{}; + uint64_t request{}; + SequenceHandle handle{}; + uint64_t partner{}; + int32_t cursor{}; + int64_t microseconds{}; +}; + +//! Internal P4 scheduler: prepared text tokens, greedy sampling and length termination only. +//! Parent runtime and stream must outlive close/destruction. Public submission and close are thread-safe. +class ContinuousScheduler +{ +public: + using Observer = std::function; + ContinuousScheduler(std::unique_ptr backend, size_t maxQueued = 8, + size_t maxQueuedBytes = 256 * 1024, Observer observer = {}); + ~ContinuousScheduler(); + SchedulerTicket submit(std::vector const& prompt, int32_t maxOutput); + void close(); + bool healthy() const + { + return mHealthy.load(); + } + +private: + struct Request + { + uint64_t id{}; + std::vector prompt; + int32_t maxOutput{}; + std::shared_ptr> cancelled; + std::promise promise; + SequenceHandle handle{}; + }; + void run() noexcept; + void boundary(); + void emit(SchedulerEvent::Kind kind, Request const* request = nullptr, uint64_t partner = 0); + void terminal( + std::unique_ptr& request, SchedulerStatus status, bool release, std::exception_ptr error = {}); + std::unique_ptr mBackend; + size_t const mMaxQueued; + size_t const mMaxQueuedBytes; + Observer mObserver; + std::mutex mMutex; + std::mutex mCloseMutex; + std::condition_variable mWake; + std::deque> mQueue; + size_t mQueuedBytes{}; + uint64_t mNextId{}; + bool mClosing{}; + std::atomic mHealthy{true}; + std::array, 2> mActive; + int32_t mNextPrefill{}; + std::thread mWorker; +}; +} // namespace rt +} // namespace trt_edgellm diff --git a/cpp/runtime/greedySchedulerBackend.cpp b/cpp/runtime/greedySchedulerBackend.cpp new file mode 100644 index 000000000..bb2a3f16a --- /dev/null +++ b/cpp/runtime/greedySchedulerBackend.cpp @@ -0,0 +1,94 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ +#include "runtime/greedySchedulerBackend.h" +#include "common/checkMacros.h" +#include +#include + +namespace trt_edgellm +{ +namespace rt +{ +GreedySchedulerBackend::GreedySchedulerBackend(LLMRankRuntime& runtime, cudaStream_t stream, int32_t vocabularySize) + : mSteps(runtime, stream) + , mStream(stream) + , mVocabulary(vocabularySize) + , mHostLogits({2, vocabularySize}, DeviceType::kCPU, nvinfer1::DataType::kFLOAT) +{ + CUDA_CHECK(cudaGetDevice(&mDevice)); +} +void GreedySchedulerBackend::start() +{ + CUDA_CHECK(cudaSetDevice(mDevice)); +} +SequenceHandle GreedySchedulerBackend::acquire(uint64_t id, std::vector prompt, int32_t maxOutput) +{ + for (auto token : prompt) + { + if (token < 0 || token >= mVocabulary) + { + throw std::invalid_argument("Token outside vocabulary"); + } + } + SequenceOptions options; + options.maxOutputTokens = maxOutput; + options.temperature = 0.0F; + return mSteps.acquire(id, std::move(prompt), std::move(options)); +} +SequenceState const& GreedySchedulerBackend::state(SequenceHandle handle) const +{ + return mSteps.state(handle); +} +void GreedySchedulerBackend::sample(Tensor const& logits, std::array const& handles, int32_t count) +{ + if (logits.getDataType() != nvinfer1::DataType::kFLOAT || logits.getShape()[0] != count + || logits.getShape()[1] != mVocabulary) + { + throw std::runtime_error("Unexpected scheduler logits"); + } + CUDA_CHECK(cudaMemcpyAsync(mHostLogits.rawPointer(), logits.rawPointer(), count * mVocabulary * sizeof(float), + cudaMemcpyDeviceToHost, mStream)); + CUDA_CHECK(cudaStreamSynchronize(mStream)); + auto const* data = mHostLogits.dataPointer(); + for (int32_t row = 0; row < count; ++row) + { + auto const* values = data + row * mVocabulary; + int32_t best = 0; + for (int32_t token = 0; token < mVocabulary; ++token) + { + if (!std::isfinite(values[token])) + { + throw std::runtime_error("Nonfinite scheduler logits"); + } + if (values[token] > values[best]) + { + best = token; + } + } + mSteps.acceptToken(handles[row], best); + } +} +void GreedySchedulerBackend::prefill(SequenceHandle handle) +{ + auto const& logits = mSteps.beginPrefillChunk(handle); + mSteps.completeStep(); + if (state(handle).phase() == SequencePhase::kAwaitingSample) + { + sample(logits, {handle, {}}, 1); + } +} +void GreedySchedulerBackend::decode(std::array const& handles, int32_t count) +{ + auto const& logits = mSteps.beginDecode(handles, count); + mSteps.completeStep(); + sample(logits, handles, count); +} +void GreedySchedulerBackend::release(SequenceHandle handle) +{ + mSteps.finish(handle); + mSteps.release(handle); +} +} // namespace rt +} // namespace trt_edgellm diff --git a/cpp/runtime/greedySchedulerBackend.h b/cpp/runtime/greedySchedulerBackend.h new file mode 100644 index 000000000..37c26b1f6 --- /dev/null +++ b/cpp/runtime/greedySchedulerBackend.h @@ -0,0 +1,42 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ +#pragma once +#include "runtime/continuousScheduler.h" +#include "runtime/sequenceStepRuntime.h" + +namespace trt_edgellm +{ +namespace rt +{ +//! P4-only greedy adapter. Borrowed runtime and explicit stream outlive this exclusive lease. +class GreedySchedulerBackend : public SchedulerBackend +{ +public: + GreedySchedulerBackend(LLMRankRuntime& runtime, cudaStream_t stream, int32_t vocabularySize); + void start() override; + void invalidate() noexcept override + { + mSteps.poison(); + } + bool validToken(int32_t token) const override + { + return token >= 0 && token < mVocabulary; + } + SequenceHandle acquire(uint64_t id, std::vector prompt, int32_t maxOutput) override; + SequenceState const& state(SequenceHandle handle) const override; + void prefill(SequenceHandle handle) override; + void decode(std::array const& handles, int32_t count) override; + void release(SequenceHandle handle) override; + +private: + void sample(Tensor const& logits, std::array const& handles, int32_t count); + SequenceStepRuntime mSteps; + cudaStream_t mStream; + int32_t mVocabulary; + int mDevice{}; + Tensor mHostLogits; +}; +} // namespace rt +} // namespace trt_edgellm diff --git a/cpp/runtime/sequenceStepRuntime.cpp b/cpp/runtime/sequenceStepRuntime.cpp index 77a0f5d98..e44bd2e29 100644 --- a/cpp/runtime/sequenceStepRuntime.cpp +++ b/cpp/runtime/sequenceStepRuntime.cpp @@ -118,6 +118,11 @@ SequenceStepRuntime::~SequenceStepRuntime() } } +void SequenceStepRuntime::poison() noexcept +{ + mLease->poisoned = true; +} + bool SequenceStepRuntime::healthy() const noexcept { return !mLease->poisoned; diff --git a/cpp/runtime/sequenceStepRuntime.h b/cpp/runtime/sequenceStepRuntime.h index ca51d7ed6..eb0f3ab6b 100644 --- a/cpp/runtime/sequenceStepRuntime.h +++ b/cpp/runtime/sequenceStepRuntime.h @@ -41,6 +41,8 @@ class SequenceStepRuntime void finish(SequenceHandle handle); void release(SequenceHandle handle); bool healthy() const noexcept; + //! Retain the parent lease after an external sampling/worker failure. Recreate the parent to recover. + void poison() noexcept; private: struct Lease; diff --git a/examples/llm/continuousBatchingProbe.cpp b/examples/llm/continuousBatchingProbe.cpp index e2d0ce3b8..8bb02a8de 100644 --- a/examples/llm/continuousBatchingProbe.cpp +++ b/examples/llm/continuousBatchingProbe.cpp @@ -6,6 +6,7 @@ #include "common/bindingNames.h" #include "common/trtUtils.h" #include "kernels/posEncoding/initializeCosSinCache.h" +#include "runtime/greedySchedulerBackend.h" #include "runtime/llmRankRuntime.h" #include "runtime/sequenceStepRuntime.h" @@ -884,6 +885,98 @@ bool runChunkTests(LLMRankRuntime& runtime, tokenizer::Tokenizer& tokenizer, cud std::cout << "P3_QUALITY_GATE passed=" << passed << std::endl; return passed; } +bool runSchedulerTests(LLMRankRuntime& runtime, tokenizer::Tokenizer& tokenizer, cudaStream_t stream) +{ + ContinuousBatchingProbe observer(runtime, stream); + auto const words = tokenizer.encode("A red fox crosses a blue river. One two three four. "); + auto prompt = [&](int32_t length) { + std::vector tokens; + for (int32_t i = 0; i < length; ++i) + { + tokens.push_back(words.at(i % words.size())); + } + return tokens; + }; + auto const pa = prompt(65), pb = prompt(1025), pc = prompt(129); + std::vector events; + events.reserve(1024); + bool armed = false, submitted = false; + ContinuousScheduler* owner = nullptr; + SchedulerTicket b, c; + auto backend = std::make_unique(runtime, stream, observer.vocabularySize()); + ContinuousScheduler scheduler(std::move(backend), 8, 65536, [&](SchedulerEvent const& event) { + if (event.kind == SchedulerEvent::Kind::kIdle) + { + return; + } + events.push_back(event); + if (armed && !submitted && event.kind == SchedulerEvent::Kind::kDecode) + { + submitted = true; + b = owner->submit(pb, 12); + c = owner->submit(pc, 12); + } + }); + owner = &scheduler; + auto result = [&](SchedulerTicket const& ticket) { + require(ticket.result().wait_for(std::chrono::seconds(60)) == std::future_status::ready, "Stranded ticket"); + auto value = ticket.result().get(); + require(value.status == SchedulerStatus::kCompleted, "Scheduler request failed"); + return value.tokens; + }; + auto const ra = result(scheduler.submit(pa, 6)); + auto const rb = result(scheduler.submit(pb, 12)); + auto const rc = result(scheduler.submit(pc, 12)); + // Configure the diagnostic through submission's mutex happens-before boundary. + armed = true; + auto a = scheduler.submit(pa, 6); + auto const actualA = result(a); + auto const actualB = result(b); + auto const actualC = result(c); + scheduler.close(); + bool const exact = actualA == ra && actualB == rb && actualC == rc; + SequenceHandle ah{}, ch{}; + bool aReleased = false, reused = false, bContinued = false, pair = false, staggered = false; + for (auto const& e : events) + { + std::cout << "P4_EVENT kind=" << static_cast(e.kind) << " request=" << e.request + << " slot=" << e.handle.slot << " generation=" << e.handle.generation << " partner=" << e.partner + << " cursor=" << e.cursor << " us=" << e.microseconds << std::endl; + if (e.request == a.id() && e.kind == SchedulerEvent::Kind::kAdmit) + { + ah = e.handle; + } + if (e.request == a.id() && e.kind == SchedulerEvent::Kind::kDecode) + { + staggered = true; + } + if (e.request == b.id() && e.kind == SchedulerEvent::Kind::kAdmit) + { + require(staggered, "B admitted before A decode"); + } + if (e.request == a.id() && e.kind == SchedulerEvent::Kind::kRelease) + { + aReleased = true; + } + if (e.request == c.id() && e.kind == SchedulerEvent::Kind::kAdmit) + { + ch = e.handle; + reused = aReleased && ah.slot == ch.slot && ch.generation > ah.generation; + } + if (reused && e.request == b.id() && e.kind == SchedulerEvent::Kind::kPrefill) + { + bContinued = true; + } + if (e.partner && (e.request == b.id() || e.request == c.id())) + { + pair = true; + } + } + bool const passed = exact && reused && bContinued && pair && scheduler.healthy(); + std::cout << "P4_SCHEDULER_GATE passed=" << passed << " exact_outputs=" << exact << " reused=" << reused + << " b_continued=" << bContinued << " paired_decode=" << pair << std::endl; + return passed; +} } // namespace rt } // namespace trt_edgellm @@ -891,10 +984,11 @@ int main(int argc, char** argv) { if (argc != 3 && (argc != 4 - || (std::string(argv[3]) != "--steps" + || (std::string(argv[3]) != "--scheduler" && std::string(argv[3]) != "--steps" && (std::string(argv[3]) != "--chunks" && std::string(argv[3]) != "--chunks-extra")))) { - std::cerr << "Usage: continuous_batching_probe ENGINE_DIR CHECKPOINT_DIR [--steps|--chunks|--chunks-extra]\n"; + std::cerr << "Usage: continuous_batching_probe ENGINE_DIR CHECKPOINT_DIR " + "[--steps|--chunks|--chunks-extra|--scheduler]\n"; return 2; } cudaStream_t stream{}; @@ -909,11 +1003,13 @@ int main(int argc, char** argv) trt_edgellm::rt::require(tokenizer.loadFromHF(argv[1], false), "Tokenizer loading failed"); trt_edgellm::rt::LLMRankRuntime runtime(argv[1], "", {}, std::nullopt, stream, trt_edgellm::rt::ParallelMapping{}, tokenizer, trt_edgellm::rt::ContextCacheConfig{}, argv[2], ""); - passed = argc == 4 ? (std::string(argv[3]).find("--chunks") == 0 - ? trt_edgellm::rt::runChunkTests( - runtime, tokenizer, stream, std::string(argv[3]) == "--chunks-extra") - : trt_edgellm::rt::runStepTests(runtime, tokenizer, stream)) - : trt_edgellm::rt::runProbe(runtime, tokenizer, stream); + passed = argc == 4 && std::string(argv[3]) == "--scheduler" + ? trt_edgellm::rt::runSchedulerTests(runtime, tokenizer, stream) + : argc == 4 ? (std::string(argv[3]).find("--chunks") == 0 + ? trt_edgellm::rt::runChunkTests( + runtime, tokenizer, stream, std::string(argv[3]) == "--chunks-extra") + : trt_edgellm::rt::runStepTests(runtime, tokenizer, stream)) + : trt_edgellm::rt::runProbe(runtime, tokenizer, stream); } CUDA_CHECK(cudaStreamDestroy(stream)); std::cout << "PROBE_RESULT passed=" << passed << std::endl; diff --git a/examples/llm/continuousScheduler.md b/examples/llm/continuousScheduler.md new file mode 100644 index 000000000..49118dde3 --- /dev/null +++ b/examples/llm/continuousScheduler.md @@ -0,0 +1,82 @@ +# P4: continuous native scheduler + +`ContinuousScheduler` owns one worker and a `SchedulerBackend`. The TensorRT +adapter `GreedySchedulerBackend` owns one exclusive `SequenceStepRuntime` lease +on an existing parent runtime and explicit stream. Construct it before legacy +inference; keep the parent and stream alive through scheduler destruction. +Do not invoke any parent API during this lease. `close()` joins the worker; +the lease remains held until the scheduler is destroyed. + +## Submission and ownership + +`submit(preparedTokenIds, maxOutputTokens)` returns a ticket immediately after +bounded CPU admission. The scheduler copies the once-formatted/tokenized prompt +only after checking the queue bounds. Defaults are eight queued requests and +256 KiB of queued token payload, in addition to at most two resident requests. +The count cap also bounds metadata; each prompt is at most 6144 tokens and +prompt plus output at most 8192. Results contain at most the requested output +cap. Consumers own completed futures; retaining arbitrary completed tickets is +caller-owned memory, not an internal result backlog. + +FIFO admission, step execution and release belong exclusively to the worker. +Each round executes one decode for every ready row, followed by at most one +qualified prefill chunk. Prefills alternate when both slots need prompt work, +so each gets one chunk per two rounds. Every chunk uses `beginPrefillChunk` +from the very first step; the P3 cap and short-tail restrictions are unchanged. +There is no mixed prefill/decode engine invocation and no execution overlap. + +Safe boundaries before/after forwards reclaim finished or cancelled slots and +admit queued work immediately. Cancelled queue heads are skipped in a bounded +scan. Physical rows never move; two-row decode handles are sorted into physical +order. Idle workers sleep on a condition variable, woken by submit or close. + +`ticket.cancel()` sets that submission's atomic flag, so a stale ticket cannot +cancel a reused slot. Active cancellation is observed at safe boundaries; +queued cancellation is settled on admission or shutdown, and does not remove a +queued item immediately while both slots remain occupied. This is P4 basic +plumbing, not P5's complete cancellation/deadline contract. `close()` rejects +new work, cancels outstanding requests and joins the worker. Concurrent calls +to close are serialized. A completion racing cancellation may finish normally. + +Any worker exception fails all outstanding tickets and rejects further +submission. The TensorRT adapter also poisons the parent lease, including errors +in D2H/sampling after forward completion. Recovery requires destroying and +recreating the parent runtime. This deliberately conservative behavior is not +fine-grained P5 fault recovery. Never call close or destroy the scheduler from +its worker/observer callback. + +## Deliberately restricted generation + +The P4 adapter accepts only token IDs and an output limit. It selects finite +FP32 logits greedily and stops at the length cap. It does not honor EOS, +temperature/top-p/top-k, RNG, stop strings, thinking, logprobs or streaming. +There is no public options argument that silently ignores those settings. +The temporary sampler copies logits to one preallocated pinned two-row buffer +and scans on the CPU; efficient GPU sampling and tuning belong to later phases. +No state-sized copies, per-token heap allocation, graphs or second model are +introduced by the scheduler. Diagnostic observers may allocate; production +callers should omit them. The backend validToken method must read immutable +metadata safely from submission threads; all other backend methods are owned +by the worker after construction. + +## Verification + +`continuous_batching_probe ENGINE_DIR CHECKPOINT_DIR --scheduler` compares +serial greedy reference outputs with automatic staggered A/B/C execution: +A=65 prompt/6 output tokens, B=1025/12, C=129/12. B and C are submitted from the +diagnostic observer after A's first decode. The observer only submits work; +all admission, selection, forwarding and slot release are automatic. +`P4_EVENT` records steady-clock microseconds, request, slot, generation, partner +and prompt cursor (kAdmit=0, kPrefill=1, kDecode=2, kRelease=3, kIdle=4). +`P4_SCHEDULER_GATE` requires exact serial token equality, later B admission, +C reuse after A release, continued B prefill after reuse and paired B/C decode. +These are mechanism/greedy-output checks, not broad numerical or performance +qualification. P3's state/logit matrix remains the chunk-policy evidence. + +The CPU tests in `unittests/cpp/runtime/continuousSchedulerTest.cpp` exercise +queue count/byte bounds, token limits, FIFO reuse, two-prefill fairness, +cancellation, failure settlement, shutdown, idle sleep and concurrent +producers with repeated reuse. See the OpenClaw P4 results for exact runs, +failed fixture coverage, sanitizer checks, source identity and restoration. +HTTP still admits one sequence; independent request policy is P5, HTTP/SSE is +P6, performance/graphs are P7 and qualified deployment is P8. diff --git a/unittests/cpp/runtime/continuousSchedulerTest.cpp b/unittests/cpp/runtime/continuousSchedulerTest.cpp new file mode 100644 index 000000000..aac09e11c --- /dev/null +++ b/unittests/cpp/runtime/continuousSchedulerTest.cpp @@ -0,0 +1,292 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ +#include "runtime/continuousScheduler.h" +#include "runtime/state/prefillChunk.h" +#include +#include + +using namespace trt_edgellm::rt; +using namespace std::chrono_literals; +namespace +{ +class FakeBackend : public SchedulerBackend +{ +public: + SequenceSlots slots{2, 6144, 8192}; + std::shared_future gate; + bool fail{}; + void start() override + { + if (gate.valid()) + { + gate.wait(); + } + } + SequenceHandle acquire(uint64_t id, std::vector prompt, int32_t maxOutput) override + { + SequenceOptions options; + options.maxOutputTokens = maxOutput; + return slots.acquire(id, std::move(prompt), options); + } + SequenceState const& state(SequenceHandle h) const override + { + return slots.get(h); + } + void prefill(SequenceHandle h) override + { + if (fail) + { + throw std::runtime_error("injected forward failure"); + } + auto const& s = state(h); + slots.commitPrompt(h, nextPrefillChunkSize(static_cast(s.prompt().size()) - s.promptCursor())); + if (state(h).phase() == SequencePhase::kAwaitingSample) + { + slots.acceptToken(h, 7); + } + } + void decode(std::array const& handles, int32_t count) override + { + if (count == 2 && handles[0].slot >= handles[1].slot) + { + throw std::runtime_error("unordered views"); + } + for (int i = 0; i < count; ++i) + { + slots.commitDecode(handles[i]); + slots.acceptToken(handles[i], 7); + } + } + void release(SequenceHandle h) override + { + slots.finish(h); + slots.release(h); + } +}; +SchedulerResult result(SchedulerTicket const& ticket) +{ + if (ticket.result().wait_for(3s) != std::future_status::ready) + { + throw std::runtime_error("Stranded ticket"); + } + return ticket.result().get(); +} +} // namespace +TEST(ContinuousScheduler, StaggeredReuseAndDecodeFirst) +{ + std::vector events; + ContinuousScheduler* owner = nullptr; + SchedulerTicket b, c; + bool submitted = false; + ContinuousScheduler scheduler(std::make_unique(), 8, 65536, [&](auto const& e) { + events.push_back(e); + if (e.kind == SchedulerEvent::Kind::kDecode && !submitted) + { + submitted = true; + b = owner->submit(std::vector(1025, 1), 12); + c = owner->submit(std::vector(129, 2), 12); + } + }); + owner = &scheduler; + auto a = scheduler.submit({1}, 6); + EXPECT_EQ(result(a).tokens.size(), 6U); + // A completion synchronizes the callback's ticket publication. + EXPECT_EQ(result(b).tokens.size(), 12U); + EXPECT_EQ(result(c).tokens.size(), 12U); + scheduler.close(); + SequenceHandle ah{}, ch{}; + bool bContinued = false, cAdmitted = false, sawPair = false; + for (auto const& e : events) + { + if (e.kind == SchedulerEvent::Kind::kAdmit && e.request == a.id()) + { + ah = e.handle; + } + if (e.kind == SchedulerEvent::Kind::kAdmit && e.request == c.id()) + { + ch = e.handle; + cAdmitted = true; + } + if (cAdmitted && e.request == b.id() && e.kind == SchedulerEvent::Kind::kPrefill) + { + bContinued = true; + } + if (e.kind == SchedulerEvent::Kind::kDecode && e.partner) + { + sawPair = true; + } + } + EXPECT_EQ(ah.slot, ch.slot); + EXPECT_GT(ch.generation, ah.generation); + EXPECT_TRUE(bContinued); + EXPECT_TRUE(sawPair); + EXPECT_TRUE(scheduler.healthy()); +} +TEST(ContinuousScheduler, QueueBoundsCancellationAndShutdown) +{ + std::promise gate; + auto backend = std::make_unique(); + backend->gate = gate.get_future().share(); + ContinuousScheduler scheduler(std::move(backend), 2, 8); + auto a = scheduler.submit({1}, 3); + auto b = scheduler.submit({2}, 3); + EXPECT_THROW(scheduler.submit({3}, 3), std::runtime_error); + a.cancel(); + gate.set_value(); + EXPECT_EQ(result(a).status, SchedulerStatus::kCancelled); + EXPECT_EQ(result(b).status, SchedulerStatus::kCompleted); + a.cancel(); + auto c = scheduler.submit({4}, 3); + EXPECT_EQ(result(c).status, SchedulerStatus::kCompleted); + scheduler.close(); + scheduler.close(); + EXPECT_THROW(scheduler.submit({1}, 1), std::runtime_error); +} +TEST(ContinuousScheduler, ByteBoundAndValidation) +{ + std::promise gate; + auto backend = std::make_unique(); + backend->gate = gate.get_future().share(); + ContinuousScheduler scheduler(std::move(backend), 10, 4); + auto a = scheduler.submit({1}, 1); + EXPECT_THROW(scheduler.submit({1}, 1), std::runtime_error); + EXPECT_THROW(scheduler.submit({}, 1), std::invalid_argument); + EXPECT_THROW(scheduler.submit({-1}, 1), std::invalid_argument); + EXPECT_THROW(scheduler.submit({1}, 8192), std::invalid_argument); + gate.set_value(); + EXPECT_EQ(result(a).tokens.size(), 1U); +} +TEST(ContinuousScheduler, FaultSettlesActiveAndQueued) +{ + std::promise gate; + auto backend = std::make_unique(); + backend->gate = gate.get_future().share(); + backend->fail = true; + ContinuousScheduler scheduler(std::move(backend)); + auto a = scheduler.submit({1}, 3); + auto b = scheduler.submit({1}, 3); + auto c = scheduler.submit({1}, 3); + gate.set_value(); + for (auto const& ticket : {a, b, c}) + { + EXPECT_EQ(result(ticket).status, SchedulerStatus::kFailed); + } + EXPECT_FALSE(scheduler.healthy()); + EXPECT_THROW(scheduler.submit({1}, 1), std::runtime_error); +} +TEST(ContinuousScheduler, IdleWaitAndWake) +{ + std::atomic idle{0}; + ContinuousScheduler scheduler(std::make_unique(), 8, 65536, [&](auto const& e) { + if (e.kind == SchedulerEvent::Kind::kIdle) + { + ++idle; + } + }); + std::this_thread::sleep_for(30ms); + EXPECT_EQ(idle.load(), 1); + std::this_thread::sleep_for(30ms); + EXPECT_EQ(idle.load(), 1); + EXPECT_EQ(result(scheduler.submit({1}, 1)).status, SchedulerStatus::kCompleted); + scheduler.close(); +} +TEST(ContinuousScheduler, TwoPrefillsFairnessAndActiveCancellation) +{ + std::promise gate; + auto backend = std::make_unique(); + backend->gate = gate.get_future().share(); + SchedulerTicket a; + std::vector prefills; + ContinuousScheduler scheduler(std::move(backend), 8, 65536, [&](auto const& e) { + if (e.kind == SchedulerEvent::Kind::kPrefill) + { + prefills.push_back(e.request); + if (prefills.size() == 3) + { + a.cancel(); + } + } + }); + a = scheduler.submit(std::vector(1025, 1), 4); + auto b = scheduler.submit(std::vector(1025, 1), 4); + gate.set_value(); + EXPECT_EQ(result(a).status, SchedulerStatus::kCancelled); + EXPECT_EQ(result(b).status, SchedulerStatus::kCompleted); + scheduler.close(); + ASSERT_GE(prefills.size(), 3U); + EXPECT_EQ(prefills[0], a.id()); + EXPECT_EQ(prefills[1], b.id()); + EXPECT_EQ(prefills[2], a.id()); +} +TEST(ContinuousScheduler, CloseSettlesAllTickets) +{ + std::promise gate; + auto backend = std::make_unique(); + backend->gate = gate.get_future().share(); + ContinuousScheduler scheduler(std::move(backend)); + auto a = scheduler.submit({1}, 3); + auto b = scheduler.submit({1}, 3); + auto closing = std::async(std::launch::async, [&] { scheduler.close(); }); + gate.set_value(); + closing.get(); + EXPECT_NE(result(a).status, SchedulerStatus::kFailed); + EXPECT_NE(result(b).status, SchedulerStatus::kFailed); +} + +TEST(ContinuousScheduler, CancelledQueueHeadDoesNotDelayAdmission) +{ + std::promise gate; + auto backend = std::make_unique(); + backend->gate = gate.get_future().share(); + std::vector events; + ContinuousScheduler scheduler(std::move(backend), 8, 65536, [&](auto const& e) { events.push_back(e); }); + auto cancelled = scheduler.submit({1}, 3); + auto a = scheduler.submit({1}, 3); + auto b = scheduler.submit({1}, 3); + cancelled.cancel(); + gate.set_value(); + EXPECT_EQ(result(cancelled).status, SchedulerStatus::kCancelled); + EXPECT_EQ(result(a).status, SchedulerStatus::kCompleted); + EXPECT_EQ(result(b).status, SchedulerStatus::kCompleted); + scheduler.close(); + int admissions = 0; + for (auto const& e : events) + { + if (e.kind == SchedulerEvent::Kind::kAdmit) + { + ++admissions; + } + if (e.kind == SchedulerEvent::Kind::kPrefill) + { + EXPECT_EQ(admissions, 2); + break; + } + } +} +TEST(ContinuousScheduler, ConcurrentProducersAndRepeatedReuse) +{ + ContinuousScheduler scheduler(std::make_unique(), 128, 65536); + std::array, 4> producers; + for (auto& producer : producers) + { + producer = std::async(std::launch::async, [&] { + for (int i = 0; i < 25; ++i) + { + auto ticket = scheduler.submit({1, 2, 3}, 3); + if (result(ticket).tokens.size() != 3) + { + throw std::runtime_error("Lost result"); + } + ticket.cancel(); + } + }); + } + for (auto& producer : producers) + { + producer.get(); + } + scheduler.close(); + EXPECT_TRUE(scheduler.healthy()); +} From 711270ff229fd2cd0b25ad25846745c4ebaea446 Mon Sep 17 00:00:00 2001 From: ajmalrasi Date: Tue, 15 Sep 2026 15:21:32 +0530 Subject: [PATCH 4/7] feat: Isolate native sequence policy cancellation and output Signed-off-by: ajmalrasi --- cpp/runtime/continuousScheduler.cpp | 141 +++++++-- cpp/runtime/continuousScheduler.h | 72 ++++- cpp/runtime/greedySchedulerBackend.cpp | 114 +++++-- cpp/runtime/greedySchedulerBackend.h | 33 ++- cpp/runtime/state/sequenceChannel.cpp | 68 +++++ cpp/runtime/state/sequenceChannel.h | 50 ++++ cpp/runtime/state/sequencePolicy.cpp | 278 ++++++++++++++++++ cpp/runtime/state/sequencePolicy.h | 94 ++++++ cpp/runtime/state/sequenceSlots.h | 4 + examples/llm/continuousBatchingProbe.cpp | 244 ++++++++++++++- .../cpp/runtime/continuousSchedulerTest.cpp | 109 ++++++- .../cpp/runtime/state/sequencePolicyTest.cpp | 204 +++++++++++++ 12 files changed, 1339 insertions(+), 72 deletions(-) create mode 100644 cpp/runtime/state/sequenceChannel.cpp create mode 100644 cpp/runtime/state/sequenceChannel.h create mode 100644 cpp/runtime/state/sequencePolicy.cpp create mode 100644 cpp/runtime/state/sequencePolicy.h create mode 100644 unittests/cpp/runtime/state/sequencePolicyTest.cpp diff --git a/cpp/runtime/continuousScheduler.cpp b/cpp/runtime/continuousScheduler.cpp index 042d2f4e7..d263ef133 100644 --- a/cpp/runtime/continuousScheduler.cpp +++ b/cpp/runtime/continuousScheduler.cpp @@ -31,10 +31,24 @@ ContinuousScheduler::~ContinuousScheduler() SchedulerTicket ContinuousScheduler::submit(std::vector const& prompt, int32_t maxOutput) { - if (prompt.empty() || prompt.size() > 6144 || maxOutput <= 0 - || maxOutput > 8192 - static_cast(prompt.size())) + SchedulerRequestOptions options; + options.generation.maxOutputTokens = maxOutput; + options.generation.temperature = 0; + options.generation.ignoreEos = true; + return submit(prompt, options); +} +SchedulerTicket ContinuousScheduler::submit(std::vector const& prompt, SchedulerRequestOptions const& input) +{ + validateSequenceOptions(input.generation, mBackend->vocabularySize()); + auto options = input; + options.generation = mBackend->normalizeOptions(options.generation); + auto const& generation = options.generation; + size_t const metadata = validateSequenceOptions(generation, mBackend->vocabularySize()); + if (prompt.empty() || prompt.size() > 6144 + || generation.maxOutputTokens > 8192 - static_cast(prompt.size()) || options.streamRecords > 8192 + || (options.streamRecords && (!options.streamBytes || options.streamBytes > 1048576))) { - throw std::invalid_argument("P4 request exceeds qualified token bounds"); + throw std::invalid_argument("Request exceeds native bounds"); } for (auto token : prompt) { @@ -43,12 +57,15 @@ SchedulerTicket ContinuousScheduler::submit(std::vector const& prompt, throw std::invalid_argument("Token outside supported vocabulary"); } } + size_t const bytes = prompt.size() * sizeof(int32_t) + metadata + + (options.streamRecords + ? options.streamRecords * (sizeof(SequenceSample) + sizeof(size_t)) + options.streamBytes + : 0); std::lock_guard lock(mMutex); if (mClosing || !mHealthy.load()) { throw std::runtime_error("Scheduler closed or failed"); } - size_t const bytes = prompt.size() * sizeof(int32_t); if (mQueue.size() >= mMaxQueued || bytes > mMaxQueuedBytes - mQueuedBytes) { throw std::runtime_error("Scheduler queue full"); @@ -60,11 +77,19 @@ SchedulerTicket ContinuousScheduler::submit(std::vector const& prompt, auto request = std::make_unique(); request->id = ++mNextId; request->prompt = prompt; - request->maxOutput = maxOutput; + request->promptTokens = static_cast(prompt.size()); + request->options = std::move(options); + request->bytes = bytes; request->cancelled = std::make_shared>(false); + if (request->options.streamRecords) + { + request->channel + = std::make_shared(request->options.streamRecords, request->options.streamBytes); + } SchedulerTicket ticket; ticket.mId = request->id; ticket.mCancelled = request->cancelled; + ticket.mChannel = request->channel; ticket.mResult = request->promise.get_future().share(); mQueue.push_back(std::move(request)); mQueuedBytes += bytes; @@ -110,56 +135,122 @@ void ContinuousScheduler::terminal( SchedulerResult result; result.status = status; result.error = error; + result.promptTokens = request->promptTokens; if (release) { + mBackend->finalize(request->handle); + auto const text = mBackend->text(request->handle); + if (request->channel && text.size() > request->publishedBytes && status != SchedulerStatus::kSlowConsumer) + { + SequenceSample flush; + if (!request->channel->push(flush, text.substr(request->publishedBytes))) + { + result.status = SchedulerStatus::kSlowConsumer; + } + } result.tokens = mBackend->state(request->handle).output(); + result.randomCounter = mBackend->state(request->handle).randomCounter(); + if (!text.empty()) + { + result.text.assign(text.data(), text.size()); + } + result.logprobs = mBackend->logprobs(request->handle); + result.finish = mBackend->finishReason(request->handle); emit(SchedulerEvent::Kind::kRelease, request.get()); mBackend->release(request->handle); } + auto channel = request->channel; + if (channel) + { + channel->close(); + } request->promise.set_value(std::move(result)); request.reset(); } +void ContinuousScheduler::publish(std::unique_ptr& request) +{ + auto const count = mBackend->state(request->handle).output().size(); + if (count == request->publishedTokens) + { + return; + } + auto const text = mBackend->text(request->handle); + if (request->channel + && !request->channel->push(mBackend->lastSample(request->handle), text.substr(request->publishedBytes))) + { + terminal(request, SchedulerStatus::kSlowConsumer, true); + return; + } + request->publishedTokens = count; + request->publishedBytes = text.size(); +} void ContinuousScheduler::boundary() { - // No GPU work is outstanding here. Reclaim before admission, including between decode and prefill. + auto const now = std::chrono::steady_clock::now(); for (auto& request : mActive) { if (!request) { continue; } - if (request->cancelled->load()) + bool const finished = mBackend->state(request->handle).phase() == SequencePhase::kFinished; + // Natural completion is already committed by the forward before a late cancellation is observed. + if (finished) + { + publish(request); + if (request) + { + terminal(request, SchedulerStatus::kCompleted, true); + } + } + else if (request->cancelled->load()) { terminal(request, SchedulerStatus::kCancelled, true); } - else if (mBackend->state(request->handle).phase() == SequencePhase::kFinished) + else if (now >= request->options.deadline) + { + terminal(request, SchedulerStatus::kDeadline, true); + } + else { - terminal(request, SchedulerStatus::kCompleted, true); + publish(request); } } - for (auto& request : mActive) { - size_t examined = 0; - while (!request && examined++ < mMaxQueued) + std::lock_guard lock(mMutex); + for (auto it = mQueue.begin(); it != mQueue.end();) { + bool const cancelled = (*it)->cancelled->load(); + if (cancelled || now >= (*it)->options.deadline || now >= (*it)->options.queueDeadline) { - std::lock_guard lock(mMutex); - if (mClosing || mQueue.empty()) - { - break; - } - request = std::move(mQueue.front()); - mQueue.pop_front(); - mQueuedBytes -= request->prompt.size() * sizeof(int32_t); + mQueuedBytes -= (*it)->bytes; + terminal(*it, cancelled ? SchedulerStatus::kCancelled : SchedulerStatus::kDeadline, false); + it = mQueue.erase(it); } - if (request->cancelled->load()) + else + { + ++it; + } + } + } + for (auto& request : mActive) + { + if (request) + { + continue; + } + { + std::lock_guard lock(mMutex); + if (mClosing || mQueue.empty()) { - terminal(request, SchedulerStatus::kCancelled, false); - continue; + break; } - request->handle = mBackend->acquire(request->id, std::move(request->prompt), request->maxOutput); - emit(SchedulerEvent::Kind::kAdmit, request.get()); + request = std::move(mQueue.front()); + mQueue.pop_front(); + mQueuedBytes -= request->bytes; } + request->handle = mBackend->acquire(request->id, std::move(request->prompt), request->options.generation); + emit(SchedulerEvent::Kind::kAdmit, request.get()); } } void ContinuousScheduler::run() noexcept diff --git a/cpp/runtime/continuousScheduler.h b/cpp/runtime/continuousScheduler.h index 3fdac24dc..0ffba447a 100644 --- a/cpp/runtime/continuousScheduler.h +++ b/cpp/runtime/continuousScheduler.h @@ -3,11 +3,13 @@ * SPDX-License-Identifier: Apache-2.0 */ #pragma once +#include "runtime/state/sequenceChannel.h" #include "runtime/state/sequenceSlots.h" #include #include #include #include +#include #include #include #include @@ -18,7 +20,7 @@ namespace trt_edgellm namespace rt { //! Forward methods are worker-only; validToken must read immutable metadata and be thread-safe. -//! Worker-only forward boundary. Implementations complete GPU work and accept one greedy token per ready row. +//! Forward implementations finish GPU work and apply each row’s independent policy before returning. class SchedulerBackend { public: @@ -29,7 +31,35 @@ class SchedulerBackend { return token >= 0; } - virtual SequenceHandle acquire(uint64_t id, std::vector prompt, int32_t maxOutput) = 0; + virtual int32_t vocabularySize() const + { + return 248320; + } + virtual SequenceOptions normalizeOptions(SequenceOptions options) const + { + return options; + } + virtual SequenceHandle acquire(uint64_t id, std::vector prompt, SequenceOptions options) = 0; + virtual SequenceSample lastSample(SequenceHandle handle) const + { + SequenceSample sample; + sample.token = state(handle).output().back(); + return sample; + } + virtual std::string_view text(SequenceHandle) const + { + return {}; + } + virtual std::vector const& logprobs(SequenceHandle) const + { + static std::vector const empty; + return empty; + } + virtual SequenceFinish finishReason(SequenceHandle) const + { + return SequenceFinish::kLength; + } + virtual void finalize(SequenceHandle) {} virtual SequenceState const& state(SequenceHandle handle) const = 0; virtual void prefill(SequenceHandle handle) = 0; virtual void decode(std::array const& handles, int32_t count) = 0; @@ -40,13 +70,20 @@ enum class SchedulerStatus { kCompleted, kCancelled, - kFailed + kFailed, + kDeadline, + kSlowConsumer }; struct SchedulerResult { SchedulerStatus status{SchedulerStatus::kFailed}; std::vector tokens; std::exception_ptr error; + std::string text; + std::vector logprobs; + SequenceFinish finish{SequenceFinish::kNone}; + int32_t promptTokens{}; + uint64_t randomCounter{}; }; //! Ticket cancellation targets its submission, never a recycled physical slot. @@ -64,6 +101,14 @@ class SchedulerTicket mCancelled->store(true); } } + SequenceRead read(std::chrono::milliseconds timeout) const + { + if (!mChannel) + { + throw std::logic_error("Ticket has no stream"); + } + return mChannel->read(timeout); + } std::shared_future result() const { return mResult; @@ -73,6 +118,7 @@ class SchedulerTicket friend class ContinuousScheduler; uint64_t mId{}; std::shared_ptr> mCancelled; + std::shared_ptr mChannel; std::shared_future mResult; }; @@ -95,7 +141,16 @@ struct SchedulerEvent int64_t microseconds{}; }; -//! Internal P4 scheduler: prepared text tokens, greedy sampling and length termination only. +struct SchedulerRequestOptions +{ + SequenceOptions generation; + std::chrono::steady_clock::time_point queueDeadline{std::chrono::steady_clock::time_point::max()}; + std::chrono::steady_clock::time_point deadline{std::chrono::steady_clock::time_point::max()}; + size_t streamRecords{}; + size_t streamBytes{16384}; +}; + +//! Independent request policy with one execution owner and bounded admission/output channels. //! Parent runtime and stream must outlive close/destruction. Public submission and close are thread-safe. class ContinuousScheduler { @@ -105,6 +160,7 @@ class ContinuousScheduler size_t maxQueuedBytes = 256 * 1024, Observer observer = {}); ~ContinuousScheduler(); SchedulerTicket submit(std::vector const& prompt, int32_t maxOutput); + SchedulerTicket submit(std::vector const& prompt, SchedulerRequestOptions const& options); void close(); bool healthy() const { @@ -116,13 +172,19 @@ class ContinuousScheduler { uint64_t id{}; std::vector prompt; - int32_t maxOutput{}; + SchedulerRequestOptions options; + size_t bytes{}; + size_t publishedTokens{}; + size_t publishedBytes{}; + int32_t promptTokens{}; + std::shared_ptr channel; std::shared_ptr> cancelled; std::promise promise; SequenceHandle handle{}; }; void run() noexcept; void boundary(); + void publish(std::unique_ptr& request); void emit(SchedulerEvent::Kind kind, Request const* request = nullptr, uint64_t partner = 0); void terminal( std::unique_ptr& request, SchedulerStatus status, bool release, std::exception_ptr error = {}); diff --git a/cpp/runtime/greedySchedulerBackend.cpp b/cpp/runtime/greedySchedulerBackend.cpp index bb2a3f16a..3c51f39cf 100644 --- a/cpp/runtime/greedySchedulerBackend.cpp +++ b/cpp/runtime/greedySchedulerBackend.cpp @@ -4,44 +4,101 @@ */ #include "runtime/greedySchedulerBackend.h" #include "common/checkMacros.h" -#include +#include "tokenizer/tokenizer.h" +#include #include - namespace trt_edgellm { namespace rt { -GreedySchedulerBackend::GreedySchedulerBackend(LLMRankRuntime& runtime, cudaStream_t stream, int32_t vocabularySize) +SamplingSchedulerBackend::SamplingSchedulerBackend( + LLMRankRuntime& runtime, cudaStream_t stream, int32_t vocabularySize, tokenizer::Tokenizer const* tokenizer) : mSteps(runtime, stream) , mStream(stream) , mVocabulary(vocabularySize) , mHostLogits({2, vocabularySize}, DeviceType::kCPU, nvinfer1::DataType::kFLOAT) + , mSampler(vocabularySize) { CUDA_CHECK(cudaGetDevice(&mDevice)); + if (tokenizer) + { + mPieces.reserve(vocabularySize); + for (int32_t id = 0; id < vocabularySize; ++id) + { + mPieces.push_back(tokenizer->idToPiece(id, true)); + mMaxPieceBytes = std::max(mMaxPieceBytes, mPieces.back().size()); + } + mEos = tokenizer->getEosIds(); + mPrimaryEos = tokenizer->getEosId(); + mThinkStart = tokenizer->getTokenId(""); + mThinkEnd = tokenizer->getTokenId(""); + } } -void GreedySchedulerBackend::start() +void SamplingSchedulerBackend::start() { CUDA_CHECK(cudaSetDevice(mDevice)); } -SequenceHandle GreedySchedulerBackend::acquire(uint64_t id, std::vector prompt, int32_t maxOutput) +SequenceOptions SamplingSchedulerBackend::normalizeOptions(SequenceOptions options) const { - for (auto token : prompt) + if (mPieces.empty() && (!options.stopStrings.empty() || options.enableThinking)) { - if (token < 0 || token >= mVocabulary) - { - throw std::invalid_argument("Token outside vocabulary"); - } + throw std::invalid_argument("Text policy requires tokenizer metadata"); + } + if (options.eosTokenIds.empty()) + { + options.eosTokenIds = mEos; + } + if (options.primaryEosTokenId == -1) + { + options.primaryEosTokenId = mPrimaryEos; + } + if (options.thinkingStartTokenId == -1) + { + options.thinkingStartTokenId = mThinkStart; + } + if (options.thinkingEndTokenId == -1) + { + options.thinkingEndTokenId = mThinkEnd; } - SequenceOptions options; - options.maxOutputTokens = maxOutput; - options.temperature = 0.0F; - return mSteps.acquire(id, std::move(prompt), std::move(options)); + return options; } -SequenceState const& GreedySchedulerBackend::state(SequenceHandle handle) const +SequenceHandle SamplingSchedulerBackend::acquire(uint64_t id, std::vector prompt, SequenceOptions options) +{ + auto handle = mSteps.acquire(id, std::move(prompt), options); + mPolicies[handle.slot].reset(std::move(options), mMaxPieceBytes); + return handle; +} +SequenceState const& SamplingSchedulerBackend::state(SequenceHandle handle) const { return mSteps.state(handle); } -void GreedySchedulerBackend::sample(Tensor const& logits, std::array const& handles, int32_t count) +SequencePolicy const& SamplingSchedulerBackend::policy(SequenceHandle handle) const +{ + mSteps.state(handle); + return mPolicies[handle.slot]; +} +SequenceSample SamplingSchedulerBackend::lastSample(SequenceHandle handle) const +{ + return policy(handle).last(); +} +std::string_view SamplingSchedulerBackend::text(SequenceHandle handle) const +{ + return policy(handle).text(); +} +std::vector const& SamplingSchedulerBackend::logprobs(SequenceHandle handle) const +{ + return policy(handle).logprobs(); +} +SequenceFinish SamplingSchedulerBackend::finishReason(SequenceHandle handle) const +{ + return policy(handle).finish(); +} +void SamplingSchedulerBackend::finalize(SequenceHandle handle) +{ + mSteps.state(handle); + mPolicies[handle.slot].finalize(); +} +void SamplingSchedulerBackend::sample(Tensor const& logits, std::array const& handles, int32_t count) { if (logits.getDataType() != nvinfer1::DataType::kFLOAT || logits.getShape()[0] != count || logits.getShape()[1] != mVocabulary) @@ -54,23 +111,18 @@ void GreedySchedulerBackend::sample(Tensor const& logits, std::array(); for (int32_t row = 0; row < count; ++row) { - auto const* values = data + row * mVocabulary; - int32_t best = 0; - for (int32_t token = 0; token < mVocabulary; ++token) + auto handle = handles[row]; + auto const& sequence = state(handle); + auto sample = mSampler.sample(data + row * mVocabulary, sequence.options(), sequence.randomCounter()); + mSteps.acceptToken(handle, sample.token, sample.randomCounter - sequence.randomCounter()); + mPolicies[handle.slot].accept(sample, mPieces.empty() ? std::string_view{} : mPieces[sample.token]); + if (mPolicies[handle.slot].finish() != SequenceFinish::kNone) { - if (!std::isfinite(values[token])) - { - throw std::runtime_error("Nonfinite scheduler logits"); - } - if (values[token] > values[best]) - { - best = token; - } + mSteps.finish(handle); } - mSteps.acceptToken(handles[row], best); } } -void GreedySchedulerBackend::prefill(SequenceHandle handle) +void SamplingSchedulerBackend::prefill(SequenceHandle handle) { auto const& logits = mSteps.beginPrefillChunk(handle); mSteps.completeStep(); @@ -79,13 +131,13 @@ void GreedySchedulerBackend::prefill(SequenceHandle handle) sample(logits, {handle, {}}, 1); } } -void GreedySchedulerBackend::decode(std::array const& handles, int32_t count) +void SamplingSchedulerBackend::decode(std::array const& handles, int32_t count) { auto const& logits = mSteps.beginDecode(handles, count); mSteps.completeStep(); sample(logits, handles, count); } -void GreedySchedulerBackend::release(SequenceHandle handle) +void SamplingSchedulerBackend::release(SequenceHandle handle) { mSteps.finish(handle); mSteps.release(handle); diff --git a/cpp/runtime/greedySchedulerBackend.h b/cpp/runtime/greedySchedulerBackend.h index 37c26b1f6..a3b1d82fc 100644 --- a/cpp/runtime/greedySchedulerBackend.h +++ b/cpp/runtime/greedySchedulerBackend.h @@ -5,16 +5,20 @@ #pragma once #include "runtime/continuousScheduler.h" #include "runtime/sequenceStepRuntime.h" - namespace trt_edgellm { +namespace tokenizer +{ +class Tokenizer; +} namespace rt { -//! P4-only greedy adapter. Borrowed runtime and explicit stream outlive this exclusive lease. -class GreedySchedulerBackend : public SchedulerBackend +//! Independent CPU policy over one sampling-free TensorRT runtime. All forward methods belong to the scheduler worker. +class SamplingSchedulerBackend : public SchedulerBackend { public: - GreedySchedulerBackend(LLMRankRuntime& runtime, cudaStream_t stream, int32_t vocabularySize); + SamplingSchedulerBackend(LLMRankRuntime& runtime, cudaStream_t stream, int32_t vocabularySize, + tokenizer::Tokenizer const* tokenizer = nullptr); void start() override; void invalidate() noexcept override { @@ -24,19 +28,38 @@ class GreedySchedulerBackend : public SchedulerBackend { return token >= 0 && token < mVocabulary; } - SequenceHandle acquire(uint64_t id, std::vector prompt, int32_t maxOutput) override; + int32_t vocabularySize() const override + { + return mVocabulary; + } + SequenceOptions normalizeOptions(SequenceOptions options) const override; + SequenceHandle acquire(uint64_t id, std::vector prompt, SequenceOptions options) override; SequenceState const& state(SequenceHandle handle) const override; void prefill(SequenceHandle handle) override; void decode(std::array const& handles, int32_t count) override; void release(SequenceHandle handle) override; + SequenceSample lastSample(SequenceHandle handle) const override; + std::string_view text(SequenceHandle handle) const override; + std::vector const& logprobs(SequenceHandle handle) const override; + SequenceFinish finishReason(SequenceHandle handle) const override; + void finalize(SequenceHandle handle) override; private: void sample(Tensor const& logits, std::array const& handles, int32_t count); + SequencePolicy const& policy(SequenceHandle handle) const; SequenceStepRuntime mSteps; cudaStream_t mStream; int32_t mVocabulary; int mDevice{}; Tensor mHostLogits; + SequenceSampler mSampler; + std::array mPolicies; + std::vector mPieces; + size_t mMaxPieceBytes{}; + std::vector mEos; + int32_t mPrimaryEos{-1}, mThinkStart{-1}, mThinkEnd{-1}; }; +//! P4 source compatibility name; the options-based API applies full independent policy. +using GreedySchedulerBackend = SamplingSchedulerBackend; } // namespace rt } // namespace trt_edgellm diff --git a/cpp/runtime/state/sequenceChannel.cpp b/cpp/runtime/state/sequenceChannel.cpp new file mode 100644 index 000000000..8cdae0ea5 --- /dev/null +++ b/cpp/runtime/state/sequenceChannel.cpp @@ -0,0 +1,68 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ +#include "runtime/state/sequenceChannel.h" +#include +namespace trt_edgellm +{ +namespace rt +{ +SequenceChannel::SequenceChannel(size_t records, size_t bytes) + : mRecords(records) + , mBytes(bytes) +{ + if (!records || !bytes) + { + throw std::invalid_argument("Empty stream capacity"); + } +} +bool SequenceChannel::push(SequenceSample const& sample, std::string_view text) +{ + std::lock_guard lock(mMutex); + if (mClosed || mCount == mRecords.size() || text.size() > mBytes.size() - mByteCount) + { + return false; + } + auto& entry = mRecords[(mHead + mCount) % mRecords.size()]; + entry.sample = sample; + entry.bytes = text.size(); + for (size_t i = 0; i < text.size(); ++i) + { + mBytes[(mByteHead + mByteCount + i) % mBytes.size()] = text[i]; + } + ++mCount; + mByteCount += text.size(); + mWake.notify_one(); + return true; +} +SequenceRead SequenceChannel::read(std::chrono::milliseconds timeout) +{ + std::unique_lock lock(mMutex); + mWake.wait_for(lock, timeout, [&] { return mClosed || mCount; }); + if (!mCount) + { + return {std::nullopt, mClosed}; + } + auto const& entry = mRecords[mHead]; + SequenceUpdate update; + update.sample = entry.sample; + update.text.resize(entry.bytes); + for (size_t i = 0; i < entry.bytes; ++i) + { + update.text[i] = mBytes[(mByteHead + i) % mBytes.size()]; + } + mByteHead = (mByteHead + entry.bytes) % mBytes.size(); + mByteCount -= entry.bytes; + mHead = (mHead + 1) % mRecords.size(); + --mCount; + return {std::move(update), mClosed && mCount == 0}; +} +void SequenceChannel::close() +{ + std::lock_guard lock(mMutex); + mClosed = true; + mWake.notify_all(); +} +} // namespace rt +} // namespace trt_edgellm diff --git a/cpp/runtime/state/sequenceChannel.h b/cpp/runtime/state/sequenceChannel.h new file mode 100644 index 000000000..e00353aa5 --- /dev/null +++ b/cpp/runtime/state/sequenceChannel.h @@ -0,0 +1,50 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ +#pragma once +#include "runtime/state/sequencePolicy.h" +#include +#include +#include +#include + +namespace trt_edgellm +{ +namespace rt +{ +struct SequenceUpdate +{ + SequenceSample sample; + std::string text; +}; +struct SequenceRead +{ + std::optional update; + bool closed{}; +}; +//! Single-consumer bounded stream; only readers allocate per update. Terminal notification cannot be blocked by +//! fullness. +class SequenceChannel +{ +public: + SequenceChannel(size_t records, size_t bytes); + bool push(SequenceSample const& sample, std::string_view text); + SequenceRead read(std::chrono::milliseconds timeout); + void close(); + +private: + struct Entry + { + SequenceSample sample; + size_t bytes{}; + }; + std::vector mRecords; + std::vector mBytes; + std::mutex mMutex; + std::condition_variable mWake; + size_t mHead{}, mCount{}, mByteHead{}, mByteCount{}; + bool mClosed{}; +}; +} // namespace rt +} // namespace trt_edgellm diff --git a/cpp/runtime/state/sequencePolicy.cpp b/cpp/runtime/state/sequencePolicy.cpp new file mode 100644 index 000000000..5daa37f49 --- /dev/null +++ b/cpp/runtime/state/sequencePolicy.cpp @@ -0,0 +1,278 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ +#include "runtime/state/sequencePolicy.h" +#include +#include +#include +#include + +namespace trt_edgellm +{ +namespace rt +{ +size_t validateSequenceOptions(SequenceOptions const& o, int32_t vocabulary) +{ + if (o.maxOutputTokens <= 0 || o.maxOutputTokens > 8192 || !std::isfinite(o.temperature) || o.temperature < 0 + || o.temperature > 2 || !std::isfinite(o.topP) || o.topP <= 0 || o.topP > 1 || o.topK < 0 || o.topK > vocabulary + || o.numLogprobs < 0 || o.numLogprobs > 50 || o.logitBias.size() > 1024 || o.eosTokenIds.size() > 256 + || o.stopStrings.size() > 64) + { + throw std::invalid_argument("Invalid per-sequence options"); + } + auto valid = [&](int32_t id) { return id >= 0 && id < vocabulary; }; + for (auto id : o.eosTokenIds) + { + if (!valid(id)) + { + throw std::invalid_argument("Invalid EOS ID"); + } + } + for (auto id : {o.primaryEosTokenId, o.thinkingStartTokenId, o.thinkingEndTokenId}) + { + if (id != -1 && !valid(id)) + { + throw std::invalid_argument("Invalid policy token ID"); + } + } + size_t bytes = o.eosTokenIds.size() * sizeof(int32_t) + o.logitBias.size() * 64; + for (auto const& pair : o.logitBias) + { + if (!valid(pair.first) || !std::isfinite(pair.second) || pair.second < -100 || pair.second > 100) + { + throw std::invalid_argument("Invalid logit bias"); + } + } + size_t stops = 0; + for (auto const& stop : o.stopStrings) + { + if (stop.empty() || stop.size() > 4096) + { + throw std::invalid_argument("Invalid stop string"); + } + stops += stop.size() + sizeof(std::string); + } + if (stops > 16384) + { + throw std::invalid_argument("Stop strings exceed byte limit"); + } + return bytes + stops; +} +SequenceSampler::SequenceSampler(int32_t vocabulary) + : mEntries(vocabulary) +{ + if (vocabulary <= 0) + { + throw std::invalid_argument("Empty vocabulary"); + } +} +SequenceSample SequenceSampler::sample(float const* logits, SequenceOptions const& o, uint64_t counter) +{ + for (size_t i = 0; i < mEntries.size(); ++i) + { + if (!std::isfinite(logits[i])) + { + throw std::runtime_error("Nonfinite logits"); + } + auto const bias = o.logitBias.find(static_cast(i)); + mEntries[i] = {static_cast(i), logits[i] + (bias == o.logitBias.end() ? 0.0 : bias->second), 0}; + } + auto better + = [](Entry const& a, Entry const& b) { return a.logit > b.logit || (a.logit == b.logit && a.token < b.token); }; + // std::sort uses stack storage; stable_sort may allocate on every token. + std::sort(mEntries.begin(), mEntries.end(), better); + double const maximum = mEntries.front().logit; + double normalizer = 0; + for (auto const& entry : mEntries) + { + normalizer += std::exp(entry.logit - maximum); + } + double const logNormalizer = maximum + std::log(normalizer); + SequenceSample out; + out.randomCounter = counter; + out.topCount = std::min(o.numLogprobs, static_cast(mEntries.size())); + for (int32_t i = 0; i < out.topCount; ++i) + { + out.top[i] = {mEntries[i].token, mEntries[i].logit - logNormalizer}; + } + size_t selected = 0; + bool const greedy = o.temperature <= 1e-3F || o.topK == 1 + || (o.topK <= 1 && o.topP >= 1.0F - 1e-6F && std::fabs(o.temperature - 1.0F) <= 1e-3F); + if (!greedy) + { + if (counter == std::numeric_limits::max()) + { + throw std::overflow_error("RNG counter exhausted"); + } + size_t const k = o.topK > 0 ? static_cast(o.topK) : mEntries.size(); + double sum = 0; + for (size_t i = 0; i < k; ++i) + { + mEntries[i].weight = std::exp((mEntries[i].logit - maximum) / o.temperature); + sum += mEntries[i].weight; + } + size_t n = 0; + double retained = 0; + do + { + retained += mEntries[n++].weight; + } while (n < k && retained < o.topP * sum); + uint64_t bits = o.seed + 0x9e3779b97f4a7c15ULL * (counter + 1); + bits = (bits ^ (bits >> 30)) * 0xbf58476d1ce4e5b9ULL; + bits = (bits ^ (bits >> 27)) * 0x94d049bb133111ebULL; + bits ^= bits >> 31; + double const draw = static_cast(bits >> 11) * 0x1.0p-53 * retained; + double cumulative = mEntries[0].weight; + while (selected + 1 < n && draw >= cumulative) + { + cumulative += mEntries[++selected].weight; + } + out.randomCounter = counter + 1; + } + out.token = mEntries[selected].token; + out.logprob = mEntries[selected].logit - logNormalizer; + return out; +} +void SequencePolicy::reset(SequenceOptions options, size_t maxPieceBytes) +{ + mOptions = std::move(options); + mLast = {}; + mLogprobs.clear(); + if (mOptions.numLogprobs) + { + mLogprobs.reserve(mOptions.maxOutputTokens); + } + mText.clear(); + mText.reserve(static_cast(mOptions.maxOutputTokens) * std::max(size_t{1}, maxPieceBytes) * 3 + 3); + mUtf8Size = 0; + mUtf8Expected = 0; + mSafeBytes = 0; + mMaxStopBytes = 0; + mGenerated = 0; + mThinkingDone = !mOptions.enableThinking; + mFinish = SequenceFinish::kNone; + for (auto const& stop : mOptions.stopStrings) + { + mMaxStopBytes = std::max(mMaxStopBytes, stop.size()); + } +} +void SequencePolicy::appendUtf8(std::string_view piece) +{ + for (unsigned char byte : piece) + { + if (mUtf8Size) + { + unsigned char const lead = static_cast(mUtf8[0]); + bool const continuation = byte >= 0x80 && byte <= 0xbf + && !(mUtf8Size == 1 + && ((lead == 0xe0 && byte < 0xa0) || (lead == 0xed && byte >= 0xa0) || (lead == 0xf0 && byte < 0x90) + || (lead == 0xf4 && byte >= 0x90))); + if (continuation) + { + mUtf8[mUtf8Size++] = static_cast(byte); + if (mUtf8Size == mUtf8Expected) + { + mText.append(mUtf8.data(), mUtf8Size); + mUtf8Size = 0; + } + continue; + } + mText.append("\xef\xbf\xbd"); + mUtf8Size = 0; + } + if (byte < 0x80) + { + mText.push_back(static_cast(byte)); + } + else if (byte >= 0xc2 && byte <= 0xf4) + { + mUtf8[0] = static_cast(byte); + mUtf8Size = 1; + mUtf8Expected = byte < 0xe0 ? 2 : (byte < 0xf0 ? 3 : 4); + } + else + { + mText.append("\xef\xbf\xbd"); + } + } +} +void SequencePolicy::accept(SequenceSample const& sample, std::string_view piece) +{ + if (mFinish != SequenceFinish::kNone) + { + throw std::logic_error("Sampling finished policy"); + } + ++mGenerated; + if (!mThinkingDone + && (sample.token == mOptions.thinkingEndTokenId + || (mGenerated == 1 && sample.token != mOptions.thinkingStartTokenId))) + { + mThinkingDone = true; + } + mLast = sample; + mLast.thinking = !mThinkingDone; + if (mOptions.numLogprobs) + { + mLogprobs.push_back(mLast); + } + bool const eos = !mOptions.ignoreEos + && (sample.token == mOptions.primaryEosTokenId + || (std::find(mOptions.eosTokenIds.begin(), mOptions.eosTokenIds.end(), sample.token) + != mOptions.eosTokenIds.end() + && (!mOptions.enableThinking || mThinkingDone))); + if (!eos) + { + appendUtf8(piece); + } + if (eos) + { + mFinish = SequenceFinish::kEos; + } + else if (mGenerated >= mOptions.maxOutputTokens) + { + mFinish = SequenceFinish::kLength; + } + if (mFinish != SequenceFinish::kNone && mUtf8Size) + { + mText.append("\xef\xbf\xbd"); + mUtf8Size = 0; + } + size_t match = std::string::npos; + for (auto const& stop : mOptions.stopStrings) + { + match = std::min(match, mText.find(stop, mSafeBytes)); + } + if (match != std::string::npos) + { + mText.resize(match); + mFinish = SequenceFinish::kStop; + mUtf8Size = 0; + } + if (mFinish != SequenceFinish::kNone) + { + finalize(); + return; + } + mSafeBytes = mText.size() > mMaxStopBytes ? mText.size() - mMaxStopBytes : 0; + if (!mMaxStopBytes) + { + mSafeBytes = mText.size(); + } + // Never split a UTF-8 code point when holding the stop look-behind window. + while (mSafeBytes < mText.size() && (static_cast(mText[mSafeBytes]) & 0xc0) == 0x80) + { + --mSafeBytes; + } +} +void SequencePolicy::finalize() +{ + if (mUtf8Size) + { + mText.append("\xef\xbf\xbd"); + mUtf8Size = 0; + } + mSafeBytes = mText.size(); +} +} // namespace rt +} // namespace trt_edgellm diff --git a/cpp/runtime/state/sequencePolicy.h b/cpp/runtime/state/sequencePolicy.h new file mode 100644 index 000000000..1c4d1fee9 --- /dev/null +++ b/cpp/runtime/state/sequencePolicy.h @@ -0,0 +1,94 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ +#pragma once +#include "runtime/state/sequenceSlots.h" +#include +#include + +namespace trt_edgellm +{ +namespace rt +{ +enum class SequenceFinish +{ + kNone, + kLength, + kEos, + kStop +}; +struct TokenLogprob +{ + int32_t token{-1}; + double logprob{}; +}; +struct SequenceSample +{ + int32_t token{-1}; + double logprob{}; + std::array top{}; + int32_t topCount{}; + uint64_t randomCounter{}; + bool thinking{}; +}; +//! Allocation-free per-token sampler; scratch is shared by the single worker, never by requests' RNG state. +class SequenceSampler +{ +public: + explicit SequenceSampler(int32_t vocabulary); + SequenceSample sample(float const* logits, SequenceOptions const& options, uint64_t counter); + +private: + struct Entry + { + int32_t token{}; + double logit{}; + double weight{}; + }; + std::vector mEntries; +}; +//! Validate all supported options before a request is enqueued. Returns bounded dynamic metadata bytes. +size_t validateSequenceOptions(SequenceOptions const& options, int32_t vocabulary); + +//! Request-owned stopping/text state. Raw BPE pieces are sanitized incrementally before stop matching. +class SequencePolicy +{ +public: + void reset(SequenceOptions options, size_t maxPieceBytes); + void accept(SequenceSample const& sample, std::string_view piece); + void finalize(); + SequenceFinish finish() const + { + return mFinish; + } + std::string_view text() const + { + return {mText.data(), mSafeBytes}; + } + SequenceSample const& last() const + { + return mLast; + } + std::vector const& logprobs() const + { + return mLogprobs; + } + +private: + void appendUtf8(std::string_view piece); + SequenceOptions mOptions; + SequenceSample mLast; + std::vector mLogprobs; + std::string mText; + std::array mUtf8{}; + int32_t mUtf8Size{}; + int32_t mUtf8Expected{}; + size_t mSafeBytes{}; + size_t mMaxStopBytes{}; + int32_t mGenerated{}; + bool mThinkingDone{}; + SequenceFinish mFinish{SequenceFinish::kNone}; +}; +} // namespace rt +} // namespace trt_edgellm diff --git a/cpp/runtime/state/sequenceSlots.h b/cpp/runtime/state/sequenceSlots.h index aeb250477..01e9e43d5 100644 --- a/cpp/runtime/state/sequenceSlots.h +++ b/cpp/runtime/state/sequenceSlots.h @@ -31,6 +31,10 @@ struct SequenceOptions uint64_t seed{42}; int32_t numLogprobs{}; bool enableThinking{}; + bool ignoreEos{}; + int32_t primaryEosTokenId{-1}; + int32_t thinkingStartTokenId{-1}; + int32_t thinkingEndTokenId{-1}; std::vector eosTokenIds; std::vector stopStrings; std::unordered_map logitBias; diff --git a/examples/llm/continuousBatchingProbe.cpp b/examples/llm/continuousBatchingProbe.cpp index 8bb02a8de..2f18b3470 100644 --- a/examples/llm/continuousBatchingProbe.cpp +++ b/examples/llm/continuousBatchingProbe.cpp @@ -977,6 +977,241 @@ bool runSchedulerTests(LLMRankRuntime& runtime, tokenizer::Tokenizer& tokenizer, << " b_continued=" << bContinued << " paired_decode=" << pair << std::endl; return passed; } +bool runPolicyTests(LLMRankRuntime& runtime, tokenizer::Tokenizer& tokenizer, cudaStream_t stream) +{ + ContinuousBatchingProbe observer(runtime, stream); + auto const words = tokenizer.encode("A fox crosses a river. Explain the colors and count to five. "); + auto prompt = [&](int32_t length) { + std::vector tokens; + tokens.reserve(length); + for (int32_t i = 0; i < length; ++i) + { + tokens.push_back(words.at(i % words.size())); + } + return tokens; + }; + auto const pa = prompt(65), pb = prompt(513), pc = prompt(129); + SchedulerRequestOptions oa, ob, oc; + oa.generation.maxOutputTokens = 6; + oa.generation.temperature = 0.7F; + oa.generation.topK = 20; + oa.generation.topP = 0.9F; + oa.generation.seed = 123; + oa.generation.numLogprobs = 5; + oa.generation.ignoreEos = true; + ob = oa; + ob.generation.maxOutputTokens = 12; + ob.generation.temperature = 1.2F; + ob.generation.topK = 40; + ob.generation.seed = 456; + ob.generation.numLogprobs = 0; + oc = oa; + oc.generation.maxOutputTokens = 12; + oc.generation.temperature = 0; + oc.generation.topK = 1; + oc.generation.numLogprobs = 2; + std::atomic armed{false}; + bool submitted = false; + ContinuousScheduler* owner = nullptr; + SchedulerTicket b, c; + bool paired = false; + std::atomic cancelMode{false}; + bool cancelSubmitted = false; + SchedulerTicket nativeB, nativeC; + auto get = [&](SchedulerTicket const& ticket) { + require( + ticket.result().wait_for(std::chrono::seconds(60)) == std::future_status::ready, "Policy ticket stranded"); + return ticket.result().get(); + }; + { + ContinuousScheduler scheduler( + std::make_unique(runtime, stream, observer.vocabularySize(), &tokenizer), 8, + 262144, [&](SchedulerEvent const& e) { + if (cancelMode.load()) + { + if (!cancelSubmitted && e.kind == SchedulerEvent::Kind::kDecode) + { + cancelSubmitted = true; + nativeB = owner->submit(pb, ob); + nativeC = owner->submit(pc, oc); + nativeC.cancel(); + } + if (cancelSubmitted && e.request == nativeB.id() && e.kind == SchedulerEvent::Kind::kPrefill) + { + nativeB.cancel(); + } + } + if (armed.load() && e.kind == SchedulerEvent::Kind::kDecode) + { + if (!submitted) + { + submitted = true; + b = owner->submit(pb, ob); + c = owner->submit(pc, oc); + } + if (e.partner) + { + paired = true; + } + } + }); + owner = &scheduler; + auto const ra = get(scheduler.submit(pa, oa)), rb = get(scheduler.submit(pb, ob)), + rc = get(scheduler.submit(pc, oc)); + armed.store(true); + auto const a = scheduler.submit(pa, oa); + auto const aa = get(a), ab = get(b), ac = get(c); + armed.store(false); + require(aa.status == SchedulerStatus::kCompleted && ab.status == SchedulerStatus::kCompleted + && ac.status == SchedulerStatus::kCompleted, + "Mixed requests failed"); + require(aa.tokens == ra.tokens && ab.tokens == rb.tokens && ac.tokens == rc.tokens, + "Seeded serial/concurrent mismatch"); + require(aa.text == ra.text && ab.text == rb.text && ac.text == rc.text, "Independent text mismatch"); + require(aa.tokens.size() == 6 && ab.tokens.size() == 12 && ac.tokens.size() == 12 && paired, + "Mixed batch/limits missing"); + require( + aa.randomCounter == 6 && ab.randomCounter == 12 && ac.randomCounter == 0, "RNG counters crossed requests"); + require(aa.logprobs.size() == 6 && ab.logprobs.empty() && ac.logprobs.size() == 12 + && aa.logprobs.front().topCount == 5 && ac.logprobs.front().topCount == 2, + "Logprob history missing"); + for (size_t i = 0; i < aa.logprobs.size(); ++i) + { + require(aa.logprobs[i].logprob == ra.logprobs[i].logprob, "Logprob baseline mismatch"); + } + std::cout << "P5_MIXED_GATE passed=1 outputs=6,12,12 rng=6,12,0 paired=1 exact_seeded=1" << std::endl; + auto stopOptions = oc; + stopOptions.generation.ignoreEos = false; + stopOptions.generation.numLogprobs = 50; + int32_t const forced = tokenizer.encode("Z").at(0); + auto const piece = tokenizer.idToPiece(forced, true); + require(!piece.empty(), "Empty forced token piece"); + stopOptions.generation.logitBias[forced] = 100; + stopOptions.generation.stopStrings = {piece + piece}; + stopOptions.streamRecords = 8; + auto stopTicket = scheduler.submit(pa, stopOptions); + auto stopped = get(stopTicket); + require(stopped.status == SchedulerStatus::kCompleted && stopped.finish == SequenceFinish::kStop + && stopped.tokens.size() == 2 && stopped.text.empty(), + "Cross-token stop failed"); + std::string streamed; + while (true) + { + auto next = stopTicket.read(std::chrono::milliseconds(0)); + if (next.update) + { + streamed += next.update->text; + } + if (next.closed) + { + break; + } + } + require(streamed.empty(), "Stop prefix leaked to stream"); + auto eosOptions = oc; + eosOptions.generation.ignoreEos = false; + eosOptions.generation.logitBias[tokenizer.getEosId()] = 100; + auto eos = get(scheduler.submit(pa, eosOptions)); + require(eos.finish == SequenceFinish::kEos && eos.tokens.size() == 1 && eos.text.empty(), "Primary EOS failed"); + auto slowOptions = oa; + slowOptions.streamRecords = 1; + slowOptions.streamBytes = 32; + auto slow = scheduler.submit(pa, slowOptions); + auto peer = scheduler.submit(pb, ob); + require(get(slow).status == SchedulerStatus::kSlowConsumer, "Slow consumer not terminated"); + require(get(peer).tokens == rb.tokens, "Slow consumer corrupted partner"); + auto deadlineOptions = oa; + deadlineOptions.queueDeadline = std::chrono::steady_clock::now() - std::chrono::milliseconds(1); + require( + get(scheduler.submit(pa, deadlineOptions)).status == SchedulerStatus::kDeadline, "Queue deadline failed"); + auto streamOptions = oa; + streamOptions.streamRecords = 32; + auto streamTicket = scheduler.submit(pa, streamOptions); + auto streamResult = get(streamTicket); + std::string streamText; + size_t recordCount = 0; + while (true) + { + auto next = streamTicket.read(std::chrono::milliseconds(0)); + if (next.update) + { + streamText += next.update->text; + if (next.update->sample.token >= 0) + { + ++recordCount; + } + } + if (next.closed) + { + break; + } + } + require(streamText == streamResult.text && recordCount == 6, "Streaming result differs from final result"); + cancelMode.store(true); + auto cancellationLeader = scheduler.submit(pa, oa); + auto leaderResult = get(cancellationLeader); + require(leaderResult.tokens == ra.tokens, "Prefill cancellation changed partner output"); + require( + get(nativeB).status == SchedulerStatus::kCancelled && get(nativeC).status == SchedulerStatus::kCancelled, + "Native active/queued cancellation failed"); + cancelMode.store(false); + nativeB.cancel(); + require(get(scheduler.submit(pc, oc)).tokens == rc.tokens, "Stale ticket affected reused slot"); + auto activeDeadline = oa; + activeDeadline.deadline = std::chrono::steady_clock::now() + std::chrono::milliseconds(50); + require(get(scheduler.submit(prompt(2049), activeDeadline)).status == SchedulerStatus::kDeadline, + "Native active deadline failed"); + require(get(scheduler.submit(pa, oa)).tokens == ra.tokens, "Deadline cleanup changed later output"); + std::cout + << "P5_CANCEL_GATE passed=1 prefill_cancel=1 queued_cancel=1 stale_ticket=1 active_deadline=1 peer_exact=1" + << std::endl; + scheduler.close(); + require(scheduler.healthy(), "Policy tests poisoned healthy runtime"); + std::cout << "P5_STOP_STREAM_GATE passed=1 stop_tokens=2 eos_tokens=1 logprob_top=50 slow_peer_exact=1" + << std::endl; + } + // Inject a reported corrupting failure after real GPU work without deliberately damaging the device. + class InjectedFailure : public SamplingSchedulerBackend + { + public: + using SamplingSchedulerBackend::SamplingSchedulerBackend; + void prefill(SequenceHandle handle) override + { + SamplingSchedulerBackend::prefill(handle); + throw std::runtime_error("injected post-forward CUDA failure"); + } + }; + { + ContinuousScheduler failed( + std::make_unique(runtime, stream, observer.vocabularySize(), &tokenizer)); + auto ticket = failed.submit(pa, oa); + require(get(ticket).status == SchedulerStatus::kFailed, "Fault left ticket live"); + failed.close(); + require(!failed.healthy(), "Fault failed to mark scheduler unready"); + bool rejected = false; + try + { + failed.submit(pa, oa); + } + catch (std::exception const&) + { + rejected = true; + } + require(rejected, "Failed scheduler admitted request"); + } + bool poisoned = false; + try + { + SequenceStepRuntime forbidden(runtime, stream); + } + catch (std::exception const&) + { + poisoned = true; + } + require(poisoned, "Failed sampling released parent lease"); + std::cout << "P5_FAULT_GATE passed=1 injected_post_forward=1 parent_lease_poisoned=1" << std::endl; + return true; +} } // namespace rt } // namespace trt_edgellm @@ -984,11 +1219,12 @@ int main(int argc, char** argv) { if (argc != 3 && (argc != 4 - || (std::string(argv[3]) != "--scheduler" && std::string(argv[3]) != "--steps" + || (std::string(argv[3]) != "--policies" && std::string(argv[3]) != "--scheduler" + && std::string(argv[3]) != "--steps" && (std::string(argv[3]) != "--chunks" && std::string(argv[3]) != "--chunks-extra")))) { std::cerr << "Usage: continuous_batching_probe ENGINE_DIR CHECKPOINT_DIR " - "[--steps|--chunks|--chunks-extra|--scheduler]\n"; + "[--steps|--chunks|--chunks-extra|--scheduler|--policies]\n"; return 2; } cudaStream_t stream{}; @@ -1003,7 +1239,9 @@ int main(int argc, char** argv) trt_edgellm::rt::require(tokenizer.loadFromHF(argv[1], false), "Tokenizer loading failed"); trt_edgellm::rt::LLMRankRuntime runtime(argv[1], "", {}, std::nullopt, stream, trt_edgellm::rt::ParallelMapping{}, tokenizer, trt_edgellm::rt::ContextCacheConfig{}, argv[2], ""); - passed = argc == 4 && std::string(argv[3]) == "--scheduler" + passed = argc == 4 && std::string(argv[3]) == "--policies" + ? trt_edgellm::rt::runPolicyTests(runtime, tokenizer, stream) + : argc == 4 && std::string(argv[3]) == "--scheduler" ? trt_edgellm::rt::runSchedulerTests(runtime, tokenizer, stream) : argc == 4 ? (std::string(argv[3]).find("--chunks") == 0 ? trt_edgellm::rt::runChunkTests( diff --git a/unittests/cpp/runtime/continuousSchedulerTest.cpp b/unittests/cpp/runtime/continuousSchedulerTest.cpp index aac09e11c..93927fa14 100644 --- a/unittests/cpp/runtime/continuousSchedulerTest.cpp +++ b/unittests/cpp/runtime/continuousSchedulerTest.cpp @@ -24,10 +24,8 @@ class FakeBackend : public SchedulerBackend gate.wait(); } } - SequenceHandle acquire(uint64_t id, std::vector prompt, int32_t maxOutput) override + SequenceHandle acquire(uint64_t id, std::vector prompt, SequenceOptions options) override { - SequenceOptions options; - options.maxOutputTokens = maxOutput; return slots.acquire(id, std::move(prompt), options); } SequenceState const& state(SequenceHandle h) const override @@ -290,3 +288,108 @@ TEST(ContinuousScheduler, ConcurrentProducersAndRepeatedReuse) scheduler.close(); EXPECT_TRUE(scheduler.healthy()); } + +TEST(ContinuousScheduler, SlowConsumerDoesNotBlockPartner) +{ + std::promise gate; + auto backend = std::make_unique(); + backend->gate = gate.get_future().share(); + ContinuousScheduler scheduler(std::move(backend)); + SchedulerRequestOptions options; + options.generation.maxOutputTokens = 10; + options.streamRecords = 1; + options.streamBytes = 8; + auto slow = scheduler.submit({1}, options); + auto peer = scheduler.submit({1}, 8); + gate.set_value(); + EXPECT_EQ(result(slow).status, SchedulerStatus::kSlowConsumer); + EXPECT_EQ(result(peer).tokens.size(), 8U); + auto read = slow.read(0ms); + ASSERT_TRUE(read.update); + EXPECT_TRUE(read.closed); + EXPECT_TRUE(scheduler.healthy()); +} +TEST(ContinuousScheduler, QueuedCancellationWhileBothSlotsOccupied) +{ + ContinuousScheduler* owner = nullptr; + SchedulerTicket queued; + std::atomic cancelledReady{false}; + bool submitted = false; + ContinuousScheduler scheduler(std::make_unique(), 8, 65536, [&](auto const& e) { + if (e.kind == SchedulerEvent::Kind::kDecode && !submitted) + { + submitted = true; + queued = owner->submit({1}, 3); + queued.cancel(); + } + else if (submitted && e.kind == SchedulerEvent::Kind::kDecode) + { + cancelledReady = queued.result().wait_for(0ms) == std::future_status::ready; + } + }); + owner = &scheduler; + auto a = scheduler.submit({1}, 20); + auto b = scheduler.submit({1}, 20); + result(a); + result(b); + scheduler.close(); + EXPECT_EQ(result(queued).status, SchedulerStatus::kCancelled); + EXPECT_TRUE(cancelledReady); +} +TEST(ContinuousScheduler, QueueAndActiveDeadlines) +{ + std::promise gate; + auto backend = std::make_unique(); + backend->gate = gate.get_future().share(); + SchedulerRequestOptions options; + options.generation.maxOutputTokens = 5; + options.queueDeadline = std::chrono::steady_clock::now() - 1ms; + ContinuousScheduler scheduler(std::move(backend)); + auto expired = scheduler.submit({1}, options); + gate.set_value(); + EXPECT_EQ(result(expired).status, SchedulerStatus::kDeadline); + scheduler.close(); + SchedulerRequestOptions active; + active.generation.maxOutputTokens = 100; + active.deadline = std::chrono::steady_clock::now() + 20ms; + bool delayed = false; + ContinuousScheduler other(std::make_unique(), 8, 65536, [&](auto const& e) { + if (e.kind == SchedulerEvent::Kind::kPrefill && !delayed) + { + delayed = true; + std::this_thread::sleep_for(30ms); + } + }); + auto a = other.submit({1}, active); + EXPECT_EQ(result(a).status, SchedulerStatus::kDeadline); + EXPECT_EQ(result(other.submit({1}, 3)).status, SchedulerStatus::kCompleted); +} +TEST(ContinuousScheduler, DecodeCancellationAndStartupFailure) +{ + SchedulerTicket a; + std::promise gate; + auto backend = std::make_unique(); + backend->gate = gate.get_future().share(); + ContinuousScheduler scheduler(std::move(backend), 8, 65536, [&](auto const& e) { + if (e.kind == SchedulerEvent::Kind::kDecode && e.request == a.id()) + { + a.cancel(); + } + }); + a = scheduler.submit({1}, 10); + auto b = scheduler.submit({1}, 12); + gate.set_value(); + EXPECT_EQ(result(a).status, SchedulerStatus::kCancelled); + EXPECT_EQ(result(b).tokens.size(), 12U); + class BrokenStart : public FakeBackend + { + void start() override + { + throw std::runtime_error("startup failure"); + } + }; + ContinuousScheduler broken(std::make_unique()); + std::this_thread::sleep_for(10ms); + EXPECT_FALSE(broken.healthy()); + EXPECT_THROW(broken.submit({1}, 1), std::runtime_error); +} diff --git a/unittests/cpp/runtime/state/sequencePolicyTest.cpp b/unittests/cpp/runtime/state/sequencePolicyTest.cpp new file mode 100644 index 000000000..0a4fd8817 --- /dev/null +++ b/unittests/cpp/runtime/state/sequencePolicyTest.cpp @@ -0,0 +1,204 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ +#include "runtime/state/sequencePolicy.h" +#include "runtime/state/sequenceChannel.h" +#include +#include +#include +using namespace trt_edgellm::rt; +namespace +{ +SequenceSample token(int id) +{ + SequenceSample s; + s.token = id; + return s; +} +} // namespace +TEST(SequencePolicy, BiasedUnscaledLogprobsAndTopK) +{ + SequenceSampler sampler(4); + float values[]{0, 1, 2, 3}; + SequenceOptions o; + o.temperature = 0; + o.numLogprobs = 4; + o.logitBias[0] = 5; + auto s = sampler.sample(values, o, 0); + EXPECT_EQ(s.token, 0); + EXPECT_EQ(s.randomCounter, 0U); + EXPECT_EQ(s.top[0].token, 0); + EXPECT_NEAR(s.logprob, 5 - std::log(std::exp(5) + std::exp(1) + std::exp(2) + std::exp(3)), 1e-12); + o.temperature = 0.7; + o.topK = 2; + o.topP = 1; + for (uint64_t i = 0; i < 100; ++i) + { + auto sampled = sampler.sample(values, o, i); + EXPECT_TRUE(sampled.token == 0 || sampled.token == 3); + } +} +TEST(SequencePolicy, RngIndependentOfOtherRowsAndSeed) +{ + SequenceSampler sampler(4); + float values[]{0, 0, 0, 0}; + SequenceOptions a, b; + a.temperature = 0.7; + a.topK = 4; + a.seed = 7; + b = a; + b.seed = 9; + bool different = false; + for (uint64_t counter = 0; counter < 100; ++counter) + { + auto reference = sampler.sample(values, a, counter); + auto other = sampler.sample(values, b, counter); + auto replay = sampler.sample(values, a, counter); + EXPECT_EQ(reference.token, replay.token); + EXPECT_EQ(reference.randomCounter, counter + 1); + different = different || reference.token != other.token; + } + EXPECT_TRUE(different); + EXPECT_THROW(sampler.sample(values, a, std::numeric_limits::max()), std::overflow_error); +} +TEST(SequencePolicy, NucleusRetainsMinimalProbabilityPrefix) +{ + SequenceSampler sampler(4); + float values[]{3, 2, 1, 0}; + SequenceOptions o; + o.temperature = 0.8; + o.topK = 4; + o.topP = 0.1; + for (uint64_t i = 0; i < 100; ++i) + { + EXPECT_EQ(sampler.sample(values, o, i).token, 0); + } + o.topP = 1; + o.topK = 1; + EXPECT_EQ(sampler.sample(values, o, 0).randomCounter, 0U); +} +TEST(SequencePolicy, StopAcrossTokensDoesNotLeakPrefix) +{ + SequencePolicy policy; + SequenceOptions o; + o.stopStrings = {"STOP", "TOP"}; + o.maxOutputTokens = 10; + policy.reset(o, 16); + policy.accept(token(1), "hello S"); + std::string emitted(policy.text()); + policy.accept(token(2), "TO"); + emitted = std::string(policy.text()); + EXPECT_EQ(emitted.find('S'), std::string::npos); + policy.accept(token(3), "P ignored"); + EXPECT_EQ(policy.finish(), SequenceFinish::kStop); + EXPECT_EQ(policy.text(), "hello "); + EXPECT_THROW(policy.accept(token(4), "x"), std::logic_error); +} +TEST(SequencePolicy, UnicodeBoundariesAndReset) +{ + SequencePolicy policy; + SequenceOptions o; + o.maxOutputTokens = 4; + policy.reset(o, 8); + policy.accept(token(1), "\xe2"); + EXPECT_TRUE(policy.text().empty()); + policy.accept(token(2), "\x82\xac"); + EXPECT_EQ(policy.text(), "\xe2\x82\xac"); + policy.accept(token(3), "\xff"); + EXPECT_EQ(policy.text(), "\xe2\x82\xac\xef\xbf\xbd"); + policy.accept(token(4), "\xf0"); + EXPECT_EQ(policy.finish(), SequenceFinish::kLength); + EXPECT_EQ(policy.text(), "\xe2\x82\xac\xef\xbf\xbd\xef\xbf\xbd"); + o.stopStrings = {"\xe2\x82\xac"}; + policy.reset(o, 8); + policy.accept(token(1), "a\xe2"); + EXPECT_TRUE(policy.text().empty()); + policy.accept(token(2), "\x82\xac"); + EXPECT_EQ(policy.text(), "a"); + EXPECT_EQ(policy.finish(), SequenceFinish::kStop); +} +TEST(SequencePolicy, IndependentThinkingAndEos) +{ + SequenceOptions o; + o.enableThinking = true; + o.thinkingStartTokenId = 1; + o.thinkingEndTokenId = 2; + o.primaryEosTokenId = 3; + o.eosTokenIds = {3, 4}; + o.maxOutputTokens = 10; + SequencePolicy a, b; + a.reset(o, 10); + o.enableThinking = false; + b.reset(o, 10); + a.accept(token(1), ""); + EXPECT_TRUE(a.last().thinking); + a.accept(token(4), ""); + EXPECT_EQ(a.finish(), SequenceFinish::kNone); + b.accept(token(4), ""); + EXPECT_EQ(b.finish(), SequenceFinish::kEos); + a.accept(token(2), ""); + EXPECT_FALSE(a.last().thinking); + a.accept(token(4), ""); + EXPECT_EQ(a.finish(), SequenceFinish::kEos); + o.enableThinking = true; + a.reset(o, 10); + a.accept(token(1), ""); + a.accept(token(3), ""); + EXPECT_EQ(a.finish(), SequenceFinish::kEos); + a.reset(o, 10); + a.accept(token(9), "answer"); + EXPECT_FALSE(a.last().thinking); +} +TEST(SequencePolicy, ValidatesMetadataBeforeAdmission) +{ + SequenceOptions o; + o.numLogprobs = 51; + EXPECT_THROW(validateSequenceOptions(o, 8), std::invalid_argument); + o.numLogprobs = 0; + o.logitBias[1] = std::numeric_limits::quiet_NaN(); + EXPECT_THROW(validateSequenceOptions(o, 8), std::invalid_argument); + o.logitBias.clear(); + o.stopStrings = {""}; + EXPECT_THROW(validateSequenceOptions(o, 8), std::invalid_argument); + o.stopStrings.clear(); + o.topK = 9; + EXPECT_THROW(validateSequenceOptions(o, 8), std::invalid_argument); + SequenceSampler sampler(4); + float values[]{0, 1, 2, std::numeric_limits::infinity()}; + EXPECT_THROW(sampler.sample(values, SequenceOptions{}, 0), std::runtime_error); +} +TEST(SequenceChannel, WraparoundByteLimitAndTerminalOnFull) +{ + SequenceChannel channel(2, 5); + EXPECT_TRUE(channel.push(token(1), "abc")); + EXPECT_FALSE(channel.push(token(2), "def")); + auto first = channel.read(std::chrono::milliseconds(0)); + ASSERT_TRUE(first.update); + EXPECT_EQ(first.update->text, "abc"); + EXPECT_TRUE(channel.push(token(2), "def")); + EXPECT_TRUE(channel.push(token(3), "gh")); + EXPECT_FALSE(channel.push(token(4), "")); + channel.close(); + auto second = channel.read(std::chrono::milliseconds(0)); + ASSERT_TRUE(second.update); + EXPECT_EQ(second.update->text, "def"); + auto third = channel.read(std::chrono::milliseconds(0)); + ASSERT_TRUE(third.update); + EXPECT_EQ(third.update->text, "gh"); + EXPECT_TRUE(third.closed); + EXPECT_TRUE(channel.read(std::chrono::milliseconds(0)).closed); + EXPECT_FALSE(channel.push(token(5), "")); +} + +TEST(SequencePolicy, FinalUtf8ReplacementParticipatesInStopMatch) +{ + SequencePolicy policy; + SequenceOptions o; + o.maxOutputTokens = 1; + o.stopStrings = {"\xef\xbf\xbd"}; + policy.reset(o, 4); + policy.accept(token(1), "\xe2"); + EXPECT_EQ(policy.finish(), SequenceFinish::kStop); + EXPECT_TRUE(policy.text().empty()); +} From a90574f894d07d6e1fbc7eda6dffa43a0dcbb37d Mon Sep 17 00:00:00 2001 From: ajmalrasi Date: Tue, 15 Sep 2026 15:51:38 +0530 Subject: [PATCH 5/7] feat: bridge continuous scheduler into HTTP runtime Signed-off-by: ajmalrasi --- cpp/runtime/continuousScheduler.cpp | 14 ++ cpp/runtime/continuousScheduler.h | 4 +- cpp/runtime/llmInferenceRuntime.cpp | 31 +++++ cpp/runtime/llmInferenceRuntime.h | 14 ++ cpp/runtime/llmRankRuntime.h | 12 ++ cpp/runtime/multiDevice/runtimeCoordinator.h | 5 +- experimental/pybind/edgellm_pybind.cpp | 135 +++++++++++++++++++ experimental/server/runtime/engine.py | 126 ++++++++++++++++- experimental/server/runtime/engine_client.py | 48 ++++++- 9 files changed, 385 insertions(+), 4 deletions(-) diff --git a/cpp/runtime/continuousScheduler.cpp b/cpp/runtime/continuousScheduler.cpp index d263ef133..47d527a5f 100644 --- a/cpp/runtime/continuousScheduler.cpp +++ b/cpp/runtime/continuousScheduler.cpp @@ -3,6 +3,7 @@ * SPDX-License-Identifier: Apache-2.0 */ #include "runtime/continuousScheduler.h" +#include #include #include #include @@ -29,6 +30,19 @@ ContinuousScheduler::~ContinuousScheduler() close(); } +size_t ContinuousScheduler::queuedCount() const +{ + std::lock_guard lock(mMutex); + return mQueue.size(); +} + +size_t ContinuousScheduler::residentCount() const +{ + std::lock_guard lock(mMutex); + return static_cast(std::count_if( + mActive.begin(), mActive.end(), [](auto const& request) { return request != nullptr; })); +} + SchedulerTicket ContinuousScheduler::submit(std::vector const& prompt, int32_t maxOutput) { SchedulerRequestOptions options; diff --git a/cpp/runtime/continuousScheduler.h b/cpp/runtime/continuousScheduler.h index 0ffba447a..5d38050c2 100644 --- a/cpp/runtime/continuousScheduler.h +++ b/cpp/runtime/continuousScheduler.h @@ -162,6 +162,8 @@ class ContinuousScheduler SchedulerTicket submit(std::vector const& prompt, int32_t maxOutput); SchedulerTicket submit(std::vector const& prompt, SchedulerRequestOptions const& options); void close(); + size_t queuedCount() const; + size_t residentCount() const; bool healthy() const { return mHealthy.load(); @@ -192,7 +194,7 @@ class ContinuousScheduler size_t const mMaxQueued; size_t const mMaxQueuedBytes; Observer mObserver; - std::mutex mMutex; + mutable std::mutex mMutex; std::mutex mCloseMutex; std::condition_variable mWake; std::deque> mQueue; diff --git a/cpp/runtime/llmInferenceRuntime.cpp b/cpp/runtime/llmInferenceRuntime.cpp index 247526a5f..3f658d697 100644 --- a/cpp/runtime/llmInferenceRuntime.cpp +++ b/cpp/runtime/llmInferenceRuntime.cpp @@ -20,6 +20,7 @@ #include "common/checkMacros.h" #include "common/logger.h" #include "runtime/llmRankRuntime.h" +#include "runtime/greedySchedulerBackend.h" #include "runtime/multiDevice/runtimeCoordinator.h" #include @@ -237,6 +238,36 @@ bool LLMInferenceRuntime::hasDraftModel() const return rootRuntime().hasDraftModel(); } +std::unique_ptr LLMInferenceRuntime::createContinuousScheduler( + cudaStream_t stream, size_t maxQueued, size_t maxQueuedBytes) +{ + ELLM_CHECK(mCoordinator != nullptr, "Runtime coordinator is not initialized."); + ELLM_CHECK(!hasDraftModel(), "Continuous scheduling does not support speculative decoding."); + auto& runtime = rootRuntime(); + auto backend = std::make_unique( + runtime, stream, runtime.vocabularySize(), &runtime.tokenizer()); + return std::make_unique(std::move(backend), maxQueued, maxQueuedBytes); +} + +std::vector LLMInferenceRuntime::prepareContinuousPrompt(LLMGenerationRequest const& request) const +{ + ELLM_CHECK(mCoordinator != nullptr, "Runtime coordinator is not initialized."); + ELLM_CHECK(request.requests.size() == 1, "Continuous scheduling accepts one request at a time."); + auto const& row = request.requests.front(); + ELLM_CHECK(row.imageBuffers.empty() && row.audioBuffers.empty() && !row.pastTrajectory.has_value(), + "Continuous scheduling currently supports text-only requests."); + ELLM_CHECK(!request.saveSystemPromptKVCache, "Continuous scheduling does not support context-cache writes."); + auto prepared = mCoordinator->prepareRequestState(request); + ELLM_CHECK(prepared.preTokenizedInputIds.size() == 1 && !prepared.preTokenizedInputIds.front().empty(), + "Continuous scheduling could not tokenize request."); + return std::move(prepared.preTokenizedInputIds.front()); +} + +std::string LLMInferenceRuntime::continuousTokenPiece(int32_t tokenId) const +{ + return rootRuntime().tokenizer().idToPiece(tokenId, true); +} + bool LLMInferenceRuntime::ownsGlobalRank(int32_t globalRank) const noexcept { return mCoordinator ? mCoordinator->ownsGlobalRank(globalRank) : globalRank == 0; diff --git a/cpp/runtime/llmInferenceRuntime.h b/cpp/runtime/llmInferenceRuntime.h index 10d9c2ce4..0ffb5a460 100644 --- a/cpp/runtime/llmInferenceRuntime.h +++ b/cpp/runtime/llmInferenceRuntime.h @@ -21,6 +21,7 @@ #include "profiling/metrics.h" #include "runtime/config/deploymentConfig.h" #include "runtime/llmRuntimeUtils.h" +#include "runtime/continuousScheduler.h" #include "runtime/modelArtifacts.h" #include "runtime/multiDevice/parallelConfig.h" #include "runtime/preprocess/visualTokenPruner.h" @@ -116,6 +117,19 @@ class LLMInferenceRuntime std::vector> const& getBaseModelInputTokenIds() const; bool hasDraftModel() const; + //! Create the P5 native scheduler for the deliberately narrow P6 serving + //! mode: one local rank, vanilla text generation and no context cache. + //! The returned scheduler borrows this runtime and `stream`; callers must + //! close/destroy it before this object or stream is released. + std::unique_ptr createContinuousScheduler( + cudaStream_t stream, size_t maxQueued = 8, size_t maxQueuedBytes = 256 * 1024); + + //! Format and tokenize exactly once using the runtime tokenizer for a P6 + //! text request. Media and batched legacy requests are rejected rather + //! than silently dropping request state. + std::vector prepareContinuousPrompt(LLMGenerationRequest const& request) const; + std::string continuousTokenPiece(int32_t tokenId) const; + //! True when this runtime instance owns the requested global rank. bool ownsGlobalRank(int32_t globalRank) const noexcept; diff --git a/cpp/runtime/llmRankRuntime.h b/cpp/runtime/llmRankRuntime.h index 195728e9a..a7c5017b4 100644 --- a/cpp/runtime/llmRankRuntime.h +++ b/cpp/runtime/llmRankRuntime.h @@ -239,6 +239,18 @@ class LLMRankRuntime return mDecoderRegistry && mDecoderRegistry->hasSpeculativeDecoder(); } + //! Metadata borrowed by the single-rank continuous scheduler. The scheduler + //! never owns either object and must be destroyed before this runtime. + tokenizer::Tokenizer const& tokenizer() const + { + ELLM_CHECK(mTokenizer != nullptr, "LLMRankRuntime tokenizer is not initialized."); + return *mTokenizer; + } + int32_t vocabularySize() const + { + return mDeployment.base.vocabSize; + } + private: friend class ContinuousBatchingProbe; friend class SequenceStepRuntime; diff --git a/cpp/runtime/multiDevice/runtimeCoordinator.h b/cpp/runtime/multiDevice/runtimeCoordinator.h index 8c109138c..81271bcef 100644 --- a/cpp/runtime/multiDevice/runtimeCoordinator.h +++ b/cpp/runtime/multiDevice/runtimeCoordinator.h @@ -93,6 +93,10 @@ class RuntimeCoordinator LLMRankRuntime& rootRuntime(); LLMRankRuntime const& rootRuntime() const; + //! Formats and tokenizes a request without submitting it to rank workers. + //! Used only by the P6 single-rank continuous scheduler bridge. + LLMGenerationRequest prepareRequestState(LLMGenerationRequest const& request) const; + bool ownsGlobalRank(int32_t globalRank) const noexcept; bool localRanksSucceeded() const noexcept; @@ -137,7 +141,6 @@ class RuntimeCoordinator LLMGenerationRequest const& request, bool enableProfiling, bool outputThinkerEmbeddings, cudaStream_t stream); std::unique_ptr createRankRuntime(int32_t globalRank); - LLMGenerationRequest prepareRequestState(LLMGenerationRequest const& request) const; void prepareRankRequests(LLMGenerationRequest const& request); CollectiveGroup const* collectiveGroup(ParallelType type) const noexcept; CollectiveGroup* collectiveGroup(ParallelType type) noexcept; diff --git a/experimental/pybind/edgellm_pybind.cpp b/experimental/pybind/edgellm_pybind.cpp index f480485cc..3beb9cef9 100644 --- a/experimental/pybind/edgellm_pybind.cpp +++ b/experimental/pybind/edgellm_pybind.cpp @@ -35,6 +35,7 @@ #include "runtime/audioUtils.h" #include "runtime/imageUtils.h" #include "runtime/llmInferenceRuntime.h" +#include "runtime/continuousScheduler.h" #include "runtime/llmRuntimeUtils.h" #include "runtime/melSpectrogram.h" #ifdef EDGELLM_ENABLE_NEMOTRON_ASR @@ -274,6 +275,53 @@ class PyLLMRuntime return response; } + void enableContinuousBatching(size_t maxQueued, size_t maxQueuedBytes) + { + ELLM_CHECK(mScheduler == nullptr, "Continuous scheduler is already enabled"); + mScheduler = mRuntime->createContinuousScheduler(mStream.get(), maxQueued, maxQueuedBytes); + } + + std::vector prepareContinuousPrompt(LLMGenerationRequest const& request) const + { + ELLM_CHECK(mScheduler != nullptr, "Enable continuous batching before preparing requests"); + return mRuntime->prepareContinuousPrompt(request); + } + + SchedulerTicket submitContinuous(std::vector const& prompt, SchedulerRequestOptions const& options) + { + ELLM_CHECK(mScheduler != nullptr, "Enable continuous batching before submitting requests"); + return mScheduler->submit(prompt, options); + } + + py::bytes continuousTokenPiece(int32_t tokenId) const + { + return py::bytes(mRuntime->continuousTokenPiece(tokenId)); + } + + bool continuousHealthy() const + { + return mScheduler != nullptr && mScheduler->healthy(); + } + + size_t continuousQueuedRequests() const + { + return mScheduler ? mScheduler->queuedCount() : 0; + } + + size_t continuousResidentRequests() const + { + return mScheduler ? mScheduler->residentCount() : 0; + } + + void closeContinuousBatching() + { + if (mScheduler) + { + mScheduler->close(); + mScheduler.reset(); + } + } + //! Load the Qwen3-Omni audio-output stack (Talker + CodePredictor + Code2Wav). //! tokenizerDir is normally the Thinker engine dir. void loadOmni(std::string const& talkerEngineDir, std::string const& codePredictorEngineDir, @@ -409,6 +457,7 @@ class PyLLMRuntime CudaStreamWrapper mStream; std::unique_ptr mPluginHandle; std::unique_ptr mRuntime; + std::unique_ptr mScheduler; std::unique_ptr mTtsRuntime; std::unique_ptr mCode2wavRunner; }; @@ -915,6 +964,79 @@ PYBIND11_MODULE(_edgellm_runtime, m) .def_readonly("finish_reasons", &LLMGenerationResponse::finishReasons) .def_readonly("prompt_token_counts", &LLMGenerationResponse::inputTokenCounts); + // ======================================================================== + // Continuous scheduler (P6 bridge) + // ======================================================================== + py::enum_(m, "SchedulerStatus") + .value("COMPLETED", SchedulerStatus::kCompleted) + .value("CANCELLED", SchedulerStatus::kCancelled) + .value("FAILED", SchedulerStatus::kFailed) + .value("DEADLINE", SchedulerStatus::kDeadline) + .value("SLOW_CONSUMER", SchedulerStatus::kSlowConsumer); + + py::enum_(m, "SequenceFinish") + .value("NONE", SequenceFinish::kNone) + .value("LENGTH", SequenceFinish::kLength) + .value("EOS", SequenceFinish::kEos) + .value("STOP", SequenceFinish::kStop); + + py::class_(m, "NativeTokenLogprob") + .def_readonly("token", &TokenLogprob::token) + .def_readonly("logprob", &TokenLogprob::logprob); + py::class_(m, "NativeSequenceSample") + .def_readonly("token", &SequenceSample::token) + .def_readonly("logprob", &SequenceSample::logprob) + .def_readonly("top", &SequenceSample::top) + .def_readonly("top_count", &SequenceSample::topCount) + .def_readonly("random_counter", &SequenceSample::randomCounter) + .def_readonly("thinking", &SequenceSample::thinking); + py::class_(m, "NativeSequenceUpdate") + .def_readonly("sample", &SequenceUpdate::sample) + .def_readonly("text", &SequenceUpdate::text); + py::class_(m, "NativeSequenceRead") + .def_readonly("update", &SequenceRead::update) + .def_readonly("closed", &SequenceRead::closed); + py::class_(m, "NativeSchedulerResult") + .def_readonly("status", &SchedulerResult::status) + .def_readonly("tokens", &SchedulerResult::tokens) + .def_readonly("text", &SchedulerResult::text) + .def_readonly("logprobs", &SchedulerResult::logprobs) + .def_readonly("finish", &SchedulerResult::finish) + .def_readonly("prompt_tokens", &SchedulerResult::promptTokens) + .def_readonly("random_counter", &SchedulerResult::randomCounter); + py::class_(m, "ContinuousSequenceOptions") + .def(py::init<>()) + .def_readwrite("max_output_tokens", &SequenceOptions::maxOutputTokens) + .def_readwrite("temperature", &SequenceOptions::temperature) + .def_readwrite("top_p", &SequenceOptions::topP) + .def_readwrite("top_k", &SequenceOptions::topK) + .def_readwrite("seed", &SequenceOptions::seed) + .def_readwrite("num_logprobs", &SequenceOptions::numLogprobs) + .def_readwrite("enable_thinking", &SequenceOptions::enableThinking) + .def_readwrite("ignore_eos", &SequenceOptions::ignoreEos) + .def_readwrite("eos_token_ids", &SequenceOptions::eosTokenIds) + .def_readwrite("stop_strings", &SequenceOptions::stopStrings) + .def_readwrite("logit_bias", &SequenceOptions::logitBias); + py::class_(m, "ContinuousRequestOptions") + .def(py::init<>()) + .def_readwrite("generation", &SchedulerRequestOptions::generation) + .def_readwrite("stream_records", &SchedulerRequestOptions::streamRecords) + .def_readwrite("stream_bytes", &SchedulerRequestOptions::streamBytes) + .def("set_queue_timeout_ms", [](SchedulerRequestOptions& self, int64_t timeoutMs) { + self.queueDeadline = std::chrono::steady_clock::now() + std::chrono::milliseconds{timeoutMs}; + }) + .def("set_timeout_ms", [](SchedulerRequestOptions& self, int64_t timeoutMs) { + self.deadline = std::chrono::steady_clock::now() + std::chrono::milliseconds{timeoutMs}; + }); + py::class_(m, "ContinuousTicket") + .def("id", &SchedulerTicket::id) + .def("cancel", &SchedulerTicket::cancel) + .def("read", [](SchedulerTicket const& self, int64_t timeoutMs) { + return self.read(std::chrono::milliseconds{timeoutMs}); + }, py::arg("timeout_ms"), py::call_guard()) + .def("result", [](SchedulerTicket const& self) { return self.result().get(); }, + py::call_guard()); + // ======================================================================== // Runtime: unified (vanilla + Eagle speculative decoding) // ======================================================================== @@ -934,6 +1056,19 @@ PYBIND11_MODULE(_edgellm_runtime, m) py::arg("dflash_block_size") = 0, "Construct for speculative decoding") .def("handle_request", &PyLLMRuntime::handleRequest, py::arg("request"), py::call_guard(), "Process a generation request and return the response") + .def("enable_continuous_batching", &PyLLMRuntime::enableContinuousBatching, + py::arg("max_queued") = 8, py::arg("max_queued_bytes") = 256 * 1024, + py::call_guard()) + .def("prepare_continuous_prompt", &PyLLMRuntime::prepareContinuousPrompt, py::arg("request"), + py::call_guard()) + .def("submit_continuous", &PyLLMRuntime::submitContinuous, py::arg("prompt"), py::arg("options"), + py::call_guard()) + .def("continuous_token_piece", &PyLLMRuntime::continuousTokenPiece, py::arg("token_id")) + .def("continuous_healthy", &PyLLMRuntime::continuousHealthy) + .def("continuous_queued_requests", &PyLLMRuntime::continuousQueuedRequests) + .def("continuous_resident_requests", &PyLLMRuntime::continuousResidentRequests) + .def("close_continuous_batching", &PyLLMRuntime::closeContinuousBatching, + py::call_guard()) .def("load_omni", &PyLLMRuntime::loadOmni, py::arg("talker_engine_dir"), py::arg("code_predictor_engine_dir"), py::arg("code2wav_engine_dir"), py::arg("tokenizer_dir"), py::arg("checkpoint_dir") = "", "Load the Qwen3-Omni audio-output stack (Talker + CodePredictor + Code2Wav)") diff --git a/experimental/server/runtime/engine.py b/experimental/server/runtime/engine.py index 1b92752df..1f9cdf527 100644 --- a/experimental/server/runtime/engine.py +++ b/experimental/server/runtime/engine.py @@ -91,6 +91,7 @@ class SamplingParams: skip_special_tokens: bool = True reuse_context: bool = True cache_generated_tokens: bool = True + seed: int = 42 @dataclass @@ -821,10 +822,130 @@ def _load_runtime(self) -> None: self._model_dir, context_cache_config, ) - self._runtime.capture_decoding_cuda_graph() + self._continuous_batching = False + # P6 is deliberately opt-in until P7 has qualified the existing graph + # captures, allocation behaviour and singleton regression. Candidate + # service units set this explicit switch; legacy deployments retain + # their exact single-request path. + if (os.environ.get("EDGELLM_CONTINUOUS_BATCHING", "").lower() + in {"1", "true", "on"}): + self.enable_continuous_batching() + else: + self._runtime.capture_decoding_cuda_graph() self._load_omni_runtime() logger.info("Engine loaded and ready.") + @property + def continuous_batching_enabled(self) -> bool: + return self._continuous_batching + + def enable_continuous_batching(self) -> None: + """Enable P6 native scheduling for supported vanilla text bundles.""" + if self._continuous_batching: + return + if self.has_draft_model or self._max_batch_size < 2: + raise ValueError("continuous batching requires a vanilla batch-two engine") + if self._layout.visual_dir or self._layout.audio_dir: + raise ValueError("continuous batching currently supports text-only engines") + if self._context_cache_config.enabled: + raise ValueError("continuous batching does not support context cache") + self._runtime.enable_continuous_batching() + self._continuous_batching = True + + def _continuous_options(self, params: SamplingParams, *, stream: bool): + options = self._rt.ContinuousRequestOptions() + generation = self._rt.ContinuousSequenceOptions() + generation.max_output_tokens = params.max_tokens + generation.temperature = params.temperature + generation.top_p = params.top_p + generation.top_k = params.top_k + generation.seed = params.seed + generation.num_logprobs = params.num_logprobs + generation.enable_thinking = params.enable_thinking + generation.stop_strings = params.stop + generation.logit_bias = _normalize_logit_bias(params.logit_bias) + options.generation = generation + options.stream_records = min(max(params.max_tokens + 2, 2), 8192) if stream else 0 + options.stream_bytes = min(max(params.max_tokens * 16, 16384), 1 << 20) + return options + + def _submit_continuous(self, request, params: SamplingParams, *, stream: bool): + self._ensure_open() + prompt = self._runtime.prepare_continuous_prompt(request) + return self._runtime.submit_continuous( + prompt, self._continuous_options(params, stream=stream)) + + def _continuous_logprobs(self, samples) -> List[List[LogprobEntry]]: + converted = [] + for sample in samples: + entries = [] + for entry in list(sample.top)[:sample.top_count]: + piece = self._runtime.continuous_token_piece(entry.token) + entries.append(LogprobEntry(entry.token, entry.logprob, + piece.decode("utf-8", "replace"), + list(piece))) + converted.append(entries) + return converted + + def _complete_continuous_request(self, request, params: SamplingParams, + tool_config: ToolConfig, *, tool_parser: str, + reasoning_parser: str) -> CompletionOutput: + ticket = self._submit_continuous(request, params, stream=False) + result = ticket.result() + if result.status != self._rt.SchedulerStatus.COMPLETED: + raise RuntimeError(f"continuous generation failed: {result.status}") + finish_reason = { + self._rt.SequenceFinish.LENGTH: "length", + self._rt.SequenceFinish.EOS: "stop", + self._rt.SequenceFinish.STOP: "stop", + }.get(result.finish, "stop") + output = self._parse_generation_output( + result.text, list(result.tokens), result.prompt_tokens, finish_reason, tool_config, + tool_parser=tool_parser, reasoning_parser=reasoning_parser) + if params.num_logprobs: + output.logprobs = self._continuous_logprobs(result.logprobs) + return output + + def generate_continuous_stream(self, request, params: SamplingParams) -> Iterator[StreamDelta]: + """Read one P5 ticket channel; closing this iterator cancels only it.""" + state = {} + + def _cancel(): + ticket = state.get("ticket") + if ticket is not None: + ticket.cancel() + + def _iterate(): + ticket = self._submit_continuous(request, params, stream=True) + state["ticket"] = ticket + terminal = None + try: + while True: + read = ticket.read(timeout_ms=200) + if read.update is not None: + update = read.update + yield StreamDelta(text=update.text, + token_ids=[update.sample.token], + logprobs=self._continuous_logprobs([update.sample])) + if read.closed: + terminal = ticket.result() + if terminal.status != self._rt.SchedulerStatus.COMPLETED: + raise RuntimeError( + f"continuous generation failed: {terminal.status}") + reason = { + self._rt.SequenceFinish.LENGTH: "length", + self._rt.SequenceFinish.EOS: "stop", + self._rt.SequenceFinish.STOP: "stop", + }.get(terminal.finish, "stop") + yield StreamDelta(finished=True, finish_reason=reason, + prompt_tokens=terminal.prompt_tokens) + return + finally: + if terminal is None: + _cancel() + + return _CancellableIterator(_iterate(), _cancel) + def _load_omni_runtime(self) -> None: """Load the Qwen3-Omni audio-output stack when its engines exist.""" if not self._layout.has_speech: @@ -1243,6 +1364,9 @@ def close(self) -> None: self._closed = True with self._admission_sem: with self._infer_lock: + if getattr(self, "_continuous_batching", False): + self._runtime.close_continuous_batching() + self._continuous_batching = False self._runtime = None def __enter__(self) -> "LLM": diff --git a/experimental/server/runtime/engine_client.py b/experimental/server/runtime/engine_client.py index 68ae7f7d3..37461a043 100644 --- a/experimental/server/runtime/engine_client.py +++ b/experimental/server/runtime/engine_client.py @@ -190,7 +190,7 @@ def _capabilities_for(llm: Union[LLM, TTS]) -> EngineCapabilities: builder.get("max_input_len"), int) else None, max_batch_size=builder.get("max_batch_size") if isinstance( builder.get("max_batch_size"), int) else None, - max_num_seqs=1, + max_num_seqs=2 if getattr(llm, "continuous_batching_enabled", False) else 1, kv_cache_dtype=str(config.get("kv_cache_dtype", "unknown")), speculative_decoding=llm.has_draft_model, speculative_method=str(config.get("spec_decode_type", "none")), @@ -280,10 +280,14 @@ def capabilities(self) -> EngineCapabilities: @property def active_requests(self) -> int: + if getattr(self._llm, "continuous_batching_enabled", False): + return int(self._llm._runtime.continuous_resident_requests()) return self._admission.active @property def queued_requests(self) -> int: + if getattr(self._llm, "continuous_batching_enabled", False): + return int(self._llm._runtime.continuous_queued_requests()) return self._admission.waiting async def close(self) -> None: @@ -316,6 +320,18 @@ async def count_prompt_tokens( tool_config: ToolConfig, enable_thinking: bool, ) -> int: + if getattr(self._llm, "continuous_batching_enabled", False): + request = await _run_sync(partial( + self._llm._make_generation_request, messages, + SamplingParams(max_tokens=1, enable_thinking=enable_thinking), + tools=tool_config.tools, tool_choice=tool_config.tool_choice, + tool_config=tool_config)) + count = await _run_sync(partial( + self._llm._count_prepared_prompt_tokens, request)) + if count is None: + raise UnsupportedFeatureError( + "exact token counting is not available for media inputs") + return count prepared = await self.prepare_request( messages, SamplingParams(max_tokens=1, enable_thinking=enable_thinking), @@ -351,6 +367,16 @@ async def generate( tool_config = tool_config or validate_tool_request( messages, tools, tool_choice) if owned is None: + if getattr(self._llm, "continuous_batching_enabled", False): + request = await _run_sync(partial( + self._llm._make_generation_request, + messages, sampling_params, tools=tool_config.tools, + tool_choice=tool_config.tool_choice, + tool_config=tool_config)) + return await _run_sync(partial( + self._llm._complete_continuous_request, request, + sampling_params, tool_config, tool_parser=tool_parser, + reasoning_parser=reasoning_parser)) owned = await self.prepare_request( messages, sampling_params, @@ -410,6 +436,26 @@ async def stream( prepared: Optional[PreparedRequest] = None, ) -> AsyncGenerator[StreamDelta, None]: iterator = None + if (prepared is None + and getattr(self._llm, "continuous_batching_enabled", False)): + try: + request = await _run_sync(partial( + self._llm._make_generation_request, messages, + sampling_params, tools=tools, tool_choice=tool_choice)) + iterator = self._llm.generate_continuous_stream( + request, sampling_params) + async for item in _iterate_sync(iterator): + yield item + return + except (ServerError, KeyError, TypeError, ValueError): + raise + except asyncio.CancelledError: + raise + except Exception as exc: + raise EngineError(str(exc)) from exc + finally: + if iterator is not None: + await asyncio.to_thread(_close_stream, iterator) owned = prepared or await self.prepare_request( messages, sampling_params, From 87aa55e4a8f40062ed2163549fb536bc9e677099 Mon Sep 17 00:00:00 2001 From: ajmalrasi Date: Tue, 15 Sep 2026 20:30:15 +0530 Subject: [PATCH 6/7] fix: route prepared HTTP requests through continuous tickets Signed-off-by: ajmalrasi --- experimental/server/api/routes.py | 4 + experimental/server/api/serving_chat.py | 12 +- experimental/server/config.py | 8 + experimental/server/runtime/engine.py | 48 ++++-- experimental/server/runtime/engine_client.py | 91 ++++++----- .../python-unittests/test_continuous_http.py | 153 ++++++++++++++++++ 6 files changed, 267 insertions(+), 49 deletions(-) create mode 100644 tests/python-unittests/test_continuous_http.py diff --git a/experimental/server/api/routes.py b/experimental/server/api/routes.py index f4648a6b2..81216ed59 100644 --- a/experimental/server/api/routes.py +++ b/experimental/server/api/routes.py @@ -45,6 +45,8 @@ async def __call__(self, scope, receive, send): async def health(request: Request): client = request.app.state.engine_client caps = client.capabilities + if not client.healthy: + return JSONResponse(status_code=503, content={"status": "unhealthy"}) return { "status": "healthy", "model": client.model_name, @@ -71,6 +73,8 @@ async def health(request: Request): @router.get("/health/ready") async def readiness(request: Request): client = request.app.state.engine_client + if not client.healthy: + return JSONResponse(status_code=503, content={"status": "unready"}) return { "status": "ready", "active_requests": client.active_requests, diff --git a/experimental/server/api/serving_chat.py b/experimental/server/api/serving_chat.py index ebaad8447..038ce1ec3 100644 --- a/experimental/server/api/serving_chat.py +++ b/experimental/server/api/serving_chat.py @@ -70,10 +70,11 @@ def _format_logprob_steps(token_ids, steps, return None content = [] for token_id, step in zip(token_ids, steps): - top = [_entry_to_openai(entry) for entry in step] + top = [_entry_to_openai(entry) for entry in step if not entry.chosen_only] + candidates = [_entry_to_openai(entry) for entry in step] chosen = next( (candidate - for candidate in top if candidate["token_id"] == token_id), None) + for candidate in candidates if candidate["token_id"] == token_id), None) content.append({ "token": chosen["token"] if chosen else "", "token_id": token_id, @@ -154,9 +155,12 @@ def prepare_request(self, raise UnsupportedFeatureError( "presence_penalty is not supported by the Edge-LLM runtime", param="presence_penalty") - if request.seed is not None: + if (request.seed is not None + and not getattr(self._client.llm, "continuous_batching_enabled", False)): raise UnsupportedFeatureError( "seed is not supported by the Edge-LLM runtime", param="seed") + if request.seed is not None and not 0 <= request.seed < (1 << 64): + raise InvalidRequestError("seed must fit an unsigned 64-bit integer", param="seed") if request.response_format is not None: raise UnsupportedFeatureError( "response_format requires structured decoding, which is not " @@ -223,6 +227,7 @@ def prepare_request(self, num_logprobs = max(1, request.top_logprobs or 0) greedy = request.temperature == 0 sampling = SamplingParams( + seed=request.seed if request.seed is not None else 42, temperature=request.temperature, top_p=1.0 if greedy else request.top_p, top_k=1 if greedy else request.top_k, @@ -372,6 +377,7 @@ async def prepare_engine_request( tools=prepared.tool_config.tools, tool_choice=prepared.tool_config.tool_choice, tool_config=prepared.tool_config, + stream=request.stream, ) return engine_request except (KeyError, TypeError, ValueError) as exc: diff --git a/experimental/server/config.py b/experimental/server/config.py index 7be49bc83..e5f92f48c 100644 --- a/experimental/server/config.py +++ b/experimental/server/config.py @@ -165,6 +165,7 @@ class ModelConfig: """Checkpoint, build-cache, and runtime-profile options.""" model: str + engine_dir: str = "" cache_dir: str = "" engine_cache_max_size_gb: float = 50.0 clear_engine_cache: bool = False @@ -188,6 +189,7 @@ def __post_init__(self) -> None: def llm_kwargs(self) -> Dict[str, Any]: return { "model": self.model, + "engine_dir": self.engine_dir, "cache_dir": self.cache_dir, "engine_cache_max_size_gb": self.engine_cache_max_size_gb, "clear_engine_cache": self.clear_engine_cache, @@ -298,6 +300,11 @@ def create_argument_parser() -> argparse.ArgumentParser: ) model = parser.add_argument_group("Model build and runtime") + model.add_argument( + "--engine-dir", + default="", + help="Load an existing Edge-LLM engine bundle and skip engine build.", + ) model.add_argument( "--cache-dir", default="", @@ -368,6 +375,7 @@ def parse_server_config(argv: Optional[Sequence[str]] = None) -> ServerConfig: ) model = ModelConfig( model=args.model, + engine_dir=args.engine_dir, cache_dir=args.cache_dir, engine_cache_max_size_gb=args.engine_cache_max_size_gb, clear_engine_cache=args.clear_engine_cache, diff --git a/experimental/server/runtime/engine.py b/experimental/server/runtime/engine.py index 1f9cdf527..a519677ea 100644 --- a/experimental/server/runtime/engine.py +++ b/experimental/server/runtime/engine.py @@ -46,6 +46,7 @@ from typing import (TYPE_CHECKING, Any, Dict, Iterator, List, Mapping, Optional, Sequence, Union) +from ..api.errors import ServerOverloadedError from ..config import ContextCacheConfig from ..parsing.tool_calling import (ToolConfig, parse_assistant_output, validate_tool_request) @@ -107,6 +108,7 @@ class LogprobEntry: logprob: float token: str bytes: List[int] + chosen_only: bool = False def _convert_logprobs(raw) -> List[List[LogprobEntry]]: @@ -650,6 +652,7 @@ def __init__( self, model: str, *, + engine_dir: str = "", cache_dir: str = "", engine_cache_max_size_gb: float = 50.0, clear_engine_cache: bool = False, @@ -706,6 +709,18 @@ def __init__( self._closed = False self._runtime = None + if engine_dir: + if not os.path.isdir(model): + raise ValueError( + "an existing --engine-dir requires a local checkpoint " + "directory for tokenizer and external weights") + self._cache_dir = "" + self._model_dir = os.path.abspath(model) + self._draft_model_dir = "" + self._init_from_bundle(os.path.abspath(engine_dir)) + self._load_runtime() + return + from .engine_build import BuildOptions, cache_root, prepare_model options = build_options or BuildOptions( @@ -865,8 +880,8 @@ def _continuous_options(self, params: SamplingParams, *, stream: bool): generation.stop_strings = params.stop generation.logit_bias = _normalize_logit_bias(params.logit_bias) options.generation = generation - options.stream_records = min(max(params.max_tokens + 2, 2), 8192) if stream else 0 - options.stream_bytes = min(max(params.max_tokens * 16, 16384), 1 << 20) + options.stream_records = min(max(params.max_tokens + 2, 2), 64) if stream else 0 + options.stream_bytes = 65536 if stream else 0 return options def _submit_continuous(self, request, params: SamplingParams, *, stream: bool): @@ -884,14 +899,22 @@ def _continuous_logprobs(self, samples) -> List[List[LogprobEntry]]: entries.append(LogprobEntry(entry.token, entry.logprob, piece.decode("utf-8", "replace"), list(piece))) + if not any(entry.token_id == sample.token for entry in entries): + piece = self._runtime.continuous_token_piece(sample.token) + entries.append(LogprobEntry(sample.token, sample.logprob, + piece.decode("utf-8", "replace"), + list(piece), chosen_only=True)) converted.append(entries) return converted def _complete_continuous_request(self, request, params: SamplingParams, tool_config: ToolConfig, *, tool_parser: str, - reasoning_parser: str) -> CompletionOutput: - ticket = self._submit_continuous(request, params, stream=False) + reasoning_parser: str, ticket=None) -> CompletionOutput: + if ticket is None: + ticket = self._submit_continuous(request, params, stream=False) result = ticket.result() + if result.status == self._rt.SchedulerStatus.DEADLINE: + raise ServerOverloadedError("native request deadline expired") if result.status != self._rt.SchedulerStatus.COMPLETED: raise RuntimeError(f"continuous generation failed: {result.status}") finish_reason = { @@ -906,9 +929,9 @@ def _complete_continuous_request(self, request, params: SamplingParams, output.logprobs = self._continuous_logprobs(result.logprobs) return output - def generate_continuous_stream(self, request, params: SamplingParams) -> Iterator[StreamDelta]: + def generate_continuous_stream(self, request, params: SamplingParams, ticket=None) -> Iterator[StreamDelta]: """Read one P5 ticket channel; closing this iterator cancels only it.""" - state = {} + state = {"ticket": ticket} def _cancel(): ticket = state.get("ticket") @@ -916,8 +939,10 @@ def _cancel(): ticket.cancel() def _iterate(): - ticket = self._submit_continuous(request, params, stream=True) - state["ticket"] = ticket + ticket = state["ticket"] + if ticket is None: + ticket = self._submit_continuous(request, params, stream=True) + state["ticket"] = ticket terminal = None try: while True: @@ -925,10 +950,13 @@ def _iterate(): if read.update is not None: update = read.update yield StreamDelta(text=update.text, - token_ids=[update.sample.token], - logprobs=self._continuous_logprobs([update.sample])) + token_ids=([update.sample.token] if update.sample.token >= 0 else []), + logprobs=(self._continuous_logprobs([update.sample]) + if params.num_logprobs and update.sample.token >= 0 else [])) if read.closed: terminal = ticket.result() + if terminal.status == self._rt.SchedulerStatus.DEADLINE: + raise ServerOverloadedError("native request deadline expired") if terminal.status != self._rt.SchedulerStatus.COMPLETED: raise RuntimeError( f"continuous generation failed: {terminal.status}") diff --git a/experimental/server/runtime/engine_client.py b/experimental/server/runtime/engine_client.py index 37461a043..570d2b70f 100644 --- a/experimental/server/runtime/engine_client.py +++ b/experimental/server/runtime/engine_client.py @@ -51,8 +51,9 @@ class EngineCapabilities: class _AdmissionController: """Bounded queue for the runtime's single generation slot.""" - def __init__(self, max_queued_requests: int, timeout: float) -> None: - self._semaphore = asyncio.Semaphore(1) + def __init__(self, max_queued_requests: int, timeout: float, capacity: int = 1) -> None: + self._capacity = capacity + self._semaphore = asyncio.Semaphore(capacity) self._max_queued = max_queued_requests self._timeout = timeout self._active = 0 @@ -73,7 +74,7 @@ async def reserve(self) -> "_AdmissionLease": acquired = False if self._closing: raise ServerUnavailableError() - if self._active + self._waiting >= self._max_queued + 1: + if self._active + self._waiting >= self._max_queued + self._capacity: raise ServerOverloadedError() self._waiting += 1 @@ -93,7 +94,7 @@ async def reserve(self) -> "_AdmissionLease": self._semaphore.release() acquired = False raise ServerUnavailableError() - self._active = 1 + self._active += 1 return _AdmissionLease(self) except BaseException: if acquired: @@ -101,7 +102,7 @@ async def reserve(self) -> "_AdmissionLease": raise def release(self) -> None: - self._active = 0 + self._active -= 1 self._semaphore.release() async def close(self) -> None: @@ -110,7 +111,8 @@ async def close(self) -> None: if self._drained: return self._closing = True - await self._semaphore.acquire() + for _ in range(self._capacity): + await self._semaphore.acquire() self._drained = True @@ -133,8 +135,11 @@ class PreparedRequest: request: Any lease: _AdmissionLease + ticket: Any = None def release(self) -> None: + if self.ticket is not None: + self.ticket.cancel() self.lease.release() @@ -260,6 +265,8 @@ def __init__(self, self._admission = _AdmissionController( self._api_config.max_queued_requests, self._api_config.queue_timeout, + capacity=(self._api_config.max_queued_requests + 2 + if getattr(llm, "continuous_batching_enabled", False) else 1), ) self._capabilities = _capabilities_for(llm) self._close_lock = asyncio.Lock() @@ -278,6 +285,14 @@ def model_name(self) -> str: def capabilities(self) -> EngineCapabilities: return self._capabilities + @property + def healthy(self) -> bool: + if self._closed: + return False + if getattr(self._llm, "continuous_batching_enabled", False): + return bool(self._llm._runtime.continuous_healthy()) + return True + @property def active_requests(self) -> int: if getattr(self._llm, "continuous_batching_enabled", False): @@ -367,16 +382,6 @@ async def generate( tool_config = tool_config or validate_tool_request( messages, tools, tool_choice) if owned is None: - if getattr(self._llm, "continuous_batching_enabled", False): - request = await _run_sync(partial( - self._llm._make_generation_request, - messages, sampling_params, tools=tool_config.tools, - tool_choice=tool_config.tool_choice, - tool_config=tool_config)) - return await _run_sync(partial( - self._llm._complete_continuous_request, request, - sampling_params, tool_config, tool_parser=tool_parser, - reasoning_parser=reasoning_parser)) owned = await self.prepare_request( messages, sampling_params, @@ -384,6 +389,18 @@ async def generate( tool_choice=tool_config.tool_choice, tool_config=tool_config, ) + if owned.ticket is not None: + operation = partial( + self._llm._complete_continuous_request, owned.request, + sampling_params, tool_config, tool_parser=tool_parser, + reasoning_parser=reasoning_parser, ticket=owned.ticket) + worker = asyncio.create_task(asyncio.to_thread(operation)) + try: + return await asyncio.shield(worker) + except asyncio.CancelledError: + owned.ticket.cancel() + await asyncio.shield(asyncio.gather(worker, return_exceptions=True)) + raise operation = partial( self._llm._complete_prepared_request, owned.request, @@ -409,6 +426,7 @@ async def prepare_request( tools: Optional[Sequence[Dict[str, Any]]] = None, tool_choice: Optional[Union[str, Dict[str, Any]]] = None, tool_config: Optional[ToolConfig] = None, + stream: bool = False, ) -> PreparedRequest: lease = await self._admission.reserve() try: @@ -421,6 +439,20 @@ async def prepare_request( tool_choice=tool_choice, tool_config=tool_config, )) + if getattr(self._llm, "continuous_batching_enabled", False): + # No await between submission and ownership transfer can lose a ticket. + prompt = self._llm._runtime.prepare_continuous_prompt(request) + options = self._llm._continuous_options(sampling_params, stream=stream) + options.set_queue_timeout_ms(int(self._api_config.queue_timeout * 1000)) + try: + ticket = self._llm._runtime.submit_continuous(prompt, options) + except RuntimeError as exc: + if not self._llm._runtime.continuous_healthy(): + raise ServerUnavailableError() from exc + if "queue full" in str(exc): + raise ServerOverloadedError() from exc + raise + return PreparedRequest(request=request, lease=lease, ticket=ticket) return PreparedRequest(request=request, lease=lease) except BaseException: lease.release() @@ -436,33 +468,20 @@ async def stream( prepared: Optional[PreparedRequest] = None, ) -> AsyncGenerator[StreamDelta, None]: iterator = None - if (prepared is None - and getattr(self._llm, "continuous_batching_enabled", False)): - try: - request = await _run_sync(partial( - self._llm._make_generation_request, messages, - sampling_params, tools=tools, tool_choice=tool_choice)) - iterator = self._llm.generate_continuous_stream( - request, sampling_params) - async for item in _iterate_sync(iterator): - yield item - return - except (ServerError, KeyError, TypeError, ValueError): - raise - except asyncio.CancelledError: - raise - except Exception as exc: - raise EngineError(str(exc)) from exc - finally: - if iterator is not None: - await asyncio.to_thread(_close_stream, iterator) owned = prepared or await self.prepare_request( messages, sampling_params, tools=tools, tool_choice=tool_choice, + stream=True, ) try: + if owned.ticket is not None: + iterator = self._llm.generate_continuous_stream( + owned.request, sampling_params, ticket=owned.ticket) + async for item in _iterate_sync(iterator): + yield item + return iterator = self._llm.generate_stream( messages, sampling_params, diff --git a/tests/python-unittests/test_continuous_http.py b/tests/python-unittests/test_continuous_http.py new file mode 100644 index 000000000..07e1930b5 --- /dev/null +++ b/tests/python-unittests/test_continuous_http.py @@ -0,0 +1,153 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Exercise the prepared HTTP boundary without a GPU.""" +import asyncio +import threading +from types import SimpleNamespace + +import pytest + +from experimental.server.config import ApiConfig +from experimental.server.runtime import engine_client +from experimental.server.runtime.engine import LLM, SamplingParams + + +class Options(SimpleNamespace): + def set_queue_timeout_ms(self, value): + self.timeout = value + + +class Ticket: + def __init__(self): + self.cancelled = threading.Event() + + def cancel(self): + self.cancelled.set() + + +class Runtime: + def __init__(self): + self.tickets = [] + self.failure = None + + def prepare_continuous_prompt(self, request): + return [1, 2, 3] + + def submit_continuous(self, prompt, options): + if self.failure: + raise RuntimeError(self.failure) + ticket = Ticket() + self.tickets.append(ticket) + return ticket + + def continuous_healthy(self): + return self.failure != "failed" + + +class Model: + model_dir = "test" + model_id = "openclaw" + continuous_batching_enabled = True + + def __init__(self): + self._runtime = Runtime() + self._rt = SimpleNamespace(ContinuousRequestOptions=Options, + ContinuousSequenceOptions=Options) + + _continuous_options = LLM._continuous_options + + def _make_generation_request(self, *args, **kwargs): + return SimpleNamespace() + + def _complete_continuous_request(self, *args, ticket, **kwargs): + assert ticket in self._runtime.tickets + return "completed" + + def generate_continuous_stream(self, *args, ticket): + assert ticket in self._runtime.tickets + yield "streamed" + + +def client(monkeypatch): + monkeypatch.setattr(engine_client, "_capabilities_for", lambda llm: None) + return engine_client.EngineClient(Model(), ApiConfig()) + + +def test_default_stream_storage_fits_native_admission(): + options = Model()._continuous_options(SamplingParams(), stream=True) + assert options.stream_records * 1024 + options.stream_bytes + 6144 * 4 < 256 * 1024 + assert Model()._continuous_options(SamplingParams(), stream=False).stream_records == 0 + + +@pytest.mark.asyncio +async def test_prepared_requests_overlap_and_use_native_tickets(monkeypatch): + owner = client(monkeypatch) + params = SamplingParams(max_tokens=4) + first = await owner.prepare_request([], params) + second = await asyncio.wait_for(owner.prepare_request([], params, stream=True), .5) + assert first.ticket is not second.ticket + assert await owner.generate([], params, prepared=first) == "completed" + assert [item async for item in owner.stream([], params, prepared=second)] == ["streamed"] + assert first.ticket.cancelled.is_set() and second.ticket.cancelled.is_set() + assert owner._admission.active == 0 + + +@pytest.mark.asyncio +async def test_nonstream_disconnect_cancels_owned_ticket(monkeypatch): + owner = client(monkeypatch) + params = SamplingParams(max_tokens=4) + prepared = await owner.prepare_request([], params) + entered = threading.Event() + + def blocking(*args, ticket, **kwargs): + entered.set() + assert ticket.cancelled.wait(2) + + owner._llm._complete_continuous_request = blocking + task = asyncio.create_task(owner.generate([], params, prepared=prepared)) + assert await asyncio.to_thread(entered.wait, 1) + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + assert prepared.ticket.cancelled.is_set() + assert owner._admission.active == 0 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("failure,exception", [ + ("Scheduler queue full", engine_client.ServerOverloadedError), + ("failed", engine_client.ServerUnavailableError)]) +async def test_native_rejection_releases_preparation(monkeypatch, failure, exception): + owner = client(monkeypatch) + owner._llm._runtime.failure = failure + with pytest.raises(exception): + await owner.prepare_request([], SamplingParams()) + assert owner._admission.active == 0 + assert owner.healthy == (failure != "failed") + + +def test_text_only_flush_does_not_inflate_usage(): + llm = LLM.__new__(LLM) + llm._rt = SimpleNamespace(SchedulerStatus=SimpleNamespace(COMPLETED=1, DEADLINE=2), + SequenceFinish=SimpleNamespace(LENGTH=1, EOS=2, STOP=3)) + reads = iter([ + SimpleNamespace(update=SimpleNamespace(text="A", sample=SimpleNamespace(token=8)), closed=False), + SimpleNamespace(update=SimpleNamespace(text="!", sample=SimpleNamespace(token=-1)), closed=True)]) + ticket = SimpleNamespace(read=lambda **kwargs: next(reads), cancel=lambda: None, + result=lambda: SimpleNamespace(status=1, finish=1, prompt_tokens=3)) + updates = list(llm.generate_continuous_stream(None, SamplingParams(), ticket=ticket)) + assert [token for update in updates for token in update.token_ids] == [8] + assert "".join(update.text for update in updates) == "A!" + assert updates[-1].finished + + +def test_sampled_token_outside_top_logprobs_keeps_probability(): + from experimental.server.api.serving_chat import _format_logprob_steps + llm = LLM.__new__(LLM) + llm._runtime = SimpleNamespace(continuous_token_piece=lambda token: b"X" if token == 9 else b"A") + sample = SimpleNamespace(token=9, logprob=-3.0, top_count=1, + top=[SimpleNamespace(token=1, logprob=-.5)]) + result = _format_logprob_steps([9], llm._continuous_logprobs([sample]), True) + entry = result['content'][0] + assert entry['token'] == 'X' and entry['logprob'] == -3.0 + assert len(entry['top_logprobs']) == 1 and entry['top_logprobs'][0]['token_id'] == 1 From 1496d3404f0cc18d26fb51d830991d28cc50cccb Mon Sep 17 00:00:00 2001 From: ajmalrasi Date: Wed, 16 Sep 2026 01:55:04 +0530 Subject: [PATCH 7/7] feat: qualify continuous batching graphs and performance Signed-off-by: ajmalrasi --- cpp/runtime/exec/engineExecutor.cpp | 70 ++++++++++++++--- cpp/runtime/exec/engineExecutor.h | 13 +++- cpp/runtime/greedySchedulerBackend.cpp | 6 +- cpp/runtime/greedySchedulerBackend.h | 2 +- cpp/runtime/llmInferenceRuntime.cpp | 12 ++- cpp/runtime/llmInferenceRuntime.h | 6 +- cpp/runtime/llmRankRuntime.h | 5 ++ cpp/runtime/sequenceStepRuntime.cpp | 59 +++++++++++++- cpp/runtime/sequenceStepRuntime.h | 3 +- cpp/runtime/state/sequencePolicy.cpp | 10 ++- examples/llm/continuousBatchingProbe.cpp | 26 ++++--- experimental/pybind/edgellm_pybind.cpp | 35 ++++++--- experimental/server/api/routes.py | 1 + experimental/server/runtime/engine.py | 82 +++++++++++++------- experimental/server/runtime/engine_client.py | 7 ++ 15 files changed, 262 insertions(+), 75 deletions(-) diff --git a/cpp/runtime/exec/engineExecutor.cpp b/cpp/runtime/exec/engineExecutor.cpp index 7c23a0d67..66557baac 100644 --- a/cpp/runtime/exec/engineExecutor.cpp +++ b/cpp/runtime/exec/engineExecutor.cpp @@ -16,6 +16,7 @@ */ #include "runtime/exec/engineExecutor.h" +#include #include "common/bindingNames.h" #include "common/checkMacros.h" @@ -92,6 +93,10 @@ class TrtEngineExecutor final : public EngineExecutor bool prepare(int32_t profileIndex, InferenceDims const& dims, TensorMap const& map, cudaStream_t stream) override; bool execute(cudaStream_t stream) override; bool captureGraph(cudaStream_t stream) override; + ExecutionStats executionStats() const noexcept override + { + return {mCaptures.load(), mReplays.load(), mEager.load(), mProfileSwitches.load()}; + } int64_t getRequiredContextMemorySize() const override; bool setContextMemory(Tensor& sharedMem) override; int32_t getNumIOTensors() const override; @@ -104,11 +109,13 @@ class TrtEngineExecutor final : public EngineExecutor nvinfer1::ICudaEngine const& getEngine() const noexcept override; private: + std::atomic mCaptures{}, mReplays{}, mEager{}, mProfileSwitches{}; AuxStreamSet mAuxStreams{}; std::unique_ptr mRuntime; std::unique_ptr mEngine; std::unique_ptr mContext; TensorRegistry mRegistry; + std::vector mUnregisteredNames; int32_t mCurrentProfileIndex{-1}; //! A captured CUDA graph together with its binding snapshot for verification. @@ -127,6 +134,7 @@ class TrtEngineExecutor final : public EngineExecutor //! Build a full snapshot of the current binding state. BindingSnapshot snapshotBindings() const; + bool matchesBindings(BindingSnapshot const& snapshot) const; }; TrtEngineExecutor::TrtEngineExecutor(std::filesystem::path const& enginePath, TensorRegistry registry) @@ -163,6 +171,14 @@ TrtEngineExecutor::TrtEngineExecutor(std::filesystem::path const& enginePath, Te {sym(&InferenceDims::skipSoftmaxScaleLen)}}); } + for (int32_t i = 0; i < mEngine->getNbIOTensors(); ++i) + { + std::string name = mEngine->getIOTensorName(i); + if (!mRegistry.contains(name)) + { + mUnregisteredNames.push_back(std::move(name)); + } + } LOG_INFO("engine loaded successfully (%d I/O tensors)", mEngine->getNbIOTensors()); } @@ -240,6 +256,10 @@ bool TrtEngineExecutor::prepare( LOG_ERROR("failed to set optimization profile %d", profileIndex); return false; } + if (mCurrentProfileIndex != profileIndex) + { + ++mProfileSwitches; + } mCurrentProfileIndex = profileIndex; if (!mRegistry.bindAll(mContext.get(), map, dims)) @@ -259,16 +279,10 @@ bool TrtEngineExecutor::prepare( // // LoRA weights are model-dependent and populated into the TensorMap by // LoRAManager::refreshTensorMap() before prepare() is called. - int32_t const numIO = mEngine->getNbIOTensors(); - for (int32_t i = 0; i < numIO; ++i) + for (auto const& binding : mUnregisteredNames) { - char const* name = mEngine->getIOTensorName(i); - if (mRegistry.contains(name)) - { - // Already bound by bindAll above; leave alone. - continue; - } - Tensor* tensor = map.get(name); + char const* name = binding.c_str(); + Tensor* tensor = map.get(binding); if (tensor == nullptr) { LOG_ERROR( @@ -302,18 +316,20 @@ bool TrtEngineExecutor::execute(cudaStream_t stream) auto it = mGraphs.find(hash); if (it != mGraphs.end()) { - BindingSnapshot const current = snapshotBindings(); - if (current == it->second.snapshot) + if (matchesBindings(it->second.snapshot)) { cudaError_t const err = cudaGraphLaunch(it->second.exec, stream); if (err == cudaSuccess) { + ++mReplays; return true; } - LOG_WARNING("cudaGraphLaunch failed (%s), falling back to enqueueV3", cudaGetErrorString(err)); + LOG_ERROR("cudaGraphLaunch failed (%s)", cudaGetErrorString(err)); + return false; } } + ++mEager; return mContext->enqueueV3(stream); } @@ -356,6 +372,7 @@ bool TrtEngineExecutor::captureGraph(cudaStream_t stream) cg.exec = result->second; cg.snapshot = snap; mGraphs[hash] = cg; + ++mCaptures; LOG_INFO("captured graph (hash=0x%zx)", hash); return true; @@ -442,6 +459,7 @@ bool EngineExecutor::BindingSnapshot::operator==(BindingSnapshot const& rhs) con size_t TrtEngineExecutor::computeBindingHash() const { size_t seed = 0; + hash_utils::hashCombine(seed, mCurrentProfileIndex); int32_t const numIO = mEngine->getNbIOTensors(); for (int32_t i = 0; i < numIO; ++i) { @@ -459,6 +477,34 @@ size_t TrtEngineExecutor::computeBindingHash() const return seed; } +bool TrtEngineExecutor::matchesBindings(BindingSnapshot const& snapshot) const +{ + int32_t const count = mEngine->getNbIOTensors(); + if (snapshot.bindings.size() != static_cast(count)) + { + return false; + } + for (int32_t i = 0; i < count; ++i) + { + char const* name = mEngine->getIOTensorName(i); + auto const& expected = snapshot.bindings[i]; + auto const shape = mContext->getTensorShape(name); + if (expected.first != reinterpret_cast(mContext->getTensorAddress(name)) + || expected.second.nbDims != shape.nbDims) + { + return false; + } + for (int32_t d = 0; d < shape.nbDims; ++d) + { + if (expected.second.d[d] != shape.d[d]) + { + return false; + } + } + } + return true; +} + EngineExecutor::BindingSnapshot TrtEngineExecutor::snapshotBindings() const { BindingSnapshot snap; diff --git a/cpp/runtime/exec/engineExecutor.h b/cpp/runtime/exec/engineExecutor.h index 8ee630cfb..fa292567f 100644 --- a/cpp/runtime/exec/engineExecutor.h +++ b/cpp/runtime/exec/engineExecutor.h @@ -139,7 +139,8 @@ class EngineExecutor //! @brief Return a profile shape (min/opt/max) for a named binding. virtual nvinfer1::Dims getProfileShape( - char const* name, int32_t profileIndex, nvinfer1::OptProfileSelector selector) const = 0; + char const* name, int32_t profileIndex, nvinfer1::OptProfileSelector selector) const + = 0; //! @brief Attach a TRT profiler to the execution context. //! @@ -159,6 +160,16 @@ class EngineExecutor bool operator==(BindingSnapshot const& rhs) const noexcept; }; + struct ExecutionStats + { + uint64_t captures{}, replays{}, eager{}, profileSwitches{}; + }; + //! Thread-safe counters for bounded execution-path qualification. + virtual ExecutionStats executionStats() const noexcept + { + return {}; + } + protected: EngineExecutor() = default; }; diff --git a/cpp/runtime/greedySchedulerBackend.cpp b/cpp/runtime/greedySchedulerBackend.cpp index 3c51f39cf..339cf85a0 100644 --- a/cpp/runtime/greedySchedulerBackend.cpp +++ b/cpp/runtime/greedySchedulerBackend.cpp @@ -11,9 +11,9 @@ namespace trt_edgellm { namespace rt { -SamplingSchedulerBackend::SamplingSchedulerBackend( - LLMRankRuntime& runtime, cudaStream_t stream, int32_t vocabularySize, tokenizer::Tokenizer const* tokenizer) - : mSteps(runtime, stream) +SamplingSchedulerBackend::SamplingSchedulerBackend(LLMRankRuntime& runtime, cudaStream_t stream, int32_t vocabularySize, + tokenizer::Tokenizer const* tokenizer, bool captureGraphs) + : mSteps(runtime, stream, captureGraphs) , mStream(stream) , mVocabulary(vocabularySize) , mHostLogits({2, vocabularySize}, DeviceType::kCPU, nvinfer1::DataType::kFLOAT) diff --git a/cpp/runtime/greedySchedulerBackend.h b/cpp/runtime/greedySchedulerBackend.h index a3b1d82fc..86d370242 100644 --- a/cpp/runtime/greedySchedulerBackend.h +++ b/cpp/runtime/greedySchedulerBackend.h @@ -18,7 +18,7 @@ class SamplingSchedulerBackend : public SchedulerBackend { public: SamplingSchedulerBackend(LLMRankRuntime& runtime, cudaStream_t stream, int32_t vocabularySize, - tokenizer::Tokenizer const* tokenizer = nullptr); + tokenizer::Tokenizer const* tokenizer = nullptr, bool captureGraphs = false); void start() override; void invalidate() noexcept override { diff --git a/cpp/runtime/llmInferenceRuntime.cpp b/cpp/runtime/llmInferenceRuntime.cpp index 3f658d697..f8c832831 100644 --- a/cpp/runtime/llmInferenceRuntime.cpp +++ b/cpp/runtime/llmInferenceRuntime.cpp @@ -19,8 +19,8 @@ #include "common/checkMacros.h" #include "common/logger.h" -#include "runtime/llmRankRuntime.h" #include "runtime/greedySchedulerBackend.h" +#include "runtime/llmRankRuntime.h" #include "runtime/multiDevice/runtimeCoordinator.h" #include @@ -233,19 +233,25 @@ std::vector> const& LLMInferenceRuntime::getBaseModelInputT return rootRuntime().getBaseModelInputTokenIds(); } +std::array LLMInferenceRuntime::continuousExecutionStats() const +{ + auto const stats = rootRuntime().executionStats(); + return {stats.captures, stats.replays, stats.eager, stats.profileSwitches}; +} + bool LLMInferenceRuntime::hasDraftModel() const { return rootRuntime().hasDraftModel(); } std::unique_ptr LLMInferenceRuntime::createContinuousScheduler( - cudaStream_t stream, size_t maxQueued, size_t maxQueuedBytes) + cudaStream_t stream, size_t maxQueued, size_t maxQueuedBytes, bool captureGraphs) { ELLM_CHECK(mCoordinator != nullptr, "Runtime coordinator is not initialized."); ELLM_CHECK(!hasDraftModel(), "Continuous scheduling does not support speculative decoding."); auto& runtime = rootRuntime(); auto backend = std::make_unique( - runtime, stream, runtime.vocabularySize(), &runtime.tokenizer()); + runtime, stream, runtime.vocabularySize(), &runtime.tokenizer(), captureGraphs); return std::make_unique(std::move(backend), maxQueued, maxQueuedBytes); } diff --git a/cpp/runtime/llmInferenceRuntime.h b/cpp/runtime/llmInferenceRuntime.h index 0ffb5a460..5b4aec7e7 100644 --- a/cpp/runtime/llmInferenceRuntime.h +++ b/cpp/runtime/llmInferenceRuntime.h @@ -16,12 +16,13 @@ */ #pragma once +#include #include "common/tensor.h" #include "profiling/metrics.h" #include "runtime/config/deploymentConfig.h" -#include "runtime/llmRuntimeUtils.h" #include "runtime/continuousScheduler.h" +#include "runtime/llmRuntimeUtils.h" #include "runtime/modelArtifacts.h" #include "runtime/multiDevice/parallelConfig.h" #include "runtime/preprocess/visualTokenPruner.h" @@ -116,13 +117,14 @@ class LLMInferenceRuntime int32_t getBaseModelPrefillLength() const; std::vector> const& getBaseModelInputTokenIds() const; bool hasDraftModel() const; + std::array continuousExecutionStats() const; //! Create the P5 native scheduler for the deliberately narrow P6 serving //! mode: one local rank, vanilla text generation and no context cache. //! The returned scheduler borrows this runtime and `stream`; callers must //! close/destroy it before this object or stream is released. std::unique_ptr createContinuousScheduler( - cudaStream_t stream, size_t maxQueued = 8, size_t maxQueuedBytes = 256 * 1024); + cudaStream_t stream, size_t maxQueued = 8, size_t maxQueuedBytes = 256 * 1024, bool captureGraphs = false); //! Format and tokenize exactly once using the runtime tokenizer for a P6 //! text request. Media and batched legacy requests are rejected rather diff --git a/cpp/runtime/llmRankRuntime.h b/cpp/runtime/llmRankRuntime.h index a7c5017b4..f87547ee3 100644 --- a/cpp/runtime/llmRankRuntime.h +++ b/cpp/runtime/llmRankRuntime.h @@ -246,6 +246,11 @@ class LLMRankRuntime ELLM_CHECK(mTokenizer != nullptr, "LLMRankRuntime tokenizer is not initialized."); return *mTokenizer; } + auto executionStats() const noexcept + { + return mBaseExecutor->executionStats(); + } + int32_t vocabularySize() const { return mDeployment.base.vocabSize; diff --git a/cpp/runtime/sequenceStepRuntime.cpp b/cpp/runtime/sequenceStepRuntime.cpp index e44bd2e29..3513d6fbc 100644 --- a/cpp/runtime/sequenceStepRuntime.cpp +++ b/cpp/runtime/sequenceStepRuntime.cpp @@ -53,7 +53,7 @@ struct SequenceStepRuntime::View std::map rows; }; -SequenceStepRuntime::SequenceStepRuntime(LLMRankRuntime& runtime, cudaStream_t stream) +SequenceStepRuntime::SequenceStepRuntime(LLMRankRuntime& runtime, cudaStream_t stream, bool captureGraphs) : mRuntime(runtime) , mStream(stream) , mLease(std::make_unique(runtime.mHandleRequestInProgress)) @@ -103,6 +103,63 @@ SequenceStepRuntime::SequenceStepRuntime(LLMRankRuntime& runtime, cudaStream_t s mViews[index] = std::move(view); } CUDA_CHECK(cudaEventCreateWithFlags(&mComplete, cudaEventDisableTiming)); + if (captureGraphs) + { + try + { + captureDecodeViews(); + } + catch (...) + { + mLease->poisoned = true; + cudaStreamSynchronize(mStream); + cudaEventDestroy(mComplete); + mComplete = nullptr; + throw; + } + } +} + +void SequenceStepRuntime::captureDecodeViews() +{ + // Capture warmups mutate hybrid state; only disposable startup sequences may own it here. + SequenceOptions options; + options.maxOutputTokens = 2; + auto first = acquire(1, {0}, options); + auto second = acquire(2, {0}, options); + for (auto handle : {first, second}) + { + beginPrefillChunk(handle); + completeStep(); + acceptToken(handle, 0); + beginDecode({handle, {}}, 1); + completeStep(); + require(mRuntime.mBaseExecutor->captureGraph(mStream), "Selected-slot graph capture failed"); + acceptToken(handle, 0); + } + // Recreate logical inputs while retaining the same physical addresses for the paired view. + finish(first); + finish(second); + release(first); + release(second); + first = acquire(3, {0}, options); + second = acquire(4, {0}, options); + for (auto handle : {first, second}) + { + beginPrefillChunk(handle); + completeStep(); + acceptToken(handle, 0); + } + beginDecode({first, second}, 2); + completeStep(); + require(mRuntime.mBaseExecutor->captureGraph(mStream), "Paired graph capture failed"); + CUDA_CHECK(cudaStreamSynchronize(mStream)); + finish(first); + finish(second); + release(first); + release(second); + mRuntime.zeroRecurrentStates(0, mStream); + mRuntime.zeroRecurrentStates(1, mStream); } SequenceStepRuntime::~SequenceStepRuntime() diff --git a/cpp/runtime/sequenceStepRuntime.h b/cpp/runtime/sequenceStepRuntime.h index eb0f3ab6b..868a3cc17 100644 --- a/cpp/runtime/sequenceStepRuntime.h +++ b/cpp/runtime/sequenceStepRuntime.h @@ -21,7 +21,7 @@ class LLMRankRuntime; class SequenceStepRuntime { public: - SequenceStepRuntime(LLMRankRuntime& runtime, cudaStream_t stream); + SequenceStepRuntime(LLMRankRuntime& runtime, cudaStream_t stream, bool captureGraphs = false); ~SequenceStepRuntime(); SequenceStepRuntime(SequenceStepRuntime const&) = delete; SequenceStepRuntime& operator=(SequenceStepRuntime const&) = delete; @@ -48,6 +48,7 @@ class SequenceStepRuntime struct Lease; struct View; void requireIdle() const; + void captureDecodeViews(); Tensor const& enqueue(std::array const& handles, int32_t count, int32_t span, bool decode); LLMRankRuntime& mRuntime; cudaStream_t mStream; diff --git a/cpp/runtime/state/sequencePolicy.cpp b/cpp/runtime/state/sequencePolicy.cpp index 5daa37f49..6df5531ed 100644 --- a/cpp/runtime/state/sequencePolicy.cpp +++ b/cpp/runtime/state/sequencePolicy.cpp @@ -80,8 +80,12 @@ SequenceSample SequenceSampler::sample(float const* logits, SequenceOptions cons } auto better = [](Entry const& a, Entry const& b) { return a.logit > b.logit || (a.logit == b.logit && a.token < b.token); }; - // std::sort uses stack storage; stable_sort may allocate on every token. - std::sort(mEntries.begin(), mEntries.end(), better); + bool const greedy = o.temperature <= 1e-3F || o.topK == 1 + || (o.topK <= 1 && o.topP >= 1.0F - 1e-6F && std::fabs(o.temperature - 1.0F) <= 1e-3F); + // Logprob requests retain the original summation order for exact policy compatibility. + size_t const needed + = o.numLogprobs > 0 || (!greedy && o.topK == 0) ? mEntries.size() : (greedy ? 1 : static_cast(o.topK)); + std::partial_sort(mEntries.begin(), mEntries.begin() + needed, mEntries.end(), better); double const maximum = mEntries.front().logit; double normalizer = 0; for (auto const& entry : mEntries) @@ -97,8 +101,6 @@ SequenceSample SequenceSampler::sample(float const* logits, SequenceOptions cons out.top[i] = {mEntries[i].token, mEntries[i].logit - logNormalizer}; } size_t selected = 0; - bool const greedy = o.temperature <= 1e-3F || o.topK == 1 - || (o.topK <= 1 && o.topP >= 1.0F - 1e-6F && std::fabs(o.temperature - 1.0F) <= 1e-3F); if (!greedy) { if (counter == std::numeric_limits::max()) diff --git a/examples/llm/continuousBatchingProbe.cpp b/examples/llm/continuousBatchingProbe.cpp index 2f18b3470..a608ffe89 100644 --- a/examples/llm/continuousBatchingProbe.cpp +++ b/examples/llm/continuousBatchingProbe.cpp @@ -488,7 +488,7 @@ namespace trt_edgellm namespace rt { //! Validate the production step interfaces against the independent P1/legacy path. -bool runStepTests(LLMRankRuntime& runtime, tokenizer::Tokenizer& tokenizer, cudaStream_t stream) +bool runStepTests(LLMRankRuntime& runtime, tokenizer::Tokenizer& tokenizer, cudaStream_t stream, bool graphs = false) { auto const words = tokenizer.encode("A red fox crosses a blue river. One two three four. "); require(!words.empty(), "Missing fixture tokens"); @@ -522,7 +522,7 @@ bool runStepTests(LLMRankRuntime& runtime, tokenizer::Tokenizer& tokenizer, cuda { ContinuousBatchingProbe observer(runtime, stream); auto const samplerBefore = observer.samplingBytes(); - SequenceStepRuntime steps(runtime, stream); + SequenceStepRuntime steps(runtime, stream, graphs); auto rejects = [&](auto operation, char const* label) { bool rejected = false; try @@ -643,6 +643,13 @@ bool runStepTests(LLMRankRuntime& runtime, tokenizer::Tokenizer& tokenizer, cuda response.outputIds.size() == 1 && response.outputIds[0] == std::vector{greedy(firstA), greedy(nextA)}, "Legacy output changed after releasing the step lease"); std::cout << "LEGACY_RESTORED passed=1" << std::endl; + if (graphs) + { + auto const stats = runtime.executionStats(); + passed = passed && stats.captures == 3 && stats.replays >= 4; + std::cout << "P7_GRAPH_GATE passed=" << passed << " captures=" << stats.captures << " replays=" << stats.replays + << " eager=" << stats.eager << " profile_switches=" << stats.profileSwitches << std::endl; + } std::cout << "P2_STATE_GATE passed=" << passed << " full_chunk_quality=separate_P1_TAIL_BASELINE" << std::endl; return passed; } @@ -1220,11 +1227,11 @@ int main(int argc, char** argv) if (argc != 3 && (argc != 4 || (std::string(argv[3]) != "--policies" && std::string(argv[3]) != "--scheduler" - && std::string(argv[3]) != "--steps" + && std::string(argv[3]) != "--steps" && std::string(argv[3]) != "--graphs" && (std::string(argv[3]) != "--chunks" && std::string(argv[3]) != "--chunks-extra")))) { std::cerr << "Usage: continuous_batching_probe ENGINE_DIR CHECKPOINT_DIR " - "[--steps|--chunks|--chunks-extra|--scheduler|--policies]\n"; + "[--steps|--graphs|--chunks|--chunks-extra|--scheduler|--policies]\n"; return 2; } cudaStream_t stream{}; @@ -1243,11 +1250,12 @@ int main(int argc, char** argv) ? trt_edgellm::rt::runPolicyTests(runtime, tokenizer, stream) : argc == 4 && std::string(argv[3]) == "--scheduler" ? trt_edgellm::rt::runSchedulerTests(runtime, tokenizer, stream) - : argc == 4 ? (std::string(argv[3]).find("--chunks") == 0 - ? trt_edgellm::rt::runChunkTests( - runtime, tokenizer, stream, std::string(argv[3]) == "--chunks-extra") - : trt_edgellm::rt::runStepTests(runtime, tokenizer, stream)) - : trt_edgellm::rt::runProbe(runtime, tokenizer, stream); + : argc == 4 + ? (std::string(argv[3]).find("--chunks") == 0 ? trt_edgellm::rt::runChunkTests(runtime, tokenizer, + stream, std::string(argv[3]) == "--chunks-extra") + : trt_edgellm::rt::runStepTests(runtime, tokenizer, + stream, std::string(argv[3]) == "--graphs")) + : trt_edgellm::rt::runProbe(runtime, tokenizer, stream); } CUDA_CHECK(cudaStreamDestroy(stream)); std::cout << "PROBE_RESULT passed=" << passed << std::endl; diff --git a/experimental/pybind/edgellm_pybind.cpp b/experimental/pybind/edgellm_pybind.cpp index 3beb9cef9..73cb95702 100644 --- a/experimental/pybind/edgellm_pybind.cpp +++ b/experimental/pybind/edgellm_pybind.cpp @@ -33,9 +33,9 @@ #include "profiling/metrics.h" #include "runtime/audioLoader.h" #include "runtime/audioUtils.h" +#include "runtime/continuousScheduler.h" #include "runtime/imageUtils.h" #include "runtime/llmInferenceRuntime.h" -#include "runtime/continuousScheduler.h" #include "runtime/llmRuntimeUtils.h" #include "runtime/melSpectrogram.h" #ifdef EDGELLM_ENABLE_NEMOTRON_ASR @@ -275,10 +275,10 @@ class PyLLMRuntime return response; } - void enableContinuousBatching(size_t maxQueued, size_t maxQueuedBytes) + void enableContinuousBatching(size_t maxQueued, size_t maxQueuedBytes, bool captureGraphs) { ELLM_CHECK(mScheduler == nullptr, "Continuous scheduler is already enabled"); - mScheduler = mRuntime->createContinuousScheduler(mStream.get(), maxQueued, maxQueuedBytes); + mScheduler = mRuntime->createContinuousScheduler(mStream.get(), maxQueued, maxQueuedBytes, captureGraphs); } std::vector prepareContinuousPrompt(LLMGenerationRequest const& request) const @@ -298,6 +298,11 @@ class PyLLMRuntime return py::bytes(mRuntime->continuousTokenPiece(tokenId)); } + std::array continuousExecutionStats() const + { + return mRuntime->continuousExecutionStats(); + } + bool continuousHealthy() const { return mScheduler != nullptr && mScheduler->healthy(); @@ -1022,19 +1027,24 @@ PYBIND11_MODULE(_edgellm_runtime, m) .def_readwrite("generation", &SchedulerRequestOptions::generation) .def_readwrite("stream_records", &SchedulerRequestOptions::streamRecords) .def_readwrite("stream_bytes", &SchedulerRequestOptions::streamBytes) - .def("set_queue_timeout_ms", [](SchedulerRequestOptions& self, int64_t timeoutMs) { - self.queueDeadline = std::chrono::steady_clock::now() + std::chrono::milliseconds{timeoutMs}; - }) + .def("set_queue_timeout_ms", + [](SchedulerRequestOptions& self, int64_t timeoutMs) { + self.queueDeadline = std::chrono::steady_clock::now() + std::chrono::milliseconds{timeoutMs}; + }) .def("set_timeout_ms", [](SchedulerRequestOptions& self, int64_t timeoutMs) { self.deadline = std::chrono::steady_clock::now() + std::chrono::milliseconds{timeoutMs}; }); py::class_(m, "ContinuousTicket") .def("id", &SchedulerTicket::id) .def("cancel", &SchedulerTicket::cancel) - .def("read", [](SchedulerTicket const& self, int64_t timeoutMs) { - return self.read(std::chrono::milliseconds{timeoutMs}); - }, py::arg("timeout_ms"), py::call_guard()) - .def("result", [](SchedulerTicket const& self) { return self.result().get(); }, + .def( + "read", + [](SchedulerTicket const& self, int64_t timeoutMs) { + return self.read(std::chrono::milliseconds{timeoutMs}); + }, + py::arg("timeout_ms"), py::call_guard()) + .def( + "result", [](SchedulerTicket const& self) { return self.result().get(); }, py::call_guard()); // ======================================================================== @@ -1056,14 +1066,15 @@ PYBIND11_MODULE(_edgellm_runtime, m) py::arg("dflash_block_size") = 0, "Construct for speculative decoding") .def("handle_request", &PyLLMRuntime::handleRequest, py::arg("request"), py::call_guard(), "Process a generation request and return the response") - .def("enable_continuous_batching", &PyLLMRuntime::enableContinuousBatching, - py::arg("max_queued") = 8, py::arg("max_queued_bytes") = 256 * 1024, + .def("enable_continuous_batching", &PyLLMRuntime::enableContinuousBatching, py::arg("max_queued") = 8, + py::arg("max_queued_bytes") = 256 * 1024, py::arg("capture_graphs") = false, py::call_guard()) .def("prepare_continuous_prompt", &PyLLMRuntime::prepareContinuousPrompt, py::arg("request"), py::call_guard()) .def("submit_continuous", &PyLLMRuntime::submitContinuous, py::arg("prompt"), py::arg("options"), py::call_guard()) .def("continuous_token_piece", &PyLLMRuntime::continuousTokenPiece, py::arg("token_id")) + .def("continuous_execution_stats", &PyLLMRuntime::continuousExecutionStats) .def("continuous_healthy", &PyLLMRuntime::continuousHealthy) .def("continuous_queued_requests", &PyLLMRuntime::continuousQueuedRequests) .def("continuous_resident_requests", &PyLLMRuntime::continuousResidentRequests) diff --git a/experimental/server/api/routes.py b/experimental/server/api/routes.py index 81216ed59..12844d69b 100644 --- a/experimental/server/api/routes.py +++ b/experimental/server/api/routes.py @@ -50,6 +50,7 @@ async def health(request: Request): return { "status": "healthy", "model": client.model_name, + "execution": client.execution_stats, "active_requests": client.active_requests, "queued_requests": client.queued_requests, "capabilities": { diff --git a/experimental/server/runtime/engine.py b/experimental/server/runtime/engine.py index a519677ea..b7ee00733 100644 --- a/experimental/server/runtime/engine.py +++ b/experimental/server/runtime/engine.py @@ -859,12 +859,16 @@ def enable_continuous_batching(self) -> None: if self._continuous_batching: return if self.has_draft_model or self._max_batch_size < 2: - raise ValueError("continuous batching requires a vanilla batch-two engine") + raise ValueError( + "continuous batching requires a vanilla batch-two engine") if self._layout.visual_dir or self._layout.audio_dir: - raise ValueError("continuous batching currently supports text-only engines") + raise ValueError( + "continuous batching currently supports text-only engines") if self._context_cache_config.enabled: - raise ValueError("continuous batching does not support context cache") - self._runtime.enable_continuous_batching() + raise ValueError( + "continuous batching does not support context cache") + self._runtime.enable_continuous_batching(capture_graphs=os.environ.get( + "EDGELLM_CONTINUOUS_GRAPHS", "0") == "1") self._continuous_batching = True def _continuous_options(self, params: SamplingParams, *, stream: bool): @@ -880,11 +884,13 @@ def _continuous_options(self, params: SamplingParams, *, stream: bool): generation.stop_strings = params.stop generation.logit_bias = _normalize_logit_bias(params.logit_bias) options.generation = generation - options.stream_records = min(max(params.max_tokens + 2, 2), 64) if stream else 0 + options.stream_records = min(max(params.max_tokens + + 2, 2), 64) if stream else 0 options.stream_bytes = 65536 if stream else 0 return options - def _submit_continuous(self, request, params: SamplingParams, *, stream: bool): + def _submit_continuous(self, request, params: SamplingParams, *, + stream: bool): self._ensure_open() prompt = self._runtime.prepare_continuous_prompt(request) return self._runtime.submit_continuous( @@ -896,40 +902,58 @@ def _continuous_logprobs(self, samples) -> List[List[LogprobEntry]]: entries = [] for entry in list(sample.top)[:sample.top_count]: piece = self._runtime.continuous_token_piece(entry.token) - entries.append(LogprobEntry(entry.token, entry.logprob, - piece.decode("utf-8", "replace"), - list(piece))) + entries.append( + LogprobEntry(entry.token, entry.logprob, + piece.decode("utf-8", "replace"), + list(piece))) if not any(entry.token_id == sample.token for entry in entries): piece = self._runtime.continuous_token_piece(sample.token) - entries.append(LogprobEntry(sample.token, sample.logprob, - piece.decode("utf-8", "replace"), - list(piece), chosen_only=True)) + entries.append( + LogprobEntry(sample.token, + sample.logprob, + piece.decode("utf-8", "replace"), + list(piece), + chosen_only=True)) converted.append(entries) return converted - def _complete_continuous_request(self, request, params: SamplingParams, - tool_config: ToolConfig, *, tool_parser: str, - reasoning_parser: str, ticket=None) -> CompletionOutput: + def _complete_continuous_request(self, + request, + params: SamplingParams, + tool_config: ToolConfig, + *, + tool_parser: str, + reasoning_parser: str, + ticket=None) -> CompletionOutput: if ticket is None: ticket = self._submit_continuous(request, params, stream=False) result = ticket.result() if result.status == self._rt.SchedulerStatus.DEADLINE: raise ServerOverloadedError("native request deadline expired") if result.status != self._rt.SchedulerStatus.COMPLETED: - raise RuntimeError(f"continuous generation failed: {result.status}") + raise RuntimeError( + f"continuous generation failed: {result.status}") finish_reason = { self._rt.SequenceFinish.LENGTH: "length", self._rt.SequenceFinish.EOS: "stop", self._rt.SequenceFinish.STOP: "stop", }.get(result.finish, "stop") output = self._parse_generation_output( - result.text, list(result.tokens), result.prompt_tokens, finish_reason, tool_config, - tool_parser=tool_parser, reasoning_parser=reasoning_parser) + result.text, + list(result.tokens), + result.prompt_tokens, + finish_reason, + tool_config, + tool_parser=tool_parser, + reasoning_parser=reasoning_parser) if params.num_logprobs: output.logprobs = self._continuous_logprobs(result.logprobs) return output - def generate_continuous_stream(self, request, params: SamplingParams, ticket=None) -> Iterator[StreamDelta]: + def generate_continuous_stream(self, + request, + params: SamplingParams, + ticket=None) -> Iterator[StreamDelta]: """Read one P5 ticket channel; closing this iterator cancels only it.""" state = {"ticket": ticket} @@ -949,23 +973,29 @@ def _iterate(): read = ticket.read(timeout_ms=200) if read.update is not None: update = read.update - yield StreamDelta(text=update.text, - token_ids=([update.sample.token] if update.sample.token >= 0 else []), - logprobs=(self._continuous_logprobs([update.sample]) - if params.num_logprobs and update.sample.token >= 0 else [])) + yield StreamDelta( + text=update.text, + token_ids=([update.sample.token] + if update.sample.token >= 0 else []), + logprobs=(self._continuous_logprobs( + [update.sample]) if params.num_logprobs + and update.sample.token >= 0 else [])) if read.closed: terminal = ticket.result() if terminal.status == self._rt.SchedulerStatus.DEADLINE: - raise ServerOverloadedError("native request deadline expired") + raise ServerOverloadedError( + "native request deadline expired") if terminal.status != self._rt.SchedulerStatus.COMPLETED: raise RuntimeError( - f"continuous generation failed: {terminal.status}") + f"continuous generation failed: {terminal.status}" + ) reason = { self._rt.SequenceFinish.LENGTH: "length", self._rt.SequenceFinish.EOS: "stop", self._rt.SequenceFinish.STOP: "stop", }.get(terminal.finish, "stop") - yield StreamDelta(finished=True, finish_reason=reason, + yield StreamDelta(finished=True, + finish_reason=reason, prompt_tokens=terminal.prompt_tokens) return finally: diff --git a/experimental/server/runtime/engine_client.py b/experimental/server/runtime/engine_client.py index 570d2b70f..d83252908 100644 --- a/experimental/server/runtime/engine_client.py +++ b/experimental/server/runtime/engine_client.py @@ -285,6 +285,13 @@ def model_name(self) -> str: def capabilities(self) -> EngineCapabilities: return self._capabilities + @property + def execution_stats(self) -> Dict[str, int]: + if getattr(self._llm, "continuous_batching_enabled", False): + values = self._llm._runtime.continuous_execution_stats() + return dict(zip(("captures", "replays", "eager", "profile_switches"), values)) + return {} + @property def healthy(self) -> bool: if self._closed: