Visual Token Pruner#
-
class VisualTokenPruner#
Abstract prefill-time visual-token pruner.
The non-virtual pruneForPrefill() / pruneBatchForPrefill() own everything algorithms must not diverge on: the enablement guards (minVisualTokens, target-is-a-reduction), modality partitioning, and the target computation. Every subclass that shortens a request must finish by calling compactToKeepList(), which records the ordered slot-local subset and owns the buffer compaction: all text tokens are always kept; kept tokens retain their original absolute RoPE positions; decode continues at the position after the unpruned sequence (matching the HF DART reference).
Batching: prune() is always invoked with a single-slot view (PruneRequest.embeds is that slot’s contiguous [len, hidden] plane; all positions are slot-local), so algorithms are batch-agnostic. In the batched flow compactToKeepList() records the slot’s keep list instead of compacting immediately; once every slot is selected, the base repacks all batch planes to the new (pruned) row pitch in one pass. Consequently batched pruning requires the algorithm to route its result through compactToKeepList(). Direct rewrites and token merging are unsupported by the MR1 protocol.
Instances are created through createVisualTokenPruner() and reused across requests, so per-request device work buffers should be preallocated in the constructor. The caller is responsible for the runtime gates (fresh KV cache, mRoPE engine, …) and for shrinking the context lengths to the returned pruned lengths.
Subclassed by trt_edgellm::rt::DartPruner
Public Functions
-
virtual ~VisualTokenPruner() = default#
-
virtual char const *name() const noexcept = 0#
Algorithm name (matches the registry key).
- int32_t pruneForPrefill(
- std::vector<int32_t> const &hostTokenIds,
- PipelineIO &io,
- int32_t origLen,
- cudaStream_t stream
Prune the assembled prefill inputs for a batch-1 request.
May synchronize
stream(algorithm-dependent); on return the compaction kernels are enqueued onstreamand the PipelineIO tensors are reshaped to the pruned length.- Parameters:
hostTokenIds – The request’s expanded token ids (image placeholders already inserted); visual positions are
tokenId == imageTokenId.io – Pipeline buffers;
inputsEmbedsmust currently be [1, origLen, hiddenSize].origLen – The unpruned prefill length (== hostTokenIds.size()).
stream – CUDA stream all device work runs on.
- Returns:
The pruned length P (< origLen), or origLen when pruning is skipped (no/too-few visual tokens, or the target keep count is not a reduction).
- int32_t pruneBatchForPrefill(
- std::vector<std::vector<int32_t>> const &hostTokenIds,
- PipelineIO &io,
- std::vector<int32_t> &effectiveLens,
- int32_t maxLen,
- std::vector<int32_t> &prunedTokensOut,
- cudaStream_t stream
Prune the assembled prefill inputs for a batch of requests. The batch size is taken from
io.inputsEmbedsshape [batch, maxLen, hiddenSize].Selects per slot (slots without enough visual tokens are left unpruned), then repacks every batch plane of
inputsEmbeds/deepstackEmbedsto the new maximum length and gathers each pruned slot’s mRoPE rows.effectiveLensis updated in place to the per-slot pruned lengths andprunedTokensOutreceives the per-slot removed counts. ReturnsmaxLenunchanged when nothing was pruned, or when any slot’s token ids don’t cover its effective length (chunked continuation — modality partitioning is impossible).- Parameters:
hostTokenIds – Per-slot expanded token ids (size >= batch; slot i must satisfy hostTokenIds[i].size() == effectiveLens[i]).
io – Pipeline buffers;
inputsEmbedsmust currently be [batch, maxLen, hiddenSize].effectiveLens – Per-slot prefill lengths (size >= batch); updated in place.
maxLen – Current padded prefill length (== max of effectiveLens).
prunedTokensOut – Per-slot number of removed tokens (resized to batch).
stream – CUDA stream all device work runs on.
- Returns:
The new padded prefill length (== max of the updated effectiveLens).
- void compactAuxiliaryInputs(
- std::vector<std::vector<int32_t>> &hostTokenIds,
- int32_t batch,
- std::vector<int32_t> const &prunedTokens,
- OptionalInputTensor &visualFeatures,
- cudaStream_t stream
Compact the auxiliary prompt inputs that downstream consumers use to re-embed the prompt — required when the pruned request continues into a speculative-decoding draft prefill, which re-embeds host token ids and re-inserts the raw visual feature rows. (Deepstack features are not compacted: no draft strategy consumes them.)
Must be called right after a pruning pass on the same request. Per slot (using the keep lists recorded by that pass): compacts
hostTokenIds[i]in place, and gathers the kept visual feature rows into a pruner-owned buffer, rebinding the reference. Feature rows are indexed by the running image-token count over the packed batch grid, so the compacted grid’s k-th image token is served by the original ordinal of the k-th kept one. This assumes ordinal 0 is the batch’s first image token, i.e. embedding runs with zero multimodal base offsets — guaranteed by the fresh-KV-cache gate on pruning (prefix reuse would start the running count at a nonzero base offset). No-op when nothing was pruned.- Parameters:
hostTokenIds – Per-slot expanded token ids (compacted in place).
batch – Number of active slots.
prunedTokens – Per-slot removed counts from the pruning pass (validated against the recorded keep lists).
visualFeatures – Raw visual feature rows, one row per image token in the request ([totalImageTokens, dim] — Qwen-VL packing); rebound on return.
stream – CUDA stream the gathers run on.
-
void preallocateAuxiliaryBuffers()#
Preallocate the compactAuxiliaryInputs() buffers to their upper bound (feature rows are bounded by maxBatchSize x maxSupportedInputLength), so no allocation happens on the prefill path. Call once at setup when the deployment will use auxiliary compaction (spec decode); without this call the buffers grow lazily on first use instead — the feature plane is too large to always reserve for deployments that never need it.
- inline std::vector<int32_t> const &keepStartOffsets(
- inline std::vector<int32_t> const &concatenatedKeepIndices(
-
inline VisualPrunerConfig const &config() const noexcept#
-
virtual ~VisualTokenPruner() = default#
-
struct VisualPrunerConfig#
Visual-token pruning configuration (runtime-side; nothing is required in the exported engine config — the prune operates on runtime buffers only).
Public Members
-
bool enabled = {false}#
-
std::string algorithm = {"dart"}#
Pruning algorithm name, resolved through the pruner registry. Built-in: “dart” (default, duplication-aware; paper: “Stop Looking for Important Tokens
in Multimodal Language Models: Duplication Matters More”).
-
float reductionRatio = {0.25F}#
Fraction of visual tokens to remove (0.25 = keep 75%).
-
int32_t minVisualTokens = {16}#
Skip pruning when the request has fewer visual tokens than this (accuracy and break-even guard: tiny images gain nothing from pruning).
-
int32_t pivotImageTokens = {4}#
-
int32_t pivotTextTokens = {4}#
-
bool enabled = {false}#
-
struct ImageSpan#
One contiguous run of visual tokens (one image, or one video-frame block) with its per-image retention quota. Pruning within each span independently — rather than over one global candidate pool — guarantees every image keeps its proportional share of tokens (at least one), so a low-information image can never be starved by the others.
-
struct PruneRequest#
Guard-checked, modality-partitioned view of one batch-1 prefill request, prepared by VisualTokenPruner::pruneForPrefill and handed to the algorithm hook.
Public Members
-
Tensor const *embeds = {nullptr}#
Non-owning [origLen, hiddenSize] FP16 GPU view of the assembled input embeddings.
-
std::vector<int32_t> const *imagePositions = {nullptr}#
Positions of visual tokens in the sequence (ascending, non-empty).
-
std::vector<int32_t> const *textPositions = {nullptr}#
Positions of all non-visual tokens in the sequence (ascending).
-
std::vector<ImageSpan> const *imageSpans = {nullptr}#
Contiguous visual spans (one per image) with per-span retention quotas.
-
int32_t targetImageTokens = {0}#
Total number of visual tokens to retain — the sum of the per-span quotas (1 <= targetImageTokens < imagePositions->size()).
-
int32_t origLen = {0}#
Unpruned prefill length (== imagePositions->size() + textPositions->size()).
-
Tensor const *embeds = {nullptr}#
- void trt_edgellm::rt::registerVisualPruner(
- std::string const &name,
- VisualPrunerFactory factory
Register a pruning algorithm under
name(case-sensitive). The built-in “dart” is registered automatically; call this to plug in a custom algorithm before constructing the runtime. Re-registering a name replaces the previous factory.
- std::unique_ptr<VisualTokenPruner> trt_edgellm::rt::createVisualTokenPruner(
- VisualPrunerConfig const &config,
- LLMEngineConfig const &engineConfig
Instantiate the pruner named by
config.algorithm.- Throws:
std::runtime_error – if the name is not registered or the config is invalid.
- std::vector<std::string> trt_edgellm::rt::registeredVisualPrunerNames(
Names of all registered pruning algorithms (for CLI help / validation).