Skip to content

Commit 0a8b0c9

Browse files
authored
feat: add tensor-rt inference engine support (vllm-project#2000)
Signed-off-by: varungupta <varungup90@gmail.com>
1 parent 6ec2a8e commit 0a8b0c9

5 files changed

Lines changed: 489 additions & 48 deletions

File tree

pkg/plugins/gateway/algorithms/pd_disaggregation.go

Lines changed: 114 additions & 45 deletions
Original file line numberDiff line numberDiff line change
@@ -46,6 +46,7 @@ const (
4646
RouterPD types.RoutingAlgorithm = "pd"
4747
VLLMEngine string = "vllm"
4848
SGLangEngine string = "sglang"
49+
TensorRTLLM string = "trtllm"
4950
SGLangBootstrapPort int64 = 8998
5051
SGLangBootstrapPortIdentifier string = "model.aibrix.ai/sglang-bootstrap-port"
5152
LLMEngineIdentifier string = constants.ModelLabelEngine
@@ -62,11 +63,10 @@ const (
6263
defaultRequestRateHighLoadThreshold = 1.0
6364
defaultRequestRateLowLoadThreshold = 0.25
6465

65-
pdRouteValidateLLMEngineFail = "pd-validate-llm-engine-fail"
66-
pdRouteFilterPrefillDecodePodsFail = "pd-filter-prefill-decode-pods-fail"
67-
pdRoutePrefillRequestError = "pd-do-prefill-request-error"
68-
pdRoutePrefillRequestSuccess = "pd-prefill-request-success"
69-
pdRoutePrefillEmptyKVTransferParams = "pd-prefill-empty-kv-transfer-params"
66+
pdRouteValidateLLMEngineFail = "pd-validate-llm-engine-fail"
67+
pdRouteFilterPrefillDecodePodsFail = "pd-filter-prefill-decode-pods-fail"
68+
pdRoutePrefillRequestError = "pd-do-prefill-request-error"
69+
pdRoutePrefillRequestSuccess = "pd-prefill-request-success"
7070
)
7171

7272
const (
@@ -522,6 +522,7 @@ func (r *pdRouter) finalPDScore(routingCtx *types.RoutingContext,
522522

523523
return targetPrefillPod, targetDecodePod, nil
524524
}
525+
525526
func (r *pdRouter) doPrefillRequest(routingCtx *types.RoutingContext, prefillPod *v1.Pod, llmEngine string) error {
526527
// Prepare prefill request payload
527528
payload, err := r.preparePrefillPayload(routingCtx, prefillPod, llmEngine)
@@ -577,53 +578,57 @@ func (r *pdRouter) doPrefillRequest(routingCtx *types.RoutingContext, prefillPod
577578
}()
578579

579580
case VLLMEngine:
580-
defer r.prefillRequestTracker.RemovePrefillRequest(routingCtx.RequestID)
581-
582581
// For vLLM, wait synchronously to get KV transfer params from response
583-
responseData, err := r.executeHTTPRequest(apiURL, routingCtx, payload)
584-
if err != nil {
585-
klog.ErrorS(err, "prefill_request_failed",
586-
"request_id", routingCtx.RequestID,
587-
"llm_engine", llmEngine,
588-
"prefill_pod", prefillPod.Name,
589-
"prefill_pod_ip", prefillPod.Status.PodIP,
590-
"elapsed", routingCtx.Elapsed(time.Now()))
591-
return fmt.Errorf("prefill request failed for request %s, pod %s: %w", routingCtx.RequestID, prefillPod.Name, err)
592-
}
582+
return r.handleSyncPrefill(routingCtx, prefillPod, llmEngine, apiURL, payload, fields, r.updateRoutingContextWithKVTransferParams, "KV transfer params")
593583

594-
// Update routing context with KV transfer params from prefill response
595-
if err := r.updateRoutingContextWithKVTransferParams(routingCtx, responseData, prefillPod); err != nil {
596-
return fmt.Errorf("failed to update routing context with KV transfer params for request %s: %w", routingCtx.RequestID, err)
597-
}
598-
599-
routingCtx.PrefillEndTime = time.Now()
600-
fields = append(fields,
601-
"routing_time_taken", routingCtx.PrefillStartTime.Sub(routingCtx.RequestTime),
602-
"prefill_time_taken", routingCtx.PrefillEndTime.Sub(routingCtx.PrefillStartTime),
603-
"outstanding_prefill_requests", r.prefillRequestTracker.GetPrefillRequestCountsForPod(prefillPod.Name)-1)
604-
klog.InfoS("prefill_request_end", fields...)
584+
case TensorRTLLM:
585+
// For TensorRT-LLM, wait synchronously to get disaggregated_params from response.
586+
// The prefill response contains first_gen_tokens and opaque_state needed by the decode worker.
587+
return r.handleSyncPrefill(routingCtx, prefillPod, llmEngine, apiURL, payload, fields, r.updateRoutingContextWithTRTDisaggParams, "TRT disagg params")
605588

606589
default:
607-
defer r.prefillRequestTracker.RemovePrefillRequest(routingCtx.RequestID)
608-
609590
// For unknown engines, use synchronous approach as a safe default
610-
if _, err := r.executeHTTPRequest(apiURL, routingCtx, payload); err != nil {
611-
klog.ErrorS(err, "prefill_request_failed",
612-
"request_id", routingCtx.RequestID,
613-
"llm_engine", llmEngine,
614-
"prefill_pod", prefillPod.Name,
615-
"prefill_pod_ip", prefillPod.Status.PodIP,
616-
"elapsed", routingCtx.Elapsed(time.Now()))
617-
return fmt.Errorf("prefill request failed for request %s, pod %s: %w", routingCtx.RequestID, prefillPod.Name, err)
591+
return r.handleSyncPrefill(routingCtx, prefillPod, llmEngine, apiURL, payload, fields, nil, "")
592+
}
593+
594+
return nil
595+
}
596+
597+
// handleSyncPrefill executes a synchronous prefill request, optionally calling updateCtxFunc
598+
// to process the response. Pass nil for updateCtxFunc when no response processing is needed.
599+
func (r *pdRouter) handleSyncPrefill(
600+
routingCtx *types.RoutingContext,
601+
prefillPod *v1.Pod,
602+
llmEngine, apiURL string,
603+
payload []byte,
604+
fields []interface{},
605+
updateCtxFunc func(*types.RoutingContext, map[string]any, *v1.Pod) error,
606+
errorContext string) error {
607+
defer r.prefillRequestTracker.RemovePrefillRequest(routingCtx.RequestID)
608+
609+
responseData, err := r.executeHTTPRequest(apiURL, routingCtx, payload)
610+
if err != nil {
611+
klog.ErrorS(err, "prefill_request_failed",
612+
"request_id", routingCtx.RequestID,
613+
"llm_engine", llmEngine,
614+
"prefill_pod", prefillPod.Name,
615+
"prefill_pod_ip", prefillPod.Status.PodIP,
616+
"elapsed", routingCtx.Elapsed(time.Now()))
617+
return fmt.Errorf("prefill request failed for request %s, pod %s: %w", routingCtx.RequestID, prefillPod.Name, err)
618+
}
619+
620+
if updateCtxFunc != nil {
621+
if err := updateCtxFunc(routingCtx, responseData, prefillPod); err != nil {
622+
return fmt.Errorf("failed to update routing context with %s for request %s: %w", errorContext, routingCtx.RequestID, err)
618623
}
619-
routingCtx.PrefillEndTime = time.Now()
620-
fields = append(fields,
621-
"routing_time_taken", routingCtx.PrefillStartTime.Sub(routingCtx.RequestTime),
622-
"prefill_time_taken", routingCtx.PrefillEndTime.Sub(routingCtx.PrefillStartTime),
623-
"outstanding_prefill_requests", r.prefillRequestTracker.GetPrefillRequestCountsForPod(prefillPod.Name)-1)
624-
klog.InfoS("prefill_request_end", fields...)
625624
}
626625

626+
routingCtx.PrefillEndTime = time.Now()
627+
fields = append(fields,
628+
"routing_time_taken", routingCtx.PrefillStartTime.Sub(routingCtx.RequestTime),
629+
"prefill_time_taken", routingCtx.PrefillEndTime.Sub(routingCtx.PrefillStartTime),
630+
"outstanding_prefill_requests", r.prefillRequestTracker.GetPrefillRequestCountsForPod(prefillPod.Name)-1)
631+
klog.InfoS("prefill_request_end", fields...)
627632
return nil
628633
}
629634

@@ -662,9 +667,22 @@ func (r *pdRouter) preparePrefillPayload(routingCtx *types.RoutingContext, pod *
662667
}
663668
}
664669

670+
if llmEngine == TensorRTLLM {
671+
// Signal to TensorRT-LLM that this is a context-only (prefill) request.
672+
// The prefill response will return disaggregated_params containing
673+
// first_gen_tokens and opaque_state, which are injected into the decode request.
674+
completionRequest["disaggregated_params"] = map[string]any{
675+
"request_type": "context_only",
676+
}
677+
}
678+
665679
// Set prefill-specific parameters
666680
completionRequest["max_tokens"] = 1
667-
completionRequest["max_completion_tokens"] = 1
681+
if llmEngine == TensorRTLLM {
682+
delete(completionRequest, "max_completion_tokens")
683+
} else {
684+
completionRequest["max_completion_tokens"] = 1
685+
}
668686
completionRequest["stream"] = false
669687
delete(completionRequest, "stream_options")
670688

@@ -788,6 +806,57 @@ func (r *pdRouter) updateRoutingContextWithKVTransferParams(routingCtx *types.Ro
788806
return nil
789807
}
790808

809+
func (r *pdRouter) updateRoutingContextWithTRTDisaggParams(routingCtx *types.RoutingContext, responseData map[string]any, prefillPod *v1.Pod) error {
810+
// Parse the original request body
811+
var originalRequest map[string]any
812+
if err := sonic.Unmarshal(routingCtx.ReqBody, &originalRequest); err != nil {
813+
return fmt.Errorf("failed to unmarshal original request body: %w", err)
814+
}
815+
816+
// Extract disaggregated_params from prefill response.
817+
// TRT-LLM may return it at the top level or inside choices[0].
818+
var disaggParams any
819+
var exists bool
820+
821+
disaggParams, exists = responseData["disaggregated_params"]
822+
if !exists {
823+
// Fallback: check choices[0] (TRT-LLM serializes handler output as a choice)
824+
if choices, ok := responseData["choices"].([]any); ok && len(choices) > 0 {
825+
if choice, ok := choices[0].(map[string]any); ok {
826+
disaggParams, exists = choice["disaggregated_params"]
827+
}
828+
}
829+
}
830+
831+
if !exists {
832+
klog.InfoS("no disaggregated_params in TRT prefill response", "request_id", routingCtx.RequestID)
833+
return nil
834+
}
835+
836+
disaggParamsMap, ok := disaggParams.(map[string]any)
837+
if !ok {
838+
return fmt.Errorf("disaggregated_params has unexpected type %T, expected map[string]any", disaggParams)
839+
}
840+
841+
// Override request_type to generation_only for the decode request
842+
disaggParamsMap["request_type"] = "generation_only"
843+
originalRequest["disaggregated_params"] = disaggParamsMap
844+
845+
updatedReqBody, err := sonic.Marshal(originalRequest)
846+
if err != nil {
847+
return fmt.Errorf("failed to marshal updated request body: %w", err)
848+
}
849+
850+
routingCtx.ReqBody = updatedReqBody
851+
852+
klog.InfoS("updated routing context with disaggregated_params (TensorRT-LLM)",
853+
"request_id", routingCtx.RequestID,
854+
"prefill_pod", prefillPod.Name,
855+
"prefill_host", prefillPod.Status.PodIP)
856+
857+
return nil
858+
}
859+
791860
func (r *pdRouter) SubscribedMetrics() []string {
792861
return []string{}
793862
}

0 commit comments

Comments
 (0)