diff --git a/tensorrt_llm/_torch/pyexecutor/py_executor.py b/tensorrt_llm/_torch/pyexecutor/py_executor.py index 4f0fac787c03..78b314b0a065 100644 --- a/tensorrt_llm/_torch/pyexecutor/py_executor.py +++ b/tensorrt_llm/_torch/pyexecutor/py_executor.py @@ -3264,10 +3264,15 @@ def _ring_broadcast_sample_state( if not self.dist.is_last_pp_rank: # Receive tokens from previous pp rank (w.r.t model forward direction) with nvtx_range("recv_sample_state"): - sample_state.host, py_result_diffs = self.dist.recv_object( + # SampleStateTorch carries a ``use_host_stop_criteria`` flag + # decided on the last PP rank; propagate it so ``update_requests`` + # picks the same branch on all ranks. Other samplers ship None. + sample_state.host, py_result_diffs, use_host_stop_criteria = self.dist.recv_object( src=self.dist.prev_pp_rank, tag=tag, ) + if use_host_stop_criteria is not None: + sample_state.use_host_stop_criteria = use_host_stop_criteria for request, py_result_diff in zip(requests, py_result_diffs): request.py_result.apply_diff(py_result_diff) @@ -3285,7 +3290,8 @@ def _ring_broadcast_sample_state( self.wait_on_pp_send_handles(self.send_handles, microbatch_id) with nvtx_range("send_sample_state"): self.send_handles[microbatch_id] = self.dist.isend_object( - (sample_state.host, py_result_diffs), + (sample_state.host, py_result_diffs, + getattr(sample_state, "use_host_stop_criteria", None)), dest=self.dist.next_pp_rank, tag=tag, )