model: support EmbeddingGemma (#32375)

This commit is contained in:
Mick
2026-07-27 10:40:47 +08:00
committed by GitHub
parent a358374ae9
commit abb8f4b5e3
12 changed files with 209 additions and 16 deletions
+8 -2
View File
@@ -3529,18 +3529,24 @@ class Scheduler(
with self.forward_stream_ctx:
self.forward_stream.wait_stream(self.schedule_stream)
resolve_forward_inputs(batch, self.future_map)
pooler_output = self.tp_worker.forward_batch_embedding(batch)
pooler_output, can_run_cuda_graph = (
self.tp_worker.forward_batch_embedding(batch)
)
ret = EmbeddingBatchResult(
embeddings=pooler_output.embeddings,
pooled_hidden_states=pooler_output.pooled_hidden_states,
can_run_cuda_graph=can_run_cuda_graph,
)
ret.copy_to_cpu()
else:
resolve_forward_inputs(batch, self.future_map)
pooler_output = self.tp_worker.forward_batch_embedding(batch)
pooler_output, can_run_cuda_graph = (
self.tp_worker.forward_batch_embedding(batch)
)
ret = EmbeddingBatchResult(
embeddings=pooler_output.embeddings,
pooled_hidden_states=pooler_output.pooled_hidden_states,
can_run_cuda_graph=can_run_cuda_graph,
)
self._maybe_report_active_ranks()
+2 -2
View File
@@ -266,8 +266,8 @@ class BaseTpWorker(ABC):
self.model_runner,
return_hidden_states_before_norm=False,
)
output = self.model_runner.forward(forward_batch).logits_output
return output # Returns EmbeddingPoolerOutput
output = self.model_runner.forward(forward_batch)
return output.logits_output, output.can_run_graph
class TpModelWorker(BaseTpWorker):
+1 -4
View File
@@ -288,10 +288,7 @@ class EmbeddingBatchResult:
embeddings: torch.Tensor
pooled_hidden_states: Optional[torch.Tensor] = None
copy_done: Optional[torch.cuda.Event] = None
@property
def can_run_cuda_graph(self) -> bool:
return False
can_run_cuda_graph: bool = False
@torch.profiler.record_function("copy_embedding_to_cpu")
def copy_to_cpu(self):