From b48e2cb1ebc350cd0dadf1b38639ae301a4ea387 Mon Sep 17 00:00:00 2001 From: Xinyuan Tong Date: Sat, 19 Sep 2026 04:04:22 +0000 Subject: [PATCH] model: TP-wide single-owner image encoding for DeepSeek V4.1 ViT and Aligner are replicated per TP rank, encoding each image eight times with TP8 on both CP1 and CP8. Elect one owner per image and use ordered full-span broadcasts with a six-phase agreement protocol. Both CP1 and CP8 benefit while local cache hits preserve collective order. --- .../sglang/srt/managers/mm_owner_embedding.py | 486 ++++++++ python/sglang/srt/managers/mm_schedule.py | 120 +- python/sglang/srt/managers/mm_utils.py | 15 +- python/sglang/srt/models/deepseek_v4.py | 85 +- .../unit/managers/test_mm_owner_embedding.py | 1010 +++++++++++++++++ .../test_deepseek_v41_vision_cp_inputs.py | 1 + 6 files changed, 1663 insertions(+), 54 deletions(-) create mode 100644 python/sglang/srt/managers/mm_owner_embedding.py create mode 100644 test/registered/unit/managers/test_mm_owner_embedding.py diff --git a/python/sglang/srt/managers/mm_owner_embedding.py b/python/sglang/srt/managers/mm_owner_embedding.py new file mode 100644 index 000000000..7db767bef --- /dev/null +++ b/python/sglang/srt/managers/mm_owner_embedding.py @@ -0,0 +1,486 @@ +"""One owner rank encodes each image span and broadcasts it to the ranks that +run the same prefill chunk; every agreement precedes the payload it guards.""" + +from __future__ import annotations + +import logging +from contextlib import contextmanager +from typing import Any, Callable, Dict, Iterator, List, Optional, Sequence, Tuple + +import msgspec +import torch + +from sglang.srt.distributed.device_communicators.pynccl_allocator import ( + disable_symmetric_memory_context, + restore_symmetric_memory_context, +) +from sglang.srt.mem_cache.multimodal_cache import EmbeddingResult, MultiModalStaticCache + +logger = logging.getLogger(__name__) + +SpanKey = Tuple[Optional[int], int] +SpanEncoder = Callable[[List[Any]], torch.Tensor | List[torch.Tensor]] +SpanSignature = Callable[[Any, int], Tuple[Any, ...]] + +LOCAL_HIT = 0 +OWNER_CACHE_BROADCAST = 1 +OWNER_ENCODE_BROADCAST = 2 + +PHASE_PREPARE = "prepare" +PHASE_FEATURES = "features" +PHASE_FINALIZE = "finalize" + + +class MmOwnerProtocolError(RuntimeError): + """Raised with identical text on every group member after a group-agreed failure.""" + + +class ImageSpanRequest(msgspec.Struct, frozen=True): + hash: Optional[int] + span_len: int + item: Any + inside_chunk: bool + duplicates: List[Any] = [] + + +class ImageSpanKey(msgspec.Struct, frozen=True): + hash: Optional[int] + span_len: int + geometry: Optional[Tuple[Any, ...]] + + +class RankManifest(msgspec.Struct, frozen=True): + rank: int + keys: List[ImageSpanKey] + cached: List[bool] + dtype: str + width: int + rids: List[str] + error: Optional[str] = None + + +class OwnerPlan(msgspec.Struct, frozen=True): + actions: List[int] + owners: List[int] + error: Optional[str] = None + + +class RankStatus(msgspec.Struct, frozen=True): + rank: int + error: Optional[str] = None + + +def select_owner_group(parallel) -> Optional[Any]: + """The group whose members all execute the same requests, or None when a + single rank already encodes every image it sees.""" + replication = parallel.tp_size // parallel.attn_dp_size + if replication <= 1: + return None + if parallel.attn_cp_size == 1: + group = parallel.attn_tp_group + elif parallel.attn_dp_size == 1 and parallel.attn_cp_size == parallel.tp_size: + group = parallel.attn_cp_group + else: + return None + return group if group.world_size == replication else None + + +def has_owner_span_work( + mm_inputs: Sequence[Any], + extend_prefix_lens: Sequence[int], + extend_seq_lens: Sequence[int], +) -> bool: + """Host-side mirror of the per-image scheduling path: does any raw + single-span image overlap the chunk on every rank of the group.""" + for mm_input, prefix_len, extend_len in zip( + mm_inputs, extend_prefix_lens, extend_seq_lens + ): + if mm_input is None or extend_len <= 0: + continue + items = [item for item in mm_input.mm_items if item is not None] + if not items or any( + item.precomputed_embeddings is not None or len(item.offsets) != 1 + for item in items + ): + continue + for item in items: + start, end = item.offsets[0] + if end >= prefix_len and start < prefix_len + extend_len: + return True + return False + + +class MmOwnerSession(msgspec.Struct): + group: Any + device: Any + dtype: Any + width: int + rids: List[str] + signature: Any + engaged: bool + phase: str = PHASE_PREPARE + in_collective: bool = False + + def resolve( + self, + requests: Sequence[ImageSpanRequest], + cache: MultiModalStaticCache, + encode: SpanEncoder, + ) -> Dict[SpanKey, torch.Tensor]: + if not self.engaged: + raise RuntimeError( + "owner protocol reached for a chunk whose host metadata has no image span" + ) + # Owners allocate different amounts than receivers, so none of these + # buffers may come out of a symmetric pool. + saved_context = disable_symmetric_memory_context() + try: + return _resolve_owner_features(self, requests, cache, encode) + finally: + restore_symmetric_memory_context(saved_context) + + def features_ready(self) -> None: + self._complete() + self.phase = PHASE_FINALIZE + + @contextmanager + def uncaptured(self) -> Iterator[None]: + # A failure inside a collective leaves the group in an unknown state; + # no later exchange may try to agree on it. + self.in_collective = True + yield + self.in_collective = False + + @contextmanager + def fence(self) -> Iterator[None]: + try: + yield + except Exception as exc: + self._fail(exc) + raise + self._complete() + + def _fail(self, exc: BaseException) -> None: + if ( + not self.engaged + or self.in_collective + or isinstance(exc, MmOwnerProtocolError) + ): + raise exc + text = _describe(self, self.phase, exc) + if self.phase == PHASE_PREPARE: + try: + _exchange_manifest(self, _manifest(self, [], [], error=text)) + except MmOwnerProtocolError as agreed: + raise agreed from exc + _exchange_status(self, text, exc) + + def _complete(self) -> None: + if not self.engaged: + return + if self.phase == PHASE_PREPARE: + raise RuntimeError( + f"owner protocol {self.phase} completed without a manifest exchange" + ) + error = None + cause = None + try: + _synchronize(self.device) + except Exception as exc: + cause = exc + error = _describe(self, self.phase, exc) + _exchange_status(self, error, cause) + + +def _manifest( + session: MmOwnerSession, + keys: List[ImageSpanKey], + cached: List[bool], + error: Optional[str] = None, +) -> RankManifest: + return RankManifest( + rank=session.group.rank_in_group, + keys=keys, + cached=cached, + dtype=str(session.dtype), + width=session.width, + rids=list(session.rids), + error=error, + ) + + +def _exchange_manifest(session: MmOwnerSession, manifest: RankManifest) -> OwnerPlan: + group = session.group + with session.uncaptured(): + manifests = group.all_gather_object(manifest) + plan = _plan_or_error(session, manifests) if group.rank_in_group == 0 else None + plan = group.broadcast_object(plan, src=0) + session.phase = PHASE_FEATURES + if plan.error is not None: + raise MmOwnerProtocolError(plan.error) + return plan + + +def _plan_or_error(session: MmOwnerSession, manifests: List[RankManifest]) -> OwnerPlan: + try: + return _make_plan(manifests) + except Exception as exc: + return OwnerPlan(actions=[], owners=[], error=_describe(session, "plan", exc)) + + +def _exchange_status( + session: MmOwnerSession, error: Optional[str], cause: Optional[BaseException] +) -> None: + with session.uncaptured(): + statuses = session.group.all_gather_object( + RankStatus(rank=session.group.rank_in_group, error=error) + ) + _raise_first_error(statuses, cause) + + +def _resolve_owner_features( + session: MmOwnerSession, + requests: Sequence[ImageSpanRequest], + cache: MultiModalStaticCache, + encode: SpanEncoder, +) -> Dict[SpanKey, torch.Tensor]: + group = session.group + features: Dict[SpanKey, torch.Tensor] = {} + keys: List[ImageSpanKey] = [] + cached: List[bool] = [] + error = None + try: + keys, cached = _pin_local_cache(session, requests, cache, features) + except Exception as exc: + error = _describe(session, "manifest", exc) + plan = _exchange_manifest(session, _manifest(session, keys, cached, error)) + + if all(action == LOCAL_HIT for action in plan.actions): + return features + + buffers: Dict[int, torch.Tensor] = {} + error = None + try: + buffers = _prepare_transfers(session, requests, keys, plan, features, encode) + _synchronize(session.device) + except Exception as exc: + error = _describe(session, "encode", exc) + _exchange_status(session, error, None) + + with session.uncaptured(): + for index, (action, owner) in enumerate(zip(plan.actions, plan.owners)): + if action != LOCAL_HIT: + group.broadcast(buffers[index], src=owner) + + for index, key in enumerate(keys): + if plan.actions[index] == LOCAL_HIT: + continue + span = buffers[index] + features[(key.hash, key.span_len)] = span + cache.set(key.hash, EmbeddingResult(embedding=span)) + return features + + +def _pin_local_cache( + session: MmOwnerSession, + requests: Sequence[ImageSpanRequest], + cache: MultiModalStaticCache, + features: Dict[SpanKey, torch.Tensor], +) -> Tuple[List[ImageSpanKey], List[bool]]: + keys: List[ImageSpanKey] = [] + cached: List[bool] = [] + for request in requests: + if request.hash is None: + raise ValueError( + f"image span of {request.span_len} tokens has no content hash" + ) + geometry = session.signature(request.item, request.span_len) + for duplicate in request.duplicates: + other = session.signature(duplicate, request.span_len) + if other != geometry: + raise ValueError( + f"image hash {request.hash} ({request.span_len} tokens) occurs " + f"with different geometry: {geometry} vs {other}" + ) + keys.append( + ImageSpanKey( + hash=request.hash, span_len=request.span_len, geometry=geometry + ) + ) + span = _valid_cached_span(session, cache, request) + if span is not None: + features[(request.hash, request.span_len)] = span + cached.append(span is not None) + return keys, cached + + +def _valid_cached_span( + session: MmOwnerSession, + cache: MultiModalStaticCache, + request: ImageSpanRequest, +) -> Optional[torch.Tensor]: + entry = cache.get_single(request.hash) + if entry is None: + return None + span = entry.embedding + if ( + span.dim() == 2 + and span.shape[0] == request.span_len + and span.shape[1] == session.width + and span.dtype == session.dtype + and span.device == session.device + ): + return span + logger.warning( + "Discarding cached multimodal embedding that cannot serve the current " + "image span: cache_key=%s expected=(%d, %d, %s) cached=(%s, %s).", + request.hash, + request.span_len, + session.width, + session.dtype, + tuple(span.shape), + span.dtype, + ) + cache.free(request.hash, None) + return None + + +def _make_plan(manifests: List[RankManifest]) -> OwnerPlan: + for manifest in manifests: + if manifest.error is not None: + return OwnerPlan(actions=[], owners=[], error=manifest.error) + lead = manifests[0] + for manifest in manifests[1:]: + if (manifest.keys, manifest.dtype, manifest.width, manifest.rids) != ( + lead.keys, + lead.dtype, + lead.width, + lead.rids, + ): + return OwnerPlan( + actions=[], + owners=[], + error=( + "image manifest mismatch between group ranks 0 and " + f"{manifest.rank}: rids={lead.rids} vs {manifest.rids}, " + f"keys={lead.keys} vs {manifest.keys}, " + f"dtype={lead.dtype} vs {manifest.dtype}, " + f"width={lead.width} vs {manifest.width}" + ), + ) + replication = len(manifests) + actions: List[int] = [] + owners: List[int] = [] + for index, key in enumerate(lead.keys): + owner = key.hash % replication + if all(manifest.cached[index] for manifest in manifests): + action = LOCAL_HIT + elif manifests[owner].cached[index]: + action = OWNER_CACHE_BROADCAST + else: + action = OWNER_ENCODE_BROADCAST + actions.append(action) + owners.append(owner) + return OwnerPlan(actions=actions, owners=owners) + + +def _prepare_transfers( + session: MmOwnerSession, + requests: Sequence[ImageSpanRequest], + keys: List[ImageSpanKey], + plan: OwnerPlan, + features: Dict[SpanKey, torch.Tensor], + encode: SpanEncoder, +) -> Dict[int, torch.Tensor]: + rank = session.group.rank_in_group + buffers: Dict[int, torch.Tensor] = {} + owned: List[int] = [] + for index, (action, owner) in enumerate(zip(plan.actions, plan.owners)): + if action == LOCAL_HIT: + continue + if owner != rank: + try: + buffers[index] = _new_span_buffer(session, keys[index]) + except Exception as exc: + raise RuntimeError( + f"receive buffer for image hash {keys[index].hash} shape " + f"{(keys[index].span_len, session.width)} {session.dtype} " + f"failed: {type(exc).__name__}: {exc}" + ) from exc + elif action == OWNER_CACHE_BROADCAST: + key = (keys[index].hash, keys[index].span_len) + buffers[index] = features[key].contiguous() + else: + owned.append(index) + if owned: + owned_hashes = [keys[index].hash for index in owned] + try: + encoded = encode([requests[index].item for index in owned]) + except Exception as exc: + raise RuntimeError( + f"owner encode of image hashes {owned_hashes} failed: " + f"{type(exc).__name__}: {exc}" + ) from exc + spans = _split_spans(encoded, [keys[index].span_len for index in owned]) + for index, span in zip(owned, spans): + buffers[index] = _validated_span(session, keys[index], span) + return buffers + + +def _new_span_buffer(session: MmOwnerSession, key: ImageSpanKey) -> torch.Tensor: + return torch.empty( + (key.span_len, session.width), device=session.device, dtype=session.dtype + ) + + +def _split_spans( + encoded: torch.Tensor | List[torch.Tensor], span_lens: List[int] +) -> List[torch.Tensor]: + if isinstance(encoded, list): + if len(encoded) != len(span_lens): + raise ValueError( + f"encoder returned {len(encoded)} spans for {len(span_lens)} images" + ) + return [span.reshape(-1, span.shape[-1]) for span in encoded] + encoded = encoded.reshape(-1, encoded.shape[-1]) + if encoded.shape[0] != sum(span_lens): + raise ValueError( + f"encoder returned {encoded.shape[0]} rows for spans of {span_lens}" + ) + return list(torch.split(encoded, span_lens, dim=0)) + + +def _validated_span( + session: MmOwnerSession, key: ImageSpanKey, span: torch.Tensor +) -> torch.Tensor: + expected = (key.span_len, session.width) + if tuple(span.shape) != expected or span.dtype != session.dtype: + raise ValueError( + f"encoded span for hash={key.hash} has shape {tuple(span.shape)} " + f"dtype {span.dtype}; expected {expected} {session.dtype}" + ) + if span.device != session.device: + span = span.to(session.device) + return span.contiguous() + + +def _synchronize(device) -> None: + if device.type == "cuda": + torch.cuda.current_stream(device).synchronize() + + +def _describe(session: MmOwnerSession, stage: str, exc: BaseException) -> str: + return ( + f"multimodal owner protocol failed during {stage} on group rank " + f"{session.group.rank_in_group} (global rank " + f"{session.group.ranks[session.group.rank_in_group]}, rids={list(session.rids)}): " + f"{type(exc).__name__}: {exc}" + ) + + +def _raise_first_error( + statuses: List[RankStatus], cause: Optional[BaseException] +) -> None: + for status in statuses: + if status.error is not None: + raise MmOwnerProtocolError(status.error) from cause diff --git a/python/sglang/srt/managers/mm_schedule.py b/python/sglang/srt/managers/mm_schedule.py index 13dc18697..067e43330 100644 --- a/python/sglang/srt/managers/mm_schedule.py +++ b/python/sglang/srt/managers/mm_schedule.py @@ -5,6 +5,7 @@ from typing import Callable, Dict, List, Optional, Tuple import torch +from sglang.srt.managers.mm_owner_embedding import ImageSpanRequest, MmOwnerSession from sglang.srt.managers.schedule_batch import MultimodalDataItem from sglang.srt.mem_cache.multimodal_cache import EmbeddingResult, MultiModalStaticCache from sglang.srt.multimodal.evs import EVSEmbeddingResult @@ -339,43 +340,24 @@ def _batch_encode_per_image_misses( unique_misses: Dict[Tuple[Optional[int], int], Tuple[MultimodalDataItem, int]] = {} hash_to_embedding: Dict[Tuple[Optional[int], int], torch.Tensor] = {} - # Phase 1a: find overlapping items per request and collect cache misses - for req_info in per_image_requests: - chunk_start = req_info.extend_prefix_len - chunk_end = chunk_start + req_info.extend_seq_len # exclusive - overlapping = [] - if req_info.extend_seq_len > 0: - for idx, (item, (start, end)) in enumerate( - zip(req_info.items, req_info.items_offset) - ): - if end >= chunk_start and start < chunk_end: - overlapping.append((idx, item, start, end)) - req_info.overlapping = overlapping - - for _idx, item, start, end in overlapping: - expected_token_count = end - start + 1 - cache_key = (item.hash, expected_token_count) - if cache_key in hash_to_embedding: + # Phase 1a: collect cache misses over the unique overlapping spans + for span in _collect_image_span_requests(per_image_requests): + cache_key = (span.hash, span.span_len) + cached = embedding_cache.get_single(span.hash) + if cached is not None: + cached_embedding = cached.embedding + cached_token_count = _embedding_token_count(cached_embedding) + if cached_token_count == span.span_len: + hash_to_embedding[cache_key] = cached_embedding continue - cached = embedding_cache.get_single(item.hash) - if cached is not None: - cached_embedding = cached.embedding - cached_token_count = _embedding_token_count(cached_embedding) - if cached_token_count == expected_token_count: - hash_to_embedding[cache_key] = cached_embedding - else: - _discard_mismatched_cached_embedding( - item.hash, expected_token_count, cached_token_count - ) - unique_misses[cache_key] = (item, expected_token_count) - elif cache_key not in unique_misses: - if ( - start >= chunk_start - and end < chunk_end - and item.can_defer_cuda_ipc_feature_reconstruction() - ): - item.model_specific_data[BORROW_CUDA_IPC_FEATURE_KEY] = True - unique_misses[cache_key] = (item, expected_token_count) + _discard_mismatched_cached_embedding( + span.hash, span.span_len, cached_token_count + ) + elif ( + span.inside_chunk and span.item.can_defer_cuda_ipc_feature_reconstruction() + ): + span.item.model_specific_data[BORROW_CUDA_IPC_FEATURE_KEY] = True + unique_misses[cache_key] = (span.item, span.span_len) # Phase 1b: single ViT call for all unique cache misses if unique_misses: @@ -412,6 +394,52 @@ def _batch_encode_per_image_misses( return hash_to_embedding +def _collect_image_span_requests( + per_image_requests: List[PerImageRequestInfo], +) -> List[ImageSpanRequest]: + spans: Dict[ + Tuple[Optional[int], int], + Tuple[MultimodalDataItem, bool, List[MultimodalDataItem]], + ] = {} + for req_info in per_image_requests: + chunk_start = req_info.extend_prefix_len + chunk_end = chunk_start + req_info.extend_seq_len # exclusive + overlapping = [] + if req_info.extend_seq_len > 0: + for idx, (item, (start, end)) in enumerate( + zip(req_info.items, req_info.items_offset) + ): + if end >= chunk_start and start < chunk_end: + overlapping.append((idx, item, start, end)) + req_info.overlapping = overlapping + + for _idx, item, start, end in overlapping: + cache_key = (item.hash, end - start + 1) + if cache_key in spans: + spans[cache_key][2].append(item) + continue + spans[cache_key] = (item, start >= chunk_start and end < chunk_end, []) + return [ + ImageSpanRequest( + hash=item_hash, + span_len=span_len, + item=item, + inside_chunk=inside_chunk, + duplicates=duplicates, + ) + for (item_hash, span_len), (item, inside_chunk, duplicates) in spans.items() + ] + + +def _owner_span_encoder(data_embedding_func: DataEmbeddingFunc, device: torch.device): + def encode(items: List[MultimodalDataItem]): + if not _can_skip_pre_embed_feature_move(data_embedding_func): + _move_items_to_device(items, device) + return data_embedding_func(items) + + return encode + + def _get_chunked_embedding_by_item( data_embedding_func: DataEmbeddingFunc, embedding_items_per_req: List[MultimodalDataItem], @@ -537,6 +565,7 @@ def _get_chunked_prefill_embedding( extend_length: List[int], items_offset_list: List[List[Tuple[int, int]]], input_ids: torch.Tensor, + mm_owner: Optional[MmOwnerSession] = None, ) -> tuple[torch.Tensor | None, torch.Tensor]: """ Chunked prefill embedding: encode items across all requests and extract @@ -598,7 +627,22 @@ def _get_chunked_prefill_embedding( # Phase 1: batch encode all per-image cache misses in ONE ViT call hash_to_embedding: Dict[Tuple[Optional[int], int], torch.Tensor] = {} - if per_image_requests: + if per_image_requests and mm_owner is not None: + # The owner protocol must see every overlapping span before any local + # cache filtering: a rank-local hit can never skip a group collective. + span_requests = _collect_image_span_requests(per_image_requests) + if mm_owner.engaged: + hash_to_embedding = mm_owner.resolve( + span_requests, + cache=embedding_cache, + encode=_owner_span_encoder(data_embedding_func, device), + ) + elif span_requests: + raise RuntimeError( + "owner eligibility saw no image span in this chunk, but " + f"scheduling found {len(span_requests)}" + ) + elif per_image_requests: hash_to_embedding = _batch_encode_per_image_misses( data_embedding_func, per_image_requests, device ) @@ -701,6 +745,7 @@ def get_embedding_and_mask( prefix_length: List[int], extend_length: List[int], items_offset_list: List[List[Tuple[int, int]]], + mm_owner: Optional[MmOwnerSession] = None, ) -> Tuple[torch.Tensor | None, torch.Tensor | None, torch.Tensor]: """ Generate multimodal embeddings and create a mask for identifying their positions in the input sequence. @@ -741,6 +786,7 @@ def get_embedding_and_mask( extend_length, items_offset_list, input_ids, + mm_owner=mm_owner, ) if embedding is None: return None, None, input_ids diff --git a/python/sglang/srt/managers/mm_utils.py b/python/sglang/srt/managers/mm_utils.py index e07bfb8c4..43f1b5d61 100644 --- a/python/sglang/srt/managers/mm_utils.py +++ b/python/sglang/srt/managers/mm_utils.py @@ -10,6 +10,7 @@ import pickle import sys from abc import abstractmethod from collections import defaultdict +from contextlib import nullcontext from multiprocessing import shared_memory from typing import Any, Dict, List, Optional, Tuple @@ -24,6 +25,7 @@ from sglang.srt.managers.io_struct import ( TokenizedEmbeddingReqInput, TokenizedGenerateReqInput, ) +from sglang.srt.managers.mm_owner_embedding import MmOwnerSession # Preserve the existing initialization import for downstream callers. from sglang.srt.managers.mm_schedule import ( @@ -397,6 +399,7 @@ def embed_mm_inputs( data_embedding_func_mapping: Dict[Modality, DataEmbeddingFunc] = None, placeholder_tokens: dict[Modality, List[int]] = None, use_deepstack: Dict[Modality, bool] = {}, + mm_owner: Optional[MmOwnerSession] = None, ) -> Optional[torch.Tensor]: """ Embed multimodal inputs and integrate them with text token embeddings. @@ -478,6 +481,7 @@ def embed_mm_inputs( prefix_length=extend_prefix_lens, extend_length=extend_seq_lens, items_offset_list=items_offsets, + mm_owner=mm_owner, ) if use_deepstack.get(modality, None) and embedding is not None: @@ -498,7 +502,12 @@ def embed_mm_inputs( # filled with the hash values of the multimodal for the prefix matching in the radix attention. # There values are useless because their embeddings will be replaced by vision embeddings anyway. input_ids.clamp_(min=0, max=vocab_size - 1) - input_embeds = input_embedding(input_ids) + if mm_owner is not None: + # The text embedding may all-reduce across TP; a rank-local failure in + # feature preparation has to be agreed on before any rank enters it. + mm_owner.features_ready() + with mm_owner.uncaptured() if mm_owner is not None else nullcontext(): + input_embeds = input_embedding(input_ids) # deepstack embedding if use_deepstack: @@ -525,7 +534,9 @@ def embed_mm_inputs( _scatter_mm_embedding(dest=input_embeds, mask=mask, src=embedding) if use_deepstack.get(modality, None): _scatter_mm_embedding( - dest=input_deepstack_embeds, mask=mask, src=deepstack_embeddings[i] + dest=input_deepstack_embeds, + mask=mask, + src=deepstack_embeddings[i], ) return input_embeds, other_info diff --git a/python/sglang/srt/models/deepseek_v4.py b/python/sglang/srt/models/deepseek_v4.py index 5e8f4bcc7..35f3e9ae7 100644 --- a/python/sglang/srt/models/deepseek_v4.py +++ b/python/sglang/srt/models/deepseek_v4.py @@ -121,6 +121,11 @@ from sglang.srt.layers.quantization.mxfp8_input import Mxfp8SwizzledInput from sglang.srt.layers.rotary_embedding import get_rope_wrapper from sglang.srt.layers.utils import PPMissingLayer, get_layer_id from sglang.srt.layers.vocab_parallel_embedding import VocabParallelEmbedding +from sglang.srt.managers.mm_owner_embedding import ( + MmOwnerSession, + has_owner_span_work, + select_owner_group, +) from sglang.srt.managers.mm_utils import ( MultiModalityDataPaddingPatternMultimodalTokens, embed_mm_inputs, @@ -4917,6 +4922,13 @@ class DeepseekV4ForCausalLM(nn.Module): self.image_start = nn.Parameter(torch.empty(config.hidden_size)) self.image_end = nn.Parameter(torch.empty(config.hidden_size)) self.image_newline = nn.Parameter(torch.empty(config.hidden_size)) + # Ranks of this group run identical image chunks; one owner encodes each + # span and broadcasts it. None keeps the replicated encoder. + self.mm_owner_group = ( + select_owner_group(get_parallel()) + if self.vision is not None and _is_cuda + else None + ) self.model = DeepseekV4Model( config, quant_config, prefix=add_prefix("model", prefix) ) @@ -5051,7 +5063,42 @@ class DeepseekV4ForCausalLM(nn.Module): spans.append(span) return spans - def _prepare_mm_embeddings(self, input_ids, forward_batch): + def _image_span_signature(self, item, span_len: int): + h, w = int(item.n_vit_h), int(item.n_vit_w) + r = self.config.vision_downsample_ratio + expected = len(image_token_types((h + r - 1) // r, (w + r - 1) // r)) + if expected != span_len: + raise ValueError( + f"image grid {(h, w)} yields {expected} span tokens, " + f"placeholder has {span_len}" + ) + plan = item.model_specific_data.get(GPU_PLAN_KEY) + feature = item.feature + return ( + h, + w, + tuple(feature.shape) if isinstance(feature, torch.Tensor) else None, + None if plan is None else tuple(sorted(plan.items())), + ) + + def _mm_owner_session(self, forward_batch) -> Optional[MmOwnerSession]: + if self.mm_owner_group is None: + return None + return MmOwnerSession( + group=self.mm_owner_group, + device=self.image_start.device, + dtype=self.image_start.dtype, + width=self.config.hidden_size, + rids=list(forward_batch.rids or ()), + signature=self._image_span_signature, + engaged=has_owner_span_work( + forward_batch.mm_inputs, + forward_batch.extend_prefix_lens_cpu, + forward_batch.extend_seq_lens_cpu, + ), + ) + + def _prepare_mm_embeddings(self, input_ids, forward_batch, mm_owner): # Keep scheduler hash IDs intact: the shared embedder clamps its input in place. input_embeds, _ = embed_mm_inputs( mm_inputs_list=[ @@ -5063,6 +5110,7 @@ class DeepseekV4ForCausalLM(nn.Module): input_ids=input_ids.clone(), input_embedding=self.get_input_embeddings(), multimodal_model=self, + mm_owner=mm_owner, ) forward_batch.mm_input_embeds = input_embeds return input_embeds @@ -5078,24 +5126,31 @@ class DeepseekV4ForCausalLM(nn.Module): ) -> Tuple[torch.Tensor, Optional[torch.Tensor]]: if self.vision is None: return input_ids, input_embeds - if ( + has_images = ( not forward_batch.forward_mode.is_decode() and not forward_batch.forward_mode.is_target_verify() and forward_batch.mm_inputs is not None and any(x is not None for x in forward_batch.mm_inputs) - ): - if input_embeds is not None: - raise ValueError("Cannot combine input_embeds and image inputs") - input_embeds = self._prepare_mm_embeddings(input_ids, forward_batch) - if not ( - forward_batch.forward_mode.is_decode_or_idle() - or forward_batch.forward_mode.is_target_verify() - ): - # Decode/verify IDs are already vocabulary IDs; remap prompt image - # hashes for Engram and routing. - input_ids = input_ids.masked_fill( - input_ids >= MM_PAD_SHIFT_VALUE, self.config.image_token_id - ) + ) + if has_images and input_embeds is not None: + raise ValueError("Cannot combine input_embeds and image inputs") + mm_owner = self._mm_owner_session(forward_batch) if has_images else None + # Peers may only enter the body or the CP shard once every rank has + # finished all of its fallible input preparation, the remap included. + with mm_owner.fence() if mm_owner is not None else nullcontext(): + if has_images: + input_embeds = self._prepare_mm_embeddings( + input_ids, forward_batch, mm_owner + ) + if not ( + forward_batch.forward_mode.is_decode_or_idle() + or forward_batch.forward_mode.is_target_verify() + ): + # Decode/verify IDs are already vocabulary IDs; remap prompt image + # hashes for Engram and routing. + input_ids = input_ids.masked_fill( + input_ids >= MM_PAD_SHIFT_VALUE, self.config.image_token_id + ) return input_ids, input_embeds def set_dspark_layers_to_capture(self, layer_ids: List[int]) -> None: diff --git a/test/registered/unit/managers/test_mm_owner_embedding.py b/test/registered/unit/managers/test_mm_owner_embedding.py new file mode 100644 index 000000000..e4ef8b00c --- /dev/null +++ b/test/registered/unit/managers/test_mm_owner_embedding.py @@ -0,0 +1,1010 @@ +"""Owner-encoded image spans: real gloo groups, production coordinator and entrypoints.""" + +import json +import sys +from contextlib import nullcontext +from datetime import timedelta +from pathlib import Path +from tempfile import TemporaryDirectory +from types import SimpleNamespace +from unittest.mock import patch + +import pytest +import torch +import torch.distributed as dist +import torch.multiprocessing +from torch import nn + +from sglang.srt.distributed import parallel_state +from sglang.srt.layers.cp.base import init_cp_strategy +from sglang.srt.layers.cp.utils import prepare_cp_forward +from sglang.srt.managers import mm_owner_embedding, mm_schedule, mm_utils +from sglang.srt.managers.mm_owner_embedding import ( + MmOwnerProtocolError, + select_owner_group, +) +from sglang.srt.managers.schedule_batch import ( + Modality, + MultimodalDataItem, + MultimodalInputs, +) +from sglang.srt.model_executor.forward_batch_info import ForwardMode +from sglang.srt.model_executor.runner.eager_runner import EagerRunner +from sglang.srt.models.deepseek_v4 import DeepseekV4ForCausalLM +from sglang.srt.runtime_context import get_parallel +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=60, suite="base-a-test-cpu") + +HIDDEN = 8 +VOCAB = 64 +IMAGE_TOKEN_ID = 7 +CACHE_TAG = 50 +DOWNSAMPLE = 2 +GLOO_TIMEOUT = timedelta(seconds=90) + + +def _span(item_hash: int, rows: int, tag: int) -> torch.Tensor: + """Distinguishable full span: the producing rank is visible in the values.""" + base = float(item_hash * 1000 + tag * 100) + return torch.arange(rows * HIDDEN, dtype=torch.float32).view(rows, HIDDEN) + base + + +def _grid(rows: int): + """A ViT grid whose downsampled span has exactly ``rows`` tokens.""" + return (0, 0) if rows == 2 else (DOWNSAMPLE, DOWNSAMPLE * (rows - 3)) + + +def _image(item_hash: int, start: int, rows: int, grid=None) -> MultimodalDataItem: + h, w = _grid(rows) if grid is None else grid + item = MultimodalDataItem( + modality=Modality.IMAGE, + feature=torch.zeros(1), + offsets=[(start, start + rows - 1)], + model_specific_data={"n_vit_h": h, "n_vit_w": w}, + ) + item.set_hash(item_hash) + return item + + +def _request(parts, prefix_len: int = 0, extend_len=None, rid: str = "rid"): + """``parts`` mixes token ids with ``(hash, rows)`` or ``(hash, rows, grid)`` images.""" + ids, items = [], [] + for part in parts: + if isinstance(part, tuple): + item = _image(part[0], len(ids), part[1], *part[2:]) + items.append(item) + ids.extend([item.pad_value] * part[1]) + else: + ids.append(part) + if extend_len is None: + extend_len = len(ids) - prefix_len + return SimpleNamespace( + ids=ids, items=items, prefix_len=prefix_len, extend_len=extend_len, rid=rid + ) + + +def _chunk_ids(req): + return req.ids[req.prefix_len : req.prefix_len + req.extend_len] + + +def _batch(requests, mode=ForwardMode.EXTEND): + return SimpleNamespace( + forward_mode=mode, + mm_inputs=[ + MultimodalInputs(mm_items=req.items, im_token_id=IMAGE_TOKEN_ID) + if req.items + else None + for req in requests + ], + extend_prefix_lens_cpu=[req.prefix_len for req in requests], + extend_seq_lens_cpu=[req.extend_len for req in requests], + seq_lens_cpu=[req.prefix_len + req.extend_len for req in requests], + input_ids=torch.tensor( + [tok for req in requests for tok in _chunk_ids(req)], dtype=torch.long + ), + positions=torch.cat( + [ + torch.arange(req.prefix_len, req.prefix_len + req.extend_len) + for req in requests + ] + ), + mm_input_embeds=None, + attn_cp_metadata=None, + global_num_tokens_cpu=None, + out_cache_loc=None, + input_ids_global=torch.zeros(1, dtype=torch.long), + rids=[req.rid for req in requests], + ) + + +def _expected_embeds(embed: nn.Embedding, requests, tags): + """``tags`` maps image hash to the tag of the span every rank must end up with.""" + rows = [] + with torch.no_grad(): + for req in requests: + ids = torch.tensor(_chunk_ids(req), dtype=torch.long) + chunk = nn.functional.embedding(ids.clamp(max=VOCAB - 1), embed.weight) + chunk_end = req.prefix_len + req.extend_len + for item in req.items: + start, end = item.offsets[0] + lo = max(start, req.prefix_len) - start + hi = min(end + 1, chunk_end) - start + if lo >= hi: + continue + span = _span(item.hash, end - start + 1, tags[item.hash]) + dest = max(start, req.prefix_len) - req.prefix_len + chunk[dest : dest + hi - lo] = span[lo:hi] + rows.append(chunk) + return torch.cat(rows) + + +def _seed_cache(item_hash: int, rows: int, tag: int) -> None: + mm_schedule.embedding_cache.set( + item_hash, mm_schedule.EmbeddingResult(embedding=_span(item_hash, rows, tag)) + ) + + +class _RecordingBody: + def __init__(self, embed): + self.embed = embed + self.calls = [] + + def get_input_embeddings(self): + return self.embed + + def __call__(self, input_ids, positions, forward_batch, input_embeds=None): + self.calls.append( + SimpleNamespace(input_ids=input_ids, input_embeds=input_embeds) + ) + return input_embeds, input_embeds + + +class _ReducingEmbedding(nn.Embedding): + """Stands in for the TP vocab embedding: a real all-reduce on the device group.""" + + def __init__(self, coordinator): + torch.manual_seed(0) + super().__init__(VOCAB, HIDDEN) + self.coordinator = coordinator + self.calls = 0 + + def forward(self, input_ids): + self.calls += 1 + if self.coordinator is not None: + dist.all_reduce(torch.zeros(1), group=self.coordinator.device_group) + return super().forward(input_ids) + + +class _VisionStub(DeepseekV4ForCausalLM): + def __init__(self, embed, owner_group, tag: int): + nn.Module.__init__(self) + self.config = SimpleNamespace( + image_token_id=IMAGE_TOKEN_ID, + hidden_size=HIDDEN, + vision_downsample_ratio=DOWNSAMPLE, + ) + self.vision = object() + self.tp_size = 1 + self.mm_owner_group = owner_group + self.image_start = nn.Parameter(torch.zeros(HIDDEN)) + self.model = _RecordingBody(embed) + self.pp_group = SimpleNamespace(is_last_rank=True) + self.lm_head = object() + self.capture_aux_hidden_states = False + self.logits_calls = [] + self.tag = tag + self.encoded = [] + self.fail_hashes = set() + + def get_image_feature(self, items): + hashes = [item.hash for item in items] + if self.fail_hashes.intersection(hashes): + raise RuntimeError("injected encoder failure") + self.encoded.extend(hashes) + return [ + _span(item.hash, item.offsets[0][1] - item.offsets[0][0] + 1, self.tag) + for item in items + ] + + def logits_processor( + self, + input_ids, + hidden_states, + lm_head, + logits_metadata, + aux_hidden_states=None, + hidden_states_before_norm=None, + ): + self.logits_calls.append(SimpleNamespace(hidden_states=hidden_states)) + return object() + + def prepare(self, forward_batch): + with torch.no_grad(): + return self.prepare_model_inputs( + input_ids=forward_batch.input_ids, + forward_batch=forward_batch, + input_embeds=None, + ) + + +class _TracedGroup: + """Records every collective (op, source, shape, group size) around a real coordinator.""" + + def __init__(self, inner): + self.inner = inner + self.trace = [] + + @property + def world_size(self): + return self.inner.world_size + + @property + def rank_in_group(self): + return self.inner.rank_in_group + + @property + def ranks(self): + return self.inner.ranks + + def all_gather_object(self, obj): + self.trace.append(["all_gather_object", type(obj).__name__, self.world_size]) + return self.inner.all_gather_object(obj) + + def broadcast_object(self, obj=None, src=0): + self.trace.append(["broadcast_object", src, self.world_size]) + return self.inner.broadcast_object(obj, src=src) + + def broadcast(self, tensor, src=0): + self.trace.append(["broadcast", src, list(tensor.shape), self.world_size]) + return self.inner.broadcast(tensor, src=src) + + +class _ForbiddenGroup: + world_size = 4 + rank_in_group = 0 + ranks = [0, 1, 2, 3] + + def __getattr__(self, name): + raise AssertionError(f"collective {name!r} reached on a fast path") + + +def _coordinator(group_ranks, rank): + return parallel_state.GroupCoordinator( + group_ranks=group_ranks, + local_rank=rank, + torch_distributed_backend="gloo", + use_pynccl=False, + use_pymscclpp=False, + use_custom_allreduce=False, + use_torch_symm_mem_all_reduce=False, + use_hpu_communicator=False, + use_xpu_communicator=False, + use_npu_communicator=False, + use_message_queue_broadcaster=False, + group_name="mm_owner_test", + gloo_timeout=GLOO_TIMEOUT, + ) + + +def _init_rank(rank: int, world_size: int, init_file: str) -> None: + torch.set_num_threads(1) + dist.init_process_group( + backend="gloo", + init_method=Path(init_file).as_uri(), + rank=rank, + world_size=world_size, + timeout=GLOO_TIMEOUT, + ) + parallel_state._MODEL_PARALLEL_GROUP_TIMEOUT = GLOO_TIMEOUT + + +def _run_ranks(world_size: int, target): + with TemporaryDirectory() as directory: + init_file = str(Path(directory) / "gloo-init") + torch.multiprocessing.spawn( + _rank_main, + args=(world_size, init_file, directory, target), + nprocs=world_size, + ) + return [ + json.loads((Path(directory) / f"rank{rank}.json").read_text()) + for rank in range(world_size) + ] + + +def _rank_main(rank, world_size, init_file, directory, target): + _init_rank(rank, world_size, init_file) + try: + result = target(rank, world_size) + (Path(directory) / f"rank{rank}.json").write_text(json.dumps(result)) + finally: + dist.destroy_process_group() + + +def _assert_traces_agree(results, ranks): + traces = [results[rank]["trace"] for rank in ranks] + assert all(trace == traces[0] for trace in traces), traces + + +def _expect_protocol_error(run): + try: + run() + except MmOwnerProtocolError as exc: + return str(exc) + raise AssertionError("the entrypoint did not raise MmOwnerProtocolError") + + +@torch.no_grad() +def _run_cp_extend(model, forward_batch, coordinator): + def gather(output, input_tensor): + dist.all_gather_into_tensor( + output, input_tensor, group=coordinator.device_group + ) + + runner = EagerRunner.__new__(EagerRunner) + runner.model_runner = SimpleNamespace(model=model) + with ( + patch("torch.cuda.current_stream", return_value=None), + patch( + "sglang.srt.layers.cp.interleave.attn_cp_all_gather_into_tensor", + side_effect=gather, + ), + patch( + "sglang.srt.layers.cp.interleave.is_allocation_symmetric", + return_value=False, + ), + patch( + "sglang.srt.layers.cp.interleave.use_symmetric_memory", + return_value=torch.no_grad(), + ), + ): + prepare_cp_forward(forward_batch) + runner._execute_extend_cp(forward_batch, {}) + + +# Global rank 0 sits outside the owner group, so a group-local source index that +# leaks out as a global rank is caught. +A, B, C, D, E = 400, 401, 402, 403, 404 # owner = hash % 4 -> group ranks 0,1,2,3,0 +ROWS = {A: 3, B: 4, C: 2, D: 3, E: 6} + + +def _owner_lifetime_program(rank, world_size): + group = _coordinator([[0], [1, 2, 3, 4]], rank) + if rank == 0: + return {"trace": [], "encoded": [], "outside": True} + traced = _TracedGroup(group) + local = group.rank_in_group + owner_tag = {key: 1 + key % 4 for key in ROWS} + mm_schedule.init_mm_embedding_cache(1 << 20) + embed = _ReducingEmbedding(group) + model = _VisionStub(embed, traced, tag=rank) + + # A valid only on its owner, B valid only on a non-owner, C valid + # everywhere, D stale on its owner. + if local == 0: + _seed_cache(A, ROWS[A], CACHE_TAG + rank) + if local == 2: + _seed_cache(B, ROWS[B], CACHE_TAG + rank) + _seed_cache(C, ROWS[C], CACHE_TAG) + if local == 3: + _seed_cache(D, ROWS[D] + 2, CACHE_TAG + rank) + tags = {A: CACHE_TAG + 1, B: owner_tag[B], C: CACHE_TAG, D: owner_tag[D]} + requests = [ + _request([10, 11, (A, ROWS[A]), 12, (B, ROWS[B]), 13], rid="r1"), + _request([20, (C, ROWS[C]), 21, (D, ROWS[D]), 22], rid="r2"), + ] + _, embeds = model.prepare(_batch(requests)) + assert torch.equal(embeds, _expected_embeds(embed, requests, tags)) + encoded = [list(model.encoded)] + traces = [list(traced.trace)] + + # Local 2 loses every entry and admits nothing new; B's owner evicts B. + if local == 2: + mm_schedule.init_mm_embedding_cache(0) + if local == 1: + mm_schedule.embedding_cache.free(B, None) + model.encoded.clear() + traced.trace.clear() + tags = {A: CACHE_TAG + 1, B: owner_tag[B], C: owner_tag[C], D: owner_tag[D]} + requests = [ + _request([10, 11, (A, ROWS[A]), (B, ROWS[B])], rid="r3"), + _request([20, (A, ROWS[A]), (C, ROWS[C]), (D, ROWS[D])], rid="r4"), + ] + _, embeds = model.prepare(_batch(requests)) + assert torch.equal(embeds, _expected_embeds(embed, requests, tags)) + encoded.append(list(model.encoded)) + traces.append(list(traced.trace)) + + # A cold span crossing the chunk boundary. + for prefix_len, extend_len in ((0, 5), (5, 4)): + model.encoded.clear() + traced.trace.clear() + request = _request( + [30, 31, (E, ROWS[E]), 32], + prefix_len=prefix_len, + extend_len=extend_len, + rid="r5", + ) + _, embeds = model.prepare(_batch([request])) + assert torch.equal( + embeds, _expected_embeds(embed, [request], {E: owner_tag[E]}) + ) + encoded.append(list(model.encoded)) + traces.append(list(traced.trace)) + return {"trace": traces, "encoded": encoded, "outside": False} + + +def _topology_program(rank, world_size): + tp8 = _coordinator([[0, 1, 2, 3, 4, 5, 6, 7]], rank) + replicas = _coordinator([[0, 1, 2, 3], [4, 5, 6, 7]], rank) + singles = _coordinator([[r] for r in range(8)], rank) + out = {} + + # TP8, DP1, CP off. + tp_group, attn_tp, attn_cp = ( + _TracedGroup(tp8), + _TracedGroup(tp8), + _TracedGroup(singles), + ) + with get_parallel().override( + tp_size=8, + attn_dp_size=1, + attn_cp_size=1, + tp_group=tp_group, + attn_tp_group=attn_tp, + attn_cp_group=attn_cp, + ): + selected = select_owner_group(get_parallel()) + assert selected is attn_tp + mm_schedule.init_mm_embedding_cache(1 << 20) + embed = _ReducingEmbedding(tp8) + model = _VisionStub(embed, selected, tag=rank) + x, y, z = 800, 805, 810 # owners 0, 5, 2 + requests = [ + _request([10, (x, 3), 11], rid="c1"), + _request([20, (y, 2), (z, 4), 21], rid="c2"), + ] + _, embeds = model.prepare(_batch(requests)) + assert torch.equal(embeds, _expected_embeds(embed, requests, {x: 0, y: 5, z: 2})) + out["cp1"] = {"encoded": list(model.encoded), "trace": list(attn_tp.trace)} + + # TP8, DP1, CP8 through the CP runner. + tp_group, attn_tp, attn_cp = ( + _TracedGroup(tp8), + _TracedGroup(singles), + _TracedGroup(tp8), + ) + init_cp_strategy(enable_prefill_cp=True, cp_size=8, cp_strategy="interleave") + try: + with get_parallel().override( + tp_size=8, + attn_dp_size=1, + attn_cp_size=8, + attn_cp_rank=rank, + tp_group=tp_group, + attn_tp_group=attn_tp, + attn_cp_group=attn_cp, + ): + selected = select_owner_group(get_parallel()) + assert selected is attn_cp + mm_schedule.init_mm_embedding_cache(1 << 20) + torch.manual_seed(0) + plain_embed = nn.Embedding(VOCAB, HIDDEN) + model = _VisionStub(plain_embed, selected, tag=rank) + w = 803 # owner 3; rows land on ranks 6,7,0,1,2,3 so ranks 4,5 hold no image row + requests = [ + _request([10, 11, 12, 13, 14], rid="p1"), + _request([20, (w, 6), 21, 22], rid="p2"), + ] + forward_batch = _batch(requests) + full = _expected_embeds(plain_embed, requests, {w: 3}) + _run_cp_extend(model, forward_batch, tp8) + physical = forward_batch.attn_cp_metadata.per_rank_actual_token[rank] + (body,) = model.model.calls + shard = full[rank::8] + assert torch.equal(body.input_embeds[: shard.shape[0]], shard) + assert body.input_embeds.shape[0] == physical + (logits,) = model.logits_calls + assert torch.equal(logits.hidden_states, full) + assert torch.equal(forward_batch.mm_input_embeds, full) + out["cp8"] = { + "encoded": list(model.encoded), + "trace": list(attn_cp.trace), + "image_rows": int((body.input_ids == IMAGE_TOKEN_ID).sum()), + } + finally: + init_cp_strategy(enable_prefill_cp=False, cp_size=1, cp_strategy="interleave") + + # TP8, attention-DP2, CP off: one replica is text-only first. + tp_group, attn_tp, attn_cp = ( + _TracedGroup(tp8), + _TracedGroup(replicas), + _TracedGroup(singles), + ) + with get_parallel().override( + tp_size=8, + attn_dp_size=2, + attn_cp_size=1, + tp_group=tp_group, + attn_tp_group=attn_tp, + attn_cp_group=attn_cp, + ): + selected = select_owner_group(get_parallel()) + assert selected is attn_tp + mm_schedule.init_mm_embedding_cache(1 << 20) + embed = _ReducingEmbedding(replicas) + model = _VisionStub(embed, selected, tag=rank) + p, q = 900, 901 # owners: local 0 (rank 0) and local 1 (rank 5) + replica = rank // 4 + if replica == 0: + requests = [_request([10, (p, 3), 11], rid="d0")] + else: + requests = [_request([10, 11, 12], rid="d1")] + _, embeds = model.prepare(_batch(requests)) + if replica == 0: + assert torch.equal(embeds, _expected_embeds(embed, requests, {p: 0})) + else: + assert embeds is None + first = {"encoded": list(model.encoded), "trace": list(attn_tp.trace)} + model.encoded.clear() + attn_tp.trace.clear() + if replica == 0: + requests = [_request([10, (p, 3), 11], rid="d2")] + tags = {p: 0} + else: + requests = [_request([30, (q, 2), 31], rid="d3")] + tags = {q: 5} + _, embeds = model.prepare(_batch(requests)) + assert torch.equal(embeds, _expected_embeds(embed, requests, tags)) + out["dp2"] = [first, {"encoded": list(model.encoded), "trace": list(attn_tp.trace)}] + return out + + +F, G, H, I = 300, 301, 302, 303 # owners 0, 1, 2, 3 in a four-rank group +J, K = 310, 311 # owner 2 and owner 3, eight-row spans for the grid cases + + +def _failure_program(rank, world_size): + group = _coordinator([[0, 1, 2, 3]], rank) + out = {} + + def fresh(embed_cls=_ReducingEmbedding): + mm_schedule.init_mm_embedding_cache(1 << 20) + traced = _TracedGroup(group) + embed = embed_cls(group) + return traced, embed, _VisionStub(embed, traced, tag=rank) + + def inject(module, name, message, on_rank): + if rank != on_rank: + return nullcontext() + return patch.object(module, name, side_effect=RuntimeError(message)) + + def record(name, traced, embed, model, error, **extra): + out[name] = { + "error": error, + "trace": list(traced.trace), + "embed_calls": embed.calls, + "encoded": list(model.encoded), + **extra, + } + + def remap_oom(target): + original = torch.Tensor.masked_fill + + def masked_fill(self, mask, value): + if self is target: + raise torch.OutOfMemoryError("injected remap OOM") + return original(self, mask, value) + + return patch.object(torch.Tensor, "masked_fill", masked_fill) + + def sync_failure(on_call): + seen = [] + + def synchronize(device): + seen.append(device) + if len(seen) == on_call: + raise RuntimeError("injected final sync") + + return patch.object(mm_owner_embedding, "_synchronize", synchronize) + + traced, embed, model = fresh() + batch = _batch([_request([10, (G, 3), 11], rid="fail-prepare")]) + with ( + patch.object( + torch, + "as_tensor", + side_effect=torch.OutOfMemoryError("injected placeholder OOM"), + ) + if rank == 1 + else nullcontext() + ): + error = _expect_protocol_error(lambda: model.prepare(batch)) + record("prepare", traced, embed, model, error) + + traced, embed, model = fresh() + if rank == 1: + model.fail_hashes = {G} + error = _expect_protocol_error( + lambda: model.prepare(_batch([_request([10, (G, 3), 11], rid="fail-encode")])) + ) + record("encode", traced, embed, model, error) + + traced, embed, model = fresh() + with inject( + mm_owner_embedding, "_new_span_buffer", "injected allocation", on_rank=3 + ): + error = _expect_protocol_error( + lambda: model.prepare( + _batch([_request([10, (F, 3), 11], rid="fail-alloc")]) + ) + ) + record("alloc", traced, embed, model, error) + + traced, embed, model = fresh() + rows = 4 if rank == 2 else 3 + error = _expect_protocol_error( + lambda: model.prepare( + _batch([_request([10, (F, rows), 11], rid="fail-manifest")]) + ) + ) + record("manifest", traced, embed, model, error) + + traced, embed, model = fresh() + grid = (4, 3) if rank == 2 else (3, 4) + error = _expect_protocol_error( + lambda: model.prepare( + _batch([_request([10, (J, 8, grid), 11], rid="fail-grid")]) + ) + ) + record("grid", traced, embed, model, error) + + traced, embed, model = fresh() + error = _expect_protocol_error( + lambda: model.prepare( + _batch( + [ + _request([10, (K, 8, (3, 4)), 11], rid="fail-dup-a"), + _request([20, (K, 8, (4, 3)), 21], rid="fail-dup-b"), + ] + ) + ) + ) + record("duplicate", traced, embed, model, error) + + traced, embed, model = fresh() + error = _expect_protocol_error( + lambda: model.prepare( + _batch([_request([10, (F, 5, (2, 2)), 11], rid="fail-span-len")]) + ) + ) + record("span_len", traced, embed, model, error) + + traced, embed, model = fresh() + with inject( + mm_schedule, "_assemble_per_image_chunk", "injected assembly", on_rank=1 + ): + error = _expect_protocol_error( + lambda: model.prepare( + _batch([_request([10, (H, 2), 11], rid="fail-assemble")]) + ) + ) + record("assemble", traced, embed, model, error) + + traced, embed, model = fresh() + with inject(mm_utils, "_scatter_mm_embedding", "injected merge", on_rank=3): + error = _expect_protocol_error( + lambda: model.prepare( + _batch([_request([10, (I, 2), 11], rid="fail-merge")]) + ) + ) + record("merge", traced, embed, model, error) + + traced, embed, model = fresh() + batch = _batch([_request([10, (F, 3), 11], rid="fail-remap")]) + with remap_oom(batch.input_ids) if rank == 1 else nullcontext(): + error = _expect_protocol_error(lambda: model.prepare(batch)) + record("remap", traced, embed, model, error) + + traced, embed, model = fresh() + with sync_failure(on_call=3) if rank == 1 else nullcontext(): + error = _expect_protocol_error( + lambda: model.prepare( + _batch([_request([10, (F, 3), 11], rid="fail-final-sync")]) + ) + ) + record("final_sync", traced, embed, model, error) + + traced, embed, model = fresh() + with ( + patch.object( + mm_owner_embedding, + "_make_plan", + side_effect=MemoryError("injected leader plan allocation"), + ) + if rank == 0 + else nullcontext() + ): + error = _expect_protocol_error( + lambda: model.prepare(_batch([_request([10, (G, 3), 11], rid="fail-plan")])) + ) + record("leader_plan", traced, embed, model, error) + + init_cp_strategy(enable_prefill_cp=True, cp_size=4, cp_strategy="interleave") + try: + traced, embed, model = fresh(lambda group: _ReducingEmbedding(None)) + batch = _batch([_request([10, (F, 3), 11], rid="fail-remap-cp")]) + with ( + get_parallel().override( + attn_cp_size=4, attn_cp_rank=rank, attn_cp_group=traced + ), + remap_oom(batch.input_ids) if rank == 1 else nullcontext(), + ): + error = _expect_protocol_error(lambda: _run_cp_extend(model, batch, group)) + record( + "remap_cp", traced, embed, model, error, body_calls=len(model.model.calls) + ) + finally: + init_cp_strategy(enable_prefill_cp=False, cp_size=1, cp_strategy="interleave") + return out + + +MANIFEST = ["all_gather_object", "RankManifest", 4] +PLAN = ["broadcast_object", 0, 4] +STATUS = ["all_gather_object", "RankStatus", 4] + + +def _bcast(src, rows, size=4): + return ["broadcast", src, [rows, HIDDEN], size] + + +def test_owner_actions_and_cache_lifetime_across_asymmetric_ranks(): + """Owner hits broadcast, non-owner hits still receive, stale entries miss, + all-hit moves no payload, per-forward references outlive the cache.""" + results = _run_ranks(5, _owner_lifetime_program) + members = range(1, 5) + _assert_traces_agree(results, members) + encoded = [results[rank]["encoded"] for rank in members] + trace = results[1]["trace"] + + # Forward 1: A cache-broadcast, B and D owner-encoded, C a local hit. + assert [e[0] for e in encoded] == [[], [B], [], [D]] + assert trace[0] == [ + MANIFEST, + PLAN, + STATUS, + _bcast(0, ROWS[A]), + _bcast(1, ROWS[B]), + _bcast(3, ROWS[D]), + STATUS, + STATUS, + ] + + # Forward 2: evicted owners of B and C re-encode once; A recurs but moves once. + assert [e[1] for e in encoded] == [[], [B], [C], []] + assert trace[1] == [ + MANIFEST, + PLAN, + STATUS, + _bcast(0, ROWS[A]), + _bcast(1, ROWS[B]), + _bcast(2, ROWS[C]), + _bcast(3, ROWS[D]), + STATUS, + STATUS, + ] + + # Forwards 3 and 4: the chunk-crossing span is encoded fully once. + assert [e[2] for e in encoded] == [[E], [], [], []] + assert [e[3] for e in encoded] == [[], [], [], []] + assert trace[2] == [MANIFEST, PLAN, STATUS, _bcast(0, ROWS[E]), STATUS, STATUS] + assert trace[3] == [MANIFEST, PLAN, STATUS, _bcast(0, ROWS[E]), STATUS, STATUS] + assert results[0]["outside"] + + +def test_tp8_cp1_and_cp8_dedupe_and_attention_dp_replicas_stay_isolated(): + """One encode per key on TP8 with CP off and CP8; attention-DP replicas never share a collective.""" + results = _run_ranks(8, _topology_program) + + cp1 = [r["cp1"] for r in results] + assert [c["encoded"] for c in cp1] == [[800], [], [810], [], [], [805], [], []] + assert all(c["trace"] == cp1[0]["trace"] for c in cp1) + assert [op for op in cp1[0]["trace"] if op[0] == "broadcast"] == [ + _bcast(0, 3, 8), + _bcast(5, 2, 8), + _bcast(2, 4, 8), + ] + + cp8 = [r["cp8"] for r in results] + assert [c["encoded"] for c in cp8] == [[], [], [], [803], [], [], [], []] + assert all(c["trace"] == cp8[0]["trace"] for c in cp8) + assert [op for op in cp8[0]["trace"] if op[0] == "broadcast"] == [_bcast(3, 6, 8)] + assert [c["image_rows"] for c in cp8] == [1, 1, 1, 1, 0, 0, 1, 1] + + dp2 = [r["dp2"] for r in results] + first = [d[0] for d in dp2] + assert [f["encoded"] for f in first] == [[900], [], [], [], [], [], [], []] + assert all(f["trace"] == [] for f in first[4:]) + assert all(f["trace"] == first[0]["trace"] for f in first[:4]) + assert all(op[-1] == 4 for op in first[0]["trace"]) + second = [d[1] for d in dp2] + assert [s["encoded"] for s in second] == [[], [], [], [], [], [901], [], []] + assert all(s["trace"] == second[0]["trace"] for s in second[:4]) + assert all(s["trace"] == second[4]["trace"] for s in second[4:]) + assert [op for op in second[0]["trace"] if op[0] == "broadcast"] == [] + assert [op for op in second[4]["trace"] if op[0] == "broadcast"] == [ + _bcast(1, 2, 4) + ] + + +def test_failures_agree_before_payload_text_embedding_or_body(): + """Every rank raises the same error and none enters a later collective.""" + results = _run_ranks(4, _failure_program) + cases = ( + "prepare", + "encode", + "alloc", + "manifest", + "grid", + "duplicate", + "span_len", + "assemble", + "merge", + "remap", + "final_sync", + "leader_plan", + "remap_cp", + ) + for case in cases: + errors = {r[case]["error"] for r in results} + assert len(errors) == 1, (case, errors) + assert all(r[case]["trace"] == results[0][case]["trace"] for r in results), case + + def calls(case, field="embed_calls"): + return [r[case][field] for r in results] + + prepare = results[0]["prepare"] + assert "during prepare on group rank 1" in prepare["error"] + assert "injected placeholder OOM" in prepare["error"] + assert prepare["trace"] == [MANIFEST, PLAN] + assert calls("prepare") == [0, 0, 0, 0] + + encode = results[0]["encode"] + assert "during encode on group rank 1" in encode["error"] + assert "fail-encode" in encode["error"] and "301" in encode["error"] + assert encode["trace"] == [MANIFEST, PLAN, STATUS] + assert calls("encode") == [0, 0, 0, 0] + + alloc = results[0]["alloc"] + assert "during encode on group rank 3" in alloc["error"] + assert "hash 300 shape (3, 8)" in alloc["error"] + assert "injected allocation" in alloc["error"] + assert alloc["trace"] == [MANIFEST, PLAN, STATUS] + assert calls("alloc", "encoded") == [[F], [], [], []] + assert calls("alloc") == [0, 0, 0, 0] + + manifest = results[0]["manifest"] + assert "manifest mismatch between group ranks 0 and 2" in manifest["error"] + assert manifest["trace"] == [MANIFEST, PLAN] + assert calls("manifest") == [0, 0, 0, 0] + + grid = results[0]["grid"] + assert "manifest mismatch between group ranks 0 and 2" in grid["error"] + assert "(3, 4" in grid["error"] and "(4, 3" in grid["error"] + assert grid["trace"] == [MANIFEST, PLAN] + + duplicate = results[0]["duplicate"] + assert "during manifest on group rank 0" in duplicate["error"] + assert ( + f"image hash {K} (8 tokens) occurs with different geometry" + in duplicate["error"] + ) + assert duplicate["trace"] == [MANIFEST, PLAN] + + span_len = results[0]["span_len"] + assert "yields 4 span tokens, placeholder has 5" in span_len["error"] + assert span_len["trace"] == [MANIFEST, PLAN] + + assemble = results[0]["assemble"] + assert "during features on group rank 1" in assemble["error"] + assert "injected assembly" in assemble["error"] + assert assemble["trace"] == [MANIFEST, PLAN, STATUS, _bcast(2, 2), STATUS] + assert calls("assemble") == [0, 0, 0, 0] + + merge = results[0]["merge"] + assert "during finalize on group rank 3" in merge["error"] + assert merge["trace"] == [MANIFEST, PLAN, STATUS, _bcast(3, 2), STATUS, STATUS] + assert calls("merge") == [1, 1, 1, 1] + + remap = results[0]["remap"] + assert "during finalize on group rank 1" in remap["error"] + assert "injected remap OOM" in remap["error"] + assert remap["trace"] == [MANIFEST, PLAN, STATUS, _bcast(0, 3), STATUS, STATUS] + assert calls("remap") == [1, 1, 1, 1] + + final_sync = results[0]["final_sync"] + assert "during finalize on group rank 1" in final_sync["error"] + assert "injected final sync" in final_sync["error"] + assert final_sync["trace"] == [MANIFEST, PLAN, STATUS, _bcast(0, 3), STATUS, STATUS] + assert calls("final_sync") == [1, 1, 1, 1] + + leader_plan = results[0]["leader_plan"] + assert "during plan on group rank 0" in leader_plan["error"] + assert "injected leader plan allocation" in leader_plan["error"] + assert leader_plan["trace"] == [MANIFEST, PLAN] + assert calls("leader_plan") == [0, 0, 0, 0] + assert calls("leader_plan", "encoded") == [[], [], [], []] + + remap_cp = results[0]["remap_cp"] + assert "during finalize on group rank 1" in remap_cp["error"] + assert remap_cp["trace"] == [MANIFEST, PLAN, STATUS, _bcast(0, 3), STATUS, STATUS] + assert calls("remap_cp", "body_calls") == [0, 0, 0, 0] + + +def test_text_decode_prefilled_future_and_precomputed_paths_pay_no_collective(): + """Chunks without owner-encoded image rows never touch the group; an overlapping chunk does.""" + mm_schedule.init_mm_embedding_cache(1 << 20) + torch.manual_seed(0) + embed = nn.Embedding(VOCAB, HIDDEN) + owner_model = _VisionStub(embed, _ForbiddenGroup(), tag=0) + legacy_model = _VisionStub(embed, None, tag=0) + + text_only = _batch([_request([10, 11], rid="t0"), _request([20, 21, 22], rid="t1")]) + assert owner_model.prepare(text_only)[1] is None + + with_image = [_request([10, (A, 3), 11, 12], rid="v0")] + for mode in (ForwardMode.DECODE, ForwardMode.TARGET_VERIFY): + assert owner_model.prepare(_batch(with_image, mode))[1] is None + + prefilled = [_request([10, (A, 3), 11, 12], prefix_len=5, rid="pf")] + ids, embeds = owner_model.prepare(_batch(prefilled)) + _, legacy_embeds = legacy_model.prepare(_batch(prefilled)) + assert torch.equal(embeds, legacy_embeds) + assert torch.equal(ids, torch.tensor([12])) + + future = [_request([10, 11, 12, 13, (A, 3)], extend_len=2, rid="fut")] + ids, embeds = owner_model.prepare(_batch(future)) + _, legacy_embeds = legacy_model.prepare(_batch(future)) + assert torch.equal(embeds, legacy_embeds) + assert torch.equal(ids, torch.tensor([10, 11])) + + precomputed = _request([10, 11, 12, 13], rid="pc") + item = MultimodalDataItem( + modality=Modality.IMAGE, + precomputed_embeddings=torch.full((3, HIDDEN), 5.0), + offsets=[(1, 3)], + ) + item.set_hash(A) + precomputed.items = [item] + precomputed.ids[1:4] = [item.pad_value] * 3 + _, embeds = owner_model.prepare(_batch([precomputed])) + assert torch.equal(embeds[1:4], torch.full((3, HIDDEN), 5.0)) + + overlapping = [_request([10, (A, 3), 11, 12], prefix_len=2, rid="ov")] + with pytest.raises(AssertionError, match="collective"): + owner_model.prepare(_batch(overlapping)) + + +def _parallel(tp_size, attn_dp_size, attn_cp_size, attn_tp_size, attn_cp_group_size): + return SimpleNamespace( + tp_size=tp_size, + attn_dp_size=attn_dp_size, + attn_cp_size=attn_cp_size, + attn_tp_group=SimpleNamespace(world_size=attn_tp_size, name="attn_tp"), + attn_cp_group=SimpleNamespace(world_size=attn_cp_group_size, name="attn_cp"), + ) + + +def test_select_owner_group_follows_the_replication_domain(): + """R = TP / attention-DP selects the attention-TP group, TP-aliased CP its handle, else None.""" + assert select_owner_group(_parallel(8, 1, 1, 8, 1)).name == "attn_tp" + assert select_owner_group(_parallel(8, 1, 8, 1, 8)).name == "attn_cp" + assert select_owner_group(_parallel(8, 2, 1, 4, 1)).name == "attn_tp" + assert select_owner_group(_parallel(1, 1, 1, 1, 1)) is None + assert select_owner_group(_parallel(8, 2, 4, 1, 4)) is None + assert select_owner_group(_parallel(8, 1, 4, 2, 4)) is None + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__])) diff --git a/test/registered/unit/models/test_deepseek_v41_vision_cp_inputs.py b/test/registered/unit/models/test_deepseek_v41_vision_cp_inputs.py index 01b09f578..748d3cc22 100644 --- a/test/registered/unit/models/test_deepseek_v41_vision_cp_inputs.py +++ b/test/registered/unit/models/test_deepseek_v41_vision_cp_inputs.py @@ -77,6 +77,7 @@ class _VisionStub(DeepseekV4ForCausalLM): self.config = SimpleNamespace(image_token_id=IMAGE_TOKEN_ID) self.vision = object() self.tp_size = 1 + self.mm_owner_group = None self.model = _RecordingBody(embed) self.pp_group = SimpleNamespace(is_last_rank=True) self.lm_head = object()