Tiny extract ModelRunnerOutput (#15400)

This commit is contained in:
fzyzcjy
2025-12-18 22:18:45 +08:00
committed by GitHub
parent c5f4e20f2f
commit ad9616f13a
7 changed files with 34 additions and 27 deletions
@@ -614,7 +614,7 @@ class SchedulerPPMixin:
forward_batch = ForwardBatch.init_new(
model_worker_batch, self.tp_worker.model_runner
)
_, _ = self.tp_worker.model_runner.forward(
_ = self.tp_worker.model_runner.forward(
forward_batch=forward_batch, pp_proxy_tensors=pp_proxy
)
+7 -4
View File
@@ -196,7 +196,7 @@ class BaseTpWorker(ABC):
def forward_batch_embedding(self, model_worker_batch: ModelWorkerBatch):
forward_batch = ForwardBatch.init_new(model_worker_batch, self.model_runner)
logits_output, _ = self.model_runner.forward(forward_batch)
logits_output = self.model_runner.forward(forward_batch).logits_output
embeddings = logits_output.embeddings
return embeddings
@@ -397,11 +397,12 @@ class TpModelWorker(BaseTpWorker):
return self._forward_batch_generation_dllm(forward_batch)
if self.pp_group.is_last_rank:
logits_output, can_run_cuda_graph = self.model_runner.forward(
out = self.model_runner.forward(
forward_batch,
pp_proxy_tensors=pp_proxy_tensors,
skip_attn_backend_init=skip_attn_backend_init,
)
logits_output, can_run_cuda_graph = out.logits_output, out.can_run_graph
batch_result = GenerationBatchResult(
logits_output=logits_output,
can_run_cuda_graph=can_run_cuda_graph,
@@ -450,11 +451,12 @@ class TpModelWorker(BaseTpWorker):
return batch_result
else:
pp_proxy_tensors, can_run_cuda_graph = self.model_runner.forward(
out = self.model_runner.forward(
forward_batch,
pp_proxy_tensors=pp_proxy_tensors,
skip_attn_backend_init=skip_attn_backend_init,
)
pp_proxy_tensors, can_run_cuda_graph = out.logits_output, out.can_run_graph
return GenerationBatchResult(
pp_hidden_states_proxy_tensors=pp_proxy_tensors,
can_run_cuda_graph=can_run_cuda_graph,
@@ -469,9 +471,10 @@ class TpModelWorker(BaseTpWorker):
else:
model_worker_batch = batch.get_model_worker_batch(batch.seq_lens_cpu_cache)
logits_output, can_run_cuda_graph = self.model_runner.forward(
out = self.model_runner.forward(
batch.split_forward_batch, split_forward_count=batch.split_forward_count
)
logits_output, can_run_cuda_graph = out.logits_output, out.can_run_graph
if logits_output:
next_token_ids = self.model_runner.sample(logits_output, model_worker_batch)
else: