Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 7 additions & 3 deletions tensorrt_llm/_torch/pyexecutor/sampler/sampler.py
Original file line number Diff line number Diff line change
Expand Up @@ -1173,6 +1173,11 @@ class SampleStateTensorsHostTorch(SampleStateTensors):
finish_reasons: torch.Tensor | None
first_finish_reasons: torch.Tensor | None
logprobs_state: LogProbsState | None = None
single_step_greedy: bool = False
"""Marks `new_tokens` as the compact 1-D `(num_requests,)` layout instead of
`[step, slot, beam]`. Must stay on this dataclass, the only part of the sample
state pipeline parallelism transfers (`_ring_broadcast_sample_state` sends just
`host`), so a receiving rank cannot pair the buffer with a stale default."""

def finish_reasons_list(self) -> FinishReasonsList:
"""`(num_seq_slots, num_steps)`"""
Expand All @@ -1187,7 +1192,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:
Expand Down Expand Up @@ -2358,7 +2362,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

Expand Down Expand Up @@ -2746,10 +2750,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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -328,7 +328,7 @@ def sample(
)
case ("greedy", None):
tokens, softmax = greedy_search_sampling_batch(logits, return_probs=return_probs)
temperature = None
return tokens, softmax, None
case (
"beam_search",
beam_width_in,
Expand Down
4 changes: 2 additions & 2 deletions tests/unittest/_torch/sampler/test_torch_sampler.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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.
Expand Down
Loading