Fix EPD OOM by offloading precomputed_embeddings during chunked prefill (#16503)
This commit is contained in:
@@ -160,7 +160,9 @@ class EmbeddingData:
|
||||
|
||||
def get_embedding(self, is_concat=False):
|
||||
if is_concat:
|
||||
return torch.concat([embedding.cuda() for embedding in self.embedding_list])
|
||||
return torch.concat(
|
||||
[embedding.cuda() for embedding in self.embedding_list]
|
||||
).to("cpu", non_blocking=True)
|
||||
else:
|
||||
return self.embedding_list
|
||||
|
||||
|
||||
@@ -1241,6 +1241,19 @@ def general_mm_embed_routine(
|
||||
feature = getattr(mm_item, "feature", None)
|
||||
if isinstance(feature, torch.Tensor) and feature.is_cuda:
|
||||
mm_item.feature = feature.to("cpu", non_blocking=True)
|
||||
if get_global_server_args().language_only:
|
||||
precomputed_embeddings = getattr(
|
||||
mm_item, "precomputed_embeddings", None
|
||||
)
|
||||
if (
|
||||
isinstance(precomputed_embeddings, torch.Tensor)
|
||||
and precomputed_embeddings.is_cuda
|
||||
):
|
||||
mm_item.precomputed_embeddings = (
|
||||
precomputed_embeddings.to(
|
||||
"cpu", non_blocking=True
|
||||
)
|
||||
)
|
||||
forward_batch.mm_inputs = None
|
||||
forward_batch.mm_input_embeds = input_embeds
|
||||
else:
|
||||
|
||||
@@ -1639,6 +1639,14 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
||||
# The reference by CudaIpcTensorTransportProxy was cut off,
|
||||
# proactively delete to avoid slow gc.
|
||||
del pixel_values
|
||||
if get_global_server_args().language_only:
|
||||
precomputed_embeddings = getattr(
|
||||
mm_item, "precomputed_embeddings", None
|
||||
)
|
||||
if isinstance(precomputed_embeddings, torch.Tensor):
|
||||
mm_item.precomputed_embeddings = precomputed_embeddings.to(
|
||||
self.device, non_blocking=True
|
||||
)
|
||||
self.multimodal_inputs = multimodal_inputs
|
||||
self.token_type_ids = token_type_ids_tensor
|
||||
self.seq_lens_sum = sum(seq_lens)
|
||||
|
||||
Reference in New Issue
Block a user