Multiple-implementations for await_response for performance.
| 686 | |
| 687 | |
| 688 | class 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 |