Tiny extract ModelRunnerOutput (#15400)
This commit is contained in:
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user