-
Notifications
You must be signed in to change notification settings - Fork 2.6k
[None][fix] Align GPTOSS router tokenization and disagg draft scheduling #15605
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from all commits
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -703,6 +703,10 @@ def __init__( | |
| # responses and flushing them at a synchronised point in the executor | ||
| # loop avoids the mismatch. | ||
| self._pending_transfer_responses: List[Tuple[int, LlmResponse]] = [] | ||
| self._disagg_generation_trans_in_progress_requests: Dict[ | ||
| int, LlmRequest] = {} | ||
| self._pending_disagg_generation_trans_complete_requests: Dict[ | ||
| int, LlmRequest] = {} | ||
| # Same buffer-then-synced-flush pattern as _pending_transfer_responses | ||
| # above: _handle_responses and _append_iter_stats are reached from | ||
| # per-rank-divergent gates, so their tp_allgather collectives are | ||
|
|
@@ -2413,6 +2417,8 @@ def _pp_schedule_and_propagate(self, microbatch_id: int): | |
| serializable_schedule = None | ||
| wait_for_disagg_gen_transfer_progress = False | ||
| is_dp_broadcast = self.dist.tp_size > 1 and self.enable_attention_dp | ||
| if getattr(self, "kv_cache_transceiver", None): | ||
| self._sync_pending_disagg_generation_trans_complete_draft_tokens() | ||
| if self.dist.rank == 0 or (self.dist.is_first_pp_rank | ||
| and is_dp_broadcast): | ||
| scheduled_batch, fitting_disagg_gen_init_requests, num_fitting_reqs = self._schedule( | ||
|
|
@@ -3355,6 +3361,120 @@ def _commit_kv_cache_stats(self, | |
| self.kv_cache_manager.commit_scheduled_kv_cache_stats( | ||
| scheduled_batch) | ||
|
|
||
| @staticmethod | ||
| def _get_disagg_generation_transfer_id( | ||
| request: LlmRequest) -> Optional[int]: | ||
| disagg_params = getattr(request, "py_disaggregated_params", None) | ||
| transfer_id = getattr(disagg_params, "disagg_request_id", None) | ||
| if isinstance(transfer_id, int): | ||
| return transfer_id | ||
| request_id = getattr(request, "py_request_id", None) | ||
| if isinstance(request_id, int): | ||
| return request_id | ||
| request_id = getattr(request, "request_id", None) | ||
| if isinstance(request_id, int): | ||
| return request_id | ||
| return None | ||
|
|
||
| def _get_disagg_generation_trans_in_progress_requests( | ||
| self) -> Dict[int, LlmRequest]: | ||
| requests = getattr(self, | ||
| "_disagg_generation_trans_in_progress_requests", | ||
| None) | ||
| if requests is None: | ||
| requests = {} | ||
| self._disagg_generation_trans_in_progress_requests = requests | ||
| return requests | ||
|
|
||
| def _get_pending_disagg_generation_trans_complete_requests( | ||
| self) -> Dict[int, LlmRequest]: | ||
| requests = getattr( | ||
| self, "_pending_disagg_generation_trans_complete_requests", None) | ||
| if requests is None: | ||
| requests = {} | ||
| self._pending_disagg_generation_trans_complete_requests = requests | ||
| return requests | ||
|
|
||
| def _track_disagg_generation_trans_in_progress( | ||
| self, requests: Iterable[LlmRequest]) -> None: | ||
| in_progress = self._get_disagg_generation_trans_in_progress_requests() | ||
| for request in requests: | ||
| transfer_id = self._get_disagg_generation_transfer_id(request) | ||
| if transfer_id is not None: | ||
| in_progress[transfer_id] = request | ||
|
|
||
| def _queue_disagg_generation_trans_complete_draft_sync( | ||
| self, requests: Iterable[LlmRequest]) -> None: | ||
| pending = self._get_pending_disagg_generation_trans_complete_requests() | ||
| for request in requests: | ||
| if not getattr(request, | ||
| "is_disagg_generation_transmission_complete", False): | ||
| continue | ||
| transfer_id = self._get_disagg_generation_transfer_id(request) | ||
| if transfer_id is not None: | ||
| pending[transfer_id] = request | ||
|
|
||
| def _queue_disagg_generation_trans_complete_draft_sync_by_ids( | ||
| self, transfer_ids: Iterable[int]) -> None: | ||
| in_progress = self._get_disagg_generation_trans_in_progress_requests() | ||
| completed_requests = [] | ||
| for transfer_id in transfer_ids: | ||
| request = in_progress.pop(transfer_id, None) | ||
| if request is not None: | ||
| completed_requests.append(request) | ||
| self._queue_disagg_generation_trans_complete_draft_sync( | ||
| completed_requests) | ||
|
|
||
| @staticmethod | ||
| def _resolve_disagg_generation_trans_complete_draft_tokens( | ||
| context_phase_draft_tokens: Optional[Iterable[int]], | ||
| enable_spec_decode: bool, max_total_draft_tokens: int) -> List[int]: | ||
| draft_tokens = ([] if context_phase_draft_tokens is None else | ||
| list(context_phase_draft_tokens)) | ||
| if not draft_tokens and enable_spec_decode: | ||
| return [0] * max_total_draft_tokens | ||
| return draft_tokens | ||
|
|
||
| @staticmethod | ||
| def _sync_disagg_generation_trans_complete_draft_tokens( | ||
| requests: Iterable[LlmRequest], | ||
| enable_spec_decode: bool = False, | ||
| max_total_draft_tokens: int = 0) -> None: | ||
| for request in requests: | ||
| if not getattr(request, | ||
| "is_disagg_generation_transmission_complete", False): | ||
| continue | ||
|
|
||
| context_phase_params = request.context_phase_params | ||
| if context_phase_params is None: | ||
| continue | ||
|
|
||
| request.py_draft_tokens = PyExecutor._resolve_disagg_generation_trans_complete_draft_tokens( | ||
| context_phase_params.draft_tokens, enable_spec_decode, | ||
| max_total_draft_tokens) | ||
| request.draft_tokens = request.py_draft_tokens | ||
| request.py_draft_pages_allocated = len(request.py_draft_tokens) | ||
|
coderabbitai[bot] marked this conversation as resolved.
|
||
|
|
||
| def _sync_pending_disagg_generation_trans_complete_draft_tokens( | ||
| self) -> None: | ||
| pending = self._get_pending_disagg_generation_trans_complete_requests() | ||
| if not pending: | ||
| return | ||
| max_total_draft_tokens = getattr(self.model_engine, | ||
| "max_total_draft_tokens", | ||
| self.max_total_draft_tokens) | ||
| self._sync_disagg_generation_trans_complete_draft_tokens( | ||
| pending.values(), self.model_engine.enable_spec_decode, | ||
| max_total_draft_tokens) | ||
| pending.clear() | ||
|
|
||
| @staticmethod | ||
| def _get_generation_num_draft_tokens(request: LlmRequest) -> int: | ||
|
SimengLiu-nv marked this conversation as resolved.
|
||
| py_draft_tokens = getattr(request, "py_draft_tokens", None) | ||
| if py_draft_tokens is None: | ||
| return request.num_draft_tokens | ||
| return max(len(py_draft_tokens), request.num_draft_tokens) | ||
|
|
||
| def _get_disagg_transfer_admission_controller( | ||
| self) -> DisaggTransferAdmissionController: | ||
| controller = getattr(self, "_disagg_transfer_admission_controller", | ||
|
|
@@ -3649,6 +3769,9 @@ def _prepare_and_schedule_batch(self): | |
| continue | ||
| request.draft_tokens = [0] * self.max_total_draft_tokens | ||
|
|
||
| if self.kv_cache_transceiver: | ||
| self._sync_pending_disagg_generation_trans_complete_draft_tokens() | ||
|
|
||
| scheduled_batch, scheduler_fitting_disagg_gen_init_requests, num_fitting_reqs = self._schedule( | ||
| ) | ||
|
|
||
|
|
@@ -5200,8 +5323,9 @@ def _compute_scheduled_tokens(context_requests, generation_requests): | |
| else: | ||
| compute = max(1, remaining - reusable_in_chunk) | ||
| num_scheduled_ctx_tokens += compute | ||
| num_scheduled_gen_tokens = sum(1 + gen_req.num_draft_tokens | ||
| for gen_req in generation_requests) | ||
| num_scheduled_gen_tokens = sum( | ||
| 1 + PyExecutor._get_generation_num_draft_tokens(gen_req) | ||
| for gen_req in generation_requests) | ||
| return num_scheduled_ctx_tokens + num_scheduled_gen_tokens | ||
|
|
||
| def _waiting_requests(self, context_requests: list[LlmRequest], | ||
|
|
@@ -5839,7 +5963,10 @@ def _prepare_disagg_gen_transmission_complete(self, scheduled_batch): | |
| ctx_draft_tokens = [ | ||
| 0 | ||
| ] * self.model_engine.max_total_draft_tokens | ||
| req.py_draft_tokens = [] if ctx_draft_tokens is None else ctx_draft_tokens | ||
| req.py_draft_tokens = [] if ctx_draft_tokens is None else list( | ||
| ctx_draft_tokens) | ||
|
SimengLiu-nv marked this conversation as resolved.
|
||
| req.draft_tokens = req.py_draft_tokens | ||
| req.py_draft_pages_allocated = len(req.py_draft_tokens) | ||
| beam_width = req.py_beam_width | ||
| if not self._update_sampler_state_for_disagg_gen_request( | ||
| req, beam_width, first_gen_tokens): | ||
|
|
@@ -5981,6 +6108,8 @@ def _recv_disagg_gen_cache(self, new_gen_reqs): | |
| if self._is_disagg_gen_only_no_context_benchmark(): | ||
| for req in new_gen_reqs: | ||
| req.state = LlmRequestState.DISAGG_GENERATION_TRANS_COMPLETE | ||
| self._queue_disagg_generation_trans_complete_draft_sync( | ||
| new_gen_reqs) | ||
| return | ||
|
|
||
| if not self._uses_async_disagg_gen_transfer(): | ||
|
|
@@ -5991,10 +6120,13 @@ def _recv_disagg_gen_cache(self, new_gen_reqs): | |
| self.kv_cache_transceiver.request_and_receive_sync(req) | ||
| if req.state == LlmRequestState.DISAGG_GENERATION_TRANS_COMPLETE: | ||
| self._sync_disagg_transfer_made_progress = True | ||
| self._queue_disagg_generation_trans_complete_draft_sync( | ||
| [req]) | ||
| self._check_cache_transfer_errors("generation requests") | ||
| return | ||
|
|
||
| for req in new_gen_reqs: | ||
| self._track_disagg_generation_trans_in_progress([req]) | ||
|
Collaborator
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Entries added here are only popped via |
||
| self.kv_cache_transceiver.request_and_receive_async(req) | ||
|
|
||
| if self.kv_cache_transceiver.kv_transfer_timeout_ms is not None: | ||
|
|
@@ -6147,10 +6279,19 @@ def _check_disagg_ctx_cache_transfer_status(self, atLeastNum: int = 0): | |
| def _check_disagg_gen_cache_transfer_status(self, atLeastNum: int = 0): | ||
| result = self.kv_cache_transceiver.check_gen_transfer_status(atLeastNum) | ||
| if isinstance(result, tuple): | ||
| _, _, cancelled_reqs = result | ||
| completed_request_ids, failed_request_ids, cancelled_reqs = result | ||
| self._queue_disagg_generation_trans_complete_draft_sync_by_ids( | ||
| completed_request_ids) | ||
| in_progress = self._get_disagg_generation_trans_in_progress_requests( | ||
| ) | ||
| for request_id in failed_request_ids: | ||
| in_progress.pop(request_id, None) | ||
| user_canceled_set = set(self.canceled_req_ids) | ||
| for req in cancelled_reqs: | ||
| req_id = req.py_request_id if not req.is_child else req.parent_request_id | ||
| transfer_id = self._get_disagg_generation_transfer_id(req) | ||
| if transfer_id is not None: | ||
| in_progress.pop(transfer_id, None) | ||
| if req_id not in user_canceled_set: | ||
| req.state = LlmRequestState.DISAGG_TRANS_ERROR | ||
| if not self._is_disagg_inflight_cancel_active(): | ||
|
|
||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -188,19 +188,23 @@ def _init_block_hashing( | |
| tokens_per_block: Optional[int] = None, | ||
| custom_tokenizer: Optional[str] = None, | ||
| tokenizer_dir: Optional[str] = None, | ||
| ): | ||
| use_harmony: Optional[bool] = None, | ||
| ) -> None: | ||
| env_tokens_per_block = os.environ.get("TRTLLM_KVCACHE_AWARE_ROUTER_HASH_TOKENS_PER_BLOCK") | ||
| if env_tokens_per_block is not None: | ||
| tokens_per_block = int(env_tokens_per_block) | ||
| self._tpb_auto = tokens_per_block is None | ||
| self._tokens_per_block = 32 if tokens_per_block is None else tokens_per_block | ||
| self._tokenizers: dict = {} | ||
| self._model_types: dict[str, Optional[str]] = {} | ||
| self._custom_tokenizer = custom_tokenizer | ||
| self._tokenizer_dir = tokenizer_dir | ||
| self._use_harmony = use_harmony | ||
| logger.info( | ||
| f"BlockHashMixin: tokens_per_block={self._tokens_per_block}" | ||
| f"{' (auto, adopts worker)' if self._tpb_auto else ''}" | ||
| f", custom_tokenizer={self._custom_tokenizer}" | ||
| f", use_harmony={self._use_harmony}" | ||
| ) | ||
|
|
||
| def _get_tokenizer(self, model: str): | ||
|
|
@@ -219,6 +223,41 @@ def _get_tokenizer(self, model: str): | |
| self._tokenizers[model] = tokenizer.tokenizer | ||
| return self._tokenizers[model] | ||
|
|
||
| def _get_model_type(self, model: str) -> Optional[str]: | ||
| if model not in self._model_types: | ||
| model_path = self._tokenizer_dir or model | ||
| from tensorrt_llm.tokenizer import resolve_model_type_from_name_or_config | ||
|
|
||
| self._model_types[model] = resolve_model_type_from_name_or_config(model_path) | ||
| return self._model_types[model] | ||
|
|
||
| def _uses_harmony_tokenization(self, request: ChatCompletionRequest) -> bool: | ||
| from tensorrt_llm.tokenizer import uses_harmony_tokenization | ||
|
|
||
| return uses_harmony_tokenization(self._get_model_type(request.model), self._use_harmony) | ||
|
|
||
| @staticmethod | ||
| def _tool_dicts(request: ChatCompletionRequest) -> Optional[list[dict[str, object]]]: | ||
| if request.tools is None: | ||
| return None | ||
| return [ | ||
| tool.model_dump() if hasattr(tool, "model_dump") else tool for tool in request.tools | ||
| ] | ||
|
|
||
| def _tokenize_harmony_chat(self, request: ChatCompletionRequest) -> list[list[int]]: | ||
|
Collaborator
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Could we unify this implementation with |
||
| from tensorrt_llm.serve import harmony_adapter | ||
|
|
||
| tools = self._tool_dicts(request) if request.tools else None | ||
| result = harmony_adapter.get_harmony_adapter().openai_to_harmony_tokens( | ||
| request.messages, | ||
| tools, | ||
| reasoning_effort=harmony_adapter.maybe_transform_reasoning_effort( | ||
| request.reasoning_effort | ||
| ), | ||
| tool_choice=request.tool_choice, | ||
| ) | ||
| return [result] | ||
|
|
||
| def _encode_with_prefix_cache(self, rendered: str, key: int, tokenizer) -> list[int]: | ||
| cache = getattr(self, "_tok_prefix_cache", None) | ||
| if cache is None: | ||
|
|
@@ -241,26 +280,19 @@ def _tokenize(self, request: OpenAIRequest) -> list[list[int]]: | |
| if isinstance(request, ChatCompletionRequest): | ||
| if request.prompt_token_ids is not None: | ||
| return [request.prompt_token_ids] | ||
| if self._uses_harmony_tokenization(request): | ||
| return self._tokenize_harmony_chat(request) | ||
| tokenizer = self._get_tokenizer(request.model) | ||
| tool_dicts = ( | ||
| None | ||
| if getattr(request, "tools", None) is None | ||
| else [ | ||
| tool.model_dump() if hasattr(tool, "model_dump") else tool | ||
| for tool in request.tools | ||
| ] | ||
| ) | ||
| chat_template_kwargs = ( | ||
| request.chat_template_kwargs | ||
| if getattr(request, "chat_template_kwargs", None) | ||
| else {} | ||
| ) | ||
| chat_template_kwargs = dict(request.chat_template_kwargs or {}) | ||
| chat_template_kwargs["tools"] = self._tool_dicts(request) | ||
| chat_template_kwargs["documents"] = request.documents | ||
| if request.chat_template is not None: | ||
| chat_template_kwargs["chat_template"] = request.chat_template | ||
| rendered = tokenizer.apply_chat_template( | ||
| [msg if isinstance(msg, dict) else dict(msg) for msg in request.messages], | ||
| add_generation_prompt=request.add_generation_prompt, | ||
| tokenize=False, | ||
| return_dict=False, | ||
| tools=tool_dicts, | ||
| **chat_template_kwargs, | ||
| ) | ||
| if isinstance(rendered, str): | ||
|
|
||
Uh oh!
There was an error while loading. Please reload this page.
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
The draft-token sync itself is small, but it adds ~10 helpers, two tracking dicts, and 6 touch points in
py_executor.py(our most safety-critical file) mainly to avoid anactive_requestsscan. Couldcheck_gen_transfer_statusreturn the completed request objects (as it already does for cancelled ones) so we can set the draft tokens inline on completion? That would let us drop the_disagg_generation_trans_in_progress_requeststracking layer (and the cancel-path leak it introduces) and keep the footprint much smaller. cc @chuangz0