Decoding Strategy#
-
class DecodingStrategy#
Subclassed by trt_edgellm::rt::BlockDiffusionDecoder, trt_edgellm::rt::DFlashDecoder, trt_edgellm::rt::DSparkDecoder, trt_edgellm::rt::EagleDecoder, trt_edgellm::rt::Gemma4MTPDecoder, trt_edgellm::rt::MTPDecoder, trt_edgellm::rt::VanillaDecoder
Public Functions
-
virtual ~DecodingStrategy() noexcept = default#
-
virtual DecodingStrategyKind kind() const noexcept = 0#
-
virtual char const *name() const noexcept = 0#
-
virtual bool isSpeculative() const noexcept = 0#
- inline virtual DecodingStrategyCapabilities capabilities(
-
inline virtual DecodingKvHeadroom requiredKvHeadroom() const#
- inline virtual DecodingTokenStateContract tokenStateContract(
-
virtual bool decodeStep(DecodingInferenceContext &context) = 0#
-
virtual bool captureCudaGraphs(cudaStream_t stream) = 0#
- inline virtual bool initializeForGeneration( )#
Initialize decoder-private generation state after base prefill. Non-speculative strategies are no-ops.
- inline virtual std::vector<int32_t> const &commonMaterializedStateLengths(
Greatest per-slot logical prefix whose continuation state is materialized by every model in this strategy. Physical model-state tails may extend beyond this boundary. This reports decoding progress only; context-cache policy decides whether that prefix can be published.
-
virtual int64_t getRequiredContextMemorySize() const noexcept = 0#
- virtual bool hasSystemPromptKVCache(
- SystemPromptCacheKey const&
- virtual void restoreSystemPromptKVCache(
- SystemPromptCacheKey const&,
- int32_t residentSlot,
- cudaStream_t
-
virtual bool runSystemPromptPrefill(DecodingInferenceContext&) = 0#
- virtual void saveSystemPromptKVCache(
- SystemPromptCacheKey const&,
- std::string const&,
- std::vector<tokenizer::Rank> const&,
- int32_t,
- cudaStream_t
-
virtual ~DecodingStrategy() noexcept = default#
-
struct DecodingStrategyCapabilities#
Public Members
-
bool ownsBaseVerificationCudaGraphs = {false}#
-
bool supportsLosslessSampling = {false}#
-
int32_t maxSamplingSupport = {0}#
0 when the decoder does not require a bounded sampling support.
-
bool fallbackToVanillaForNonGreedySampling = {false}#
Preserve request sampling semantics instead of coercing a greedy-only strategy.
-
bool requiresDefaultDecoderCudaGraphs = {false}#
Capture vanilla graphs when request routing may select the default decoder.
-
bool ownsBaseVerificationCudaGraphs = {false}#
-
struct SamplingBuffers#
Public Members
-
Tensor &hostOutputSpaceIds#
Sampled indices captured before mapReducedVocabToFullVocab, i.e. still in the engine’s output vocabulary. Grammar matchers live in that space (see GuidedDecoder), so they must be advanced with these rather than the remapped full IDs.
-
Tensor &hostOutputSpaceIds#
-
struct LogprobsBuffers#
GPU/CPU buffers for logprobs computation (owned by LLMInferenceRuntime). Referenced here so decoders can call log-softmax + top-K on the output logits.
Public Members
-
struct BaseEngineResources#
Base-engine execution infrastructure: executor, tensor map, KV cache, pipeline I/O, shared resources, and CUDA-graph capture callback.
Public Members
-
EngineExecutor &executor#
-
HybridCacheManager &cacheManager#
-
PipelineIO &pipelineIO#
-
std::function<bool(InferenceDims const&, cudaStream_t)> captureGraph#
-
EngineExecutor &executor#
-
struct PreprocessResources#
Preprocessing resources: embedding lookup, step preparation, deepstack.
Public Members
-
StepPreparer &stepPreparer#
-
EmbeddingPreprocessor &embeddingPreprocessor#
-
EmbeddingData &embedding#
-
DeepstackBinding *deepstack#
-
Gemma4EmbeddingPreprocessor *gemma4Ple#
-
StepPreparer &stepPreparer#
-
struct DecodingRuntimeContext#
Public Members
-
DeploymentConfig &deployment#
-
int32_t maxRuntimeBatchSize#
-
std::filesystem::path const &checkpointDir#
-
std::filesystem::path const &draftCheckpointDir#
-
BaseEngineResources base#
-
PreprocessResources preprocess#
-
GuidedDecoder &guidedDecoder#
-
SamplingBuffers sampling#
-
LogprobsBuffers logprobs#
-
std::function<bool(void *buffer, int32_t count, cudaStream_t stream)> tokenBroadcast = {}#
Optional sampled-token synchronization callback. Rank 0 sends; peer ranks receive.
-
int32_t parallelRank = {-1}#
-
DeploymentConfig &deployment#