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()