From d2e841976c00b70aea295637bf2099ad43dbb2c6 Mon Sep 17 00:00:00 2001 From: trtllm-agent <296075020+trtllm-agent@users.noreply.github.com> Date: Mon, 17 Aug 2026 23:02:04 -0700 Subject: [PATCH] [nvbugs/6627041][fix] Carry single_step_greedy on the host sample state for PP Signed-off-by: trtllm-agent <296075020+trtllm-agent@users.noreply.github.com> --- tensorrt_llm/_torch/pyexecutor/sampler/sampler.py | 8 +++++--- tests/unittest/_torch/sampler/test_torch_sampler.py | 4 ++-- 2 files changed, 7 insertions(+), 5 deletions(-) diff --git a/tensorrt_llm/_torch/pyexecutor/sampler/sampler.py b/tensorrt_llm/_torch/pyexecutor/sampler/sampler.py index 71dbac467ad3..f92c6eb1b038 100644 --- a/tensorrt_llm/_torch/pyexecutor/sampler/sampler.py +++ b/tensorrt_llm/_torch/pyexecutor/sampler/sampler.py @@ -1173,6 +1173,9 @@ class SampleStateTensorsHostTorch(SampleStateTensors): finish_reasons: torch.Tensor | None first_finish_reasons: torch.Tensor | None logprobs_state: LogProbsState | None = None + # Describes the layout of `new_tokens`, so it must travel with it: pipeline + # parallelism only transports this host object between ranks. + single_step_greedy: bool = False def finish_reasons_list(self) -> FinishReasonsList: """`(num_seq_slots, num_steps)`""" @@ -1187,7 +1190,6 @@ def finish_reasons_list(self) -> FinishReasonsList: @dataclass(kw_only=True) class SampleStateTorch(SampleState[SampleStateTensorsHostTorch, SampleStateTensors]): beam_history_builders: list[BeamHistoryBuilder | None] | None = None - single_step_greedy: bool = False class _SideStreamCopier: @@ -2358,7 +2360,7 @@ def update_requests( assert state.host is not None # Reuse sample_async's qualification instead of rechecking every # request after the asynchronous sample completes. - if state.single_step_greedy: + if state.host.single_step_greedy: self._update_requests_single_beam_single_step(state) return @@ -2746,10 +2748,10 @@ def sample_async( finish_reasons=finish_reasons_host, first_finish_reasons=first_finish_reasons_host, logprobs_state=logprobs_state, + single_step_greedy=single_step_greedy, ), sampler_event=sampler_event, beam_history_builders=beam_history_builders, - single_step_greedy=single_step_greedy, ) @staticmethod diff --git a/tests/unittest/_torch/sampler/test_torch_sampler.py b/tests/unittest/_torch/sampler/test_torch_sampler.py index 64401627ca29..978f9934422e 100644 --- a/tests/unittest/_torch/sampler/test_torch_sampler.py +++ b/tests/unittest/_torch/sampler/test_torch_sampler.py @@ -961,8 +961,8 @@ def test_single_step_greedy_updates_finish_reasons_and_filters_completed_request new_tokens=new_tokens, finish_reasons=None, first_finish_reasons=None, + single_step_greedy=True, ), - single_step_greedy=True, ) sampler.update_requests(state) @@ -1986,7 +1986,7 @@ def _mock_filter(self, requests: ScheduledRequests) -> list[LlmRequest]: sample_state.sampler_event.synchronize() assert sample_state.host is not None host_new_tokens = sample_state.host.new_tokens - if sample_state.single_step_greedy: + if sample_state.host.single_step_greedy: # The stable greedy path copies one token per active request instead of # the full [step, slot, beam] buffer. This fixture uses dense sequence # slots, so restore that layout before comparing sampling results.