Skip to content
Merged
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
24 changes: 16 additions & 8 deletions vllm_ascend/worker/v2/model_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -490,7 +490,10 @@ def postprocess_sampled(
query_start_loc,
)

self._copy_num_computed_tokens_to_cpu()
# Skip D2H copy without MTP: num_computed_tokens_cpu is synced
# from num_computed_tokens_np in _update_seq_lens_cpu instead.
if self.speculator is not None:
self._copy_num_computed_tokens_to_cpu()

def _copy_num_computed_tokens_to_cpu(self):
# npu attention backend still need to use seq_lens_cpu,
Expand All @@ -512,13 +515,18 @@ def _update_seq_lens_cpu(
req_ids: list[str],
):
num_scheduled_tokens = scheduler_output.num_scheduled_tokens
# wait for num_computed_tokens copy to cpu stream to finish.
self.num_computed_tokens_event.synchronize()
for req_id in scheduler_output.scheduled_cached_reqs.req_ids:
req_index = self.req_states.req_id_to_index[req_id]
# num_computed_tokens_cpu has reverted by num_rejected_tokens already.
# in super postprocess method.
self.req_states.num_computed_tokens_cpu[req_index] = self.num_computed_tokens_cpu[req_index]

# MTP needs D2H copy to get reverted num_computed_tokens after rejection.
# Without MTP, num_computed_tokens_np is already correct from update_requests.
if self.speculator is not None:
self.num_computed_tokens_event.synchronize()
for req_id in scheduler_output.scheduled_cached_reqs.req_ids:
req_index = self.req_states.req_id_to_index[req_id]
self.req_states.num_computed_tokens_cpu[req_index] = self.num_computed_tokens_cpu[req_index]
else:
for req_id in scheduler_output.scheduled_cached_reqs.req_ids:
req_index = self.req_states.req_id_to_index[req_id]
self.req_states.num_computed_tokens_cpu[req_index] = self.req_states.num_computed_tokens_np[req_index]
Comment thread
xiayingqing marked this conversation as resolved.

# update seq_lens_cpu
for i, req_id in enumerate(req_ids): # type: ignore
Expand Down
Loading