EVS Framework: Support NemotronH_Nano_VL_V2 (#14051)

This commit is contained in:
Netanel Haber
2026-01-05 16:18:07 +08:00
committed by GitHub
parent b12258bfaa
commit bebd625ba1
14 changed files with 821 additions and 56 deletions
+46 -21
View File
@@ -21,8 +21,9 @@ from sglang.srt.managers.schedule_batch import (
MultimodalDataItem,
MultimodalInputs,
)
from sglang.srt.mem_cache.multimodal_cache import MultiModalStaticCache
from sglang.srt.mem_cache.multimodal_cache import EmbeddingResult, MultiModalStaticCache
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.multimodal.evs import EVSEmbeddingResult
from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import flatten_nested_list, is_npu, print_warning_once
from sglang.utils import logger
@@ -478,6 +479,11 @@ def _get_precomputed_embedding(
return None
DataEmbeddingFunc = Callable[
[List[MultimodalDataItem]], torch.Tensor | EVSEmbeddingResult
]
def get_embedding_items_per_chunk_with_extra_padding(
embedding_items_per_req: List["MultimodalDataItem"],
extend_prefix_len: int,
@@ -540,13 +546,14 @@ def get_embedding_items_per_chunk_with_extra_padding(
# TODO: To be obsoleted.
def _get_chunked_prefill_embedding(
data_embedding_func: Callable[[List[MultimodalDataItem]], torch.Tensor],
data_embedding_func: DataEmbeddingFunc,
embedding_items: List[MultimodalDataItem],
items_size: List[int],
prefix_length: List[int],
extend_length: List[int],
items_offset_list: List[List[Tuple[int, int]]],
) -> Optional[torch.Tensor]:
input_ids: torch.Tensor,
) -> tuple[torch.Tensor | None, torch.Tensor]:
# Calculate embedding for each request, try to get it from cache to avoid repeated calculation
embedding_list = []
# FIXME(Xinyuan): temporary workaround for eagle3, which may have len(items_size) > len(prefix_length)
@@ -564,7 +571,12 @@ def _get_chunked_prefill_embedding(
embedding_items_hash = MultiModalStaticCache.combine_hashes(item_hashes)
embedding_per_req = embedding_cache.get(item_hashes)
if embedding_per_req is None:
embedding_per_req = data_embedding_func(embedding_items_per_req)
embedding = data_embedding_func(embedding_items_per_req)
embedding_per_req = (
EmbeddingResult(embedding=embedding)
if isinstance(embedding, torch.Tensor)
else embedding
)
if not embedding_cache.set(embedding_items_hash, embedding_per_req):
print_warning_once(
"Multimodal embedding cache is full. This typically occurs when a single "
@@ -573,16 +585,31 @@ def _get_chunked_prefill_embedding(
"embedding size."
)
extend_prefix_len = prefix_length[i]
extend_seq_len = extend_length[i] if i < len(extend_length) else 0
if isinstance(embedding_per_req, EVSEmbeddingResult):
item = embedding_items_per_req[0]
input_ids, items_offset = (
embedding_per_req.redistribute_pruned_frames_placeholders(
input_ids,
items_offset,
item=item,
extend_prefix_len=extend_prefix_len,
extend_seq_len=extend_seq_len,
)
)
embedding_per_req_chunk, _, _ = get_embedding_chunk(
embedding=embedding_per_req,
extend_prefix_len=prefix_length[i],
extend_seq_len=extend_length[i] if i < len(extend_length) else 0,
embedding=embedding_per_req.embedding,
extend_prefix_len=extend_prefix_len,
extend_seq_len=extend_seq_len,
items_offset=items_offset,
)
embedding_list.append(embedding_per_req_chunk)
if len(embedding_list) == 0:
return None
return torch.concat(embedding_list, dim=0)
return None, input_ids
return torch.concat(embedding_list, dim=0), input_ids
def get_embedding_chunk_remove_extra_padding(
@@ -826,7 +853,7 @@ def _adjust_embedding_length(
def get_embedding_and_mask(
data_embedding_func: Callable[[List[MultimodalDataItem]], torch.Tensor],
data_embedding_func: DataEmbeddingFunc,
embedding_items: List[MultimodalDataItem],
placeholder_tensor: torch.Tensor,
input_ids: torch.Tensor,
@@ -834,7 +861,7 @@ def get_embedding_and_mask(
prefix_length: List[int],
extend_length: List[int],
items_offset_list: List[List[Tuple[int, int]]],
) -> Tuple[torch.Tensor, torch.Tensor]:
) -> Tuple[torch.Tensor | None, torch.Tensor | None, torch.Tensor]:
"""
Generate multimodal embeddings and create a mask for identifying their positions in the input sequence.
@@ -852,29 +879,31 @@ def get_embedding_and_mask(
A tuple containing:
- The generated embeddings tensor
- A boolean mask tensor indicating where these embeddings should be placed
- If EVS is used, the pruned input ids tensor; otherwise, the original input ids tensor
"""
# 1. Get embedding
embedding = _get_precomputed_embedding(
embedding_items, prefix_length, extend_length, items_offset_list
)
if embedding is None:
embedding = _get_chunked_prefill_embedding(
embedding, input_ids = _get_chunked_prefill_embedding(
data_embedding_func,
embedding_items,
items_size,
prefix_length,
extend_length,
items_offset_list,
input_ids,
)
if embedding is None:
return None, None
return None, None, input_ids
# 2. Get mask
if _is_npu:
torch.npu.current_stream().synchronize()
special_multimodal_mask = _get_multimodal_mask(input_ids, placeholder_tensor)
# 3. Adjust embedding length if needed
embedding = _adjust_embedding_length(embedding, special_multimodal_mask, logger)
return embedding, special_multimodal_mask
return embedding, special_multimodal_mask, input_ids
def embed_mm_inputs(
@@ -884,9 +913,7 @@ def embed_mm_inputs(
input_ids: torch.Tensor,
input_embedding: nn.Embedding,
multimodal_model: nn.Module = None,
data_embedding_func_mapping: Dict[
Modality, Callable[[List[MultimodalDataItem]], torch.Tensor]
] = None,
data_embedding_func_mapping: Dict[Modality, DataEmbeddingFunc] = None,
placeholder_tokens: dict[Modality, List[int]] = None,
use_deepstack: Dict[Modality, bool] = {},
) -> Optional[torch.Tensor]:
@@ -953,7 +980,7 @@ def embed_mm_inputs(
)
items_size = torch.cumsum(items_size, dim=0).tolist()
embedding, mask = get_embedding_and_mask(
embedding, mask, input_ids = get_embedding_and_mask(
data_embedding_func=embedder,
embedding_items=items,
placeholder_tensor=placeholder_tensor,
@@ -1020,9 +1047,7 @@ def general_mm_embed_routine(
forward_batch: ForwardBatch,
language_model: nn.Module,
multimodal_model: Optional[nn.Module] = None,
data_embedding_funcs: Dict[
Modality, Callable[[List[MultimodalDataItem]], torch.Tensor]
] = None,
data_embedding_funcs: Dict[Modality, DataEmbeddingFunc] = None,
placeholder_tokens: Optional[dict[Modality, List[int]]] = None,
use_deepstack: Dict[Modality, bool] = {},
**kwargs,