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(
) const noexcept#
inline virtual DecodingKvHeadroom requiredKvHeadroom() const#
inline virtual DecodingTokenStateContract tokenStateContract(
) const noexcept#
virtual bool decodeStep(DecodingInferenceContext &context) = 0#
virtual bool captureCudaGraphs(cudaStream_t stream) = 0#
inline virtual bool initializeForGeneration(
DecodingInferenceContext&
)#

Initialize decoder-private generation state after base prefill. Non-speculative strategies are no-ops.

inline virtual std::vector<int32_t> const &commonMaterializedStateLengths(
) const noexcept#

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 void setContextMemory(Tensor&) = 0#
virtual bool hasSystemPromptKVCache(
SystemPromptCacheKey const&
) const = 0#
virtual void restoreSystemPromptKVCache(
SystemPromptCacheKey const&,
int32_t residentSlot,
cudaStream_t
) = 0#
virtual bool runSystemPromptPrefill(DecodingInferenceContext&) = 0#
virtual void saveSystemPromptKVCache(
SystemPromptCacheKey const&,
std::string const&,
std::vector<tokenizer::Rank> const&,
int32_t,
cudaStream_t
) = 0#
virtual void resetForNewSequences(Tensor&, cudaStream_t) = 0#
virtual void onBatchEvict(
std::vector<int32_t> const&,
int32_t,
int32_t,
Tensor&,
cudaStream_t
) = 0#
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.

struct SamplingBuffers#

Public Members

Tensor &workspace#
Tensor &indices#
Tensor &scores#
Tensor &baseVocabMappingTable#
Tensor &hostPackedTokenIds#
Tensor &hostSelectedTokenIds#
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 &uniforms#
Tensor &hostUniforms#
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

Tensor &deviceLogprobsValues#

GPU [logprobsMaxBatch, kMaxLogprobsK] top-K log-prob values.

Tensor &deviceLogprobsIndices#

GPU [logprobsMaxBatch, kMaxLogprobsK] top-K token indices.

Tensor &hostLogprobsValues#

CPU pinned [logprobsMaxBatch, kMaxLogprobsK].

Tensor &hostLogprobsIndices#

CPU pinned [logprobsMaxBatch, kMaxLogprobsK]

Tensor &gatheredLogits#

GPU [maxBatch * maxAcceptDepth, vocab] accepted verify rows gathered before extraction. Used by the spec-decode verify paths whose accepted rows are non-contiguous in the output logits (EAGLE / MTP / DFlash / JetSpec); Gemma4 MTP’s sequential chain reads logits directly.

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#
TensorMap &tensorMap#
SharedResources &sharedResources#
HybridCacheManager &cacheManager#
PipelineIO &pipelineIO#
std::function<bool(InferenceDims const&, cudaStream_t)> captureGraph#
struct PreprocessResources#

Preprocessing resources: embedding lookup, step preparation, deepstack.

Public Members

StepPreparer &stepPreparer#
EmbeddingPreprocessor &embeddingPreprocessor#
EmbeddingData &embedding#
Tensor &idsInput#
DeepstackBinding *deepstack#
Gemma4EmbeddingPreprocessor *gemma4Ple#
struct DecodingRuntimeContext#

Public Members

DeploymentConfig &deployment#
int32_t maxRuntimeBatchSize#
std::filesystem::path const &checkpointDir#
std::filesystem::path const &draftCheckpointDir#
BaseEngineResources base#
PreprocessResources preprocess#
tokenizer::Tokenizer &tokenizer#
LogitBias &logitBias#
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}#