MCPcopy Create free account
hub / github.com/NVIDIA/TensorRT-LLM / AwaitResponseHelper

Class AwaitResponseHelper

tensorrt_llm/executor/base_worker.py:688–816  ·  view source on GitHub ↗

Multiple-implementations for await_response for performance.

Source from the content-addressed store, hash-verified

686
687
688class AwaitResponseHelper:
689 ''' Multiple-implementations for await_response for performance. '''
690
691 class HandlerKind(enum.Enum):
692 unknown = 0
693 single_process_worker = 1
694 ipc_batched = 2
695
696 def __init__(self, worker: "BaseWorker"):
697 # TODO: make worker weakref
698 self.worker = worker
699 self.handler_kind: AwaitResponseHelper.HandlerKind = AwaitResponseHelper.HandlerKind.unknown
700 self.enable_postprocprocess_parallel = self.worker.enable_postprocess_parallel
701 # The error responses when submit request failed will be put here
702 self.temp_error_responses = Queue()
703
704 def responses_handler(self, responses: List[tllm.Response]):
705 HandlerKind = AwaitResponseHelper.HandlerKind
706
707 if self.handler_kind is HandlerKind.unknown:
708 if not (self.worker.result_queue is not None
709 or self.worker.postproc_queues is not None):
710 logger_debug(f"creating await_response helper for Worker\n",
711 color="yellow")
712 # When ExecutorBindingWorker is used in the main process
713 # aka the single process mode
714 self.handler_kind = HandlerKind.single_process_worker
715 elif self.worker.result_queue is not None or self.worker.postproc_queues is not None:
716 # The ExecutorBindingProxy is used
717 logger_debug(f"creating await_response helper for IPC\n",
718 color="yellow")
719 self.handler_kind = HandlerKind.ipc_batched
720 else:
721 raise NotImplementedError
722
723 match self.handler_kind:
724 case HandlerKind.single_process_worker:
725 return self.handle_for_worker(responses)
726 case HandlerKind.ipc_batched:
727 return self.handle_for_ipc_batched(responses)
728 case _:
729 raise NotImplementedError
730
731 def __call__(self, timeout: Optional[float] = None) -> bool:
732 ''' This method should be called by a ManagedThread. '''
733 timeout = timeout or 0.1
734 responses = self.worker.engine.await_responses(
735 timeout=datetime.timedelta(seconds=timeout))
736 # filter since The _engine_response_callback may return None
737 responses = list(
738 filter(
739 lambda _: _,
740 [self.worker._engine_response_callback(r) for r in responses]))
741
742 # append the error responses to the temp_error_responses
743 while not self.temp_error_responses.empty():
744 responses.append(self.temp_error_responses.get())
745

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected