""" Multi-modality utils """ import logging from abc import abstractmethod from typing import Callable, List, Optional, Tuple import torch from torch import nn from sglang.srt.managers.schedule_batch import ( MultimodalDataItem, MultimodalInputs, global_server_args_dict, ) from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.utils import print_warning_once logger = logging.getLogger(__name__) class MultiModalityDataPaddingPattern: """ Data tokens (like image tokens) often need special handling during padding to maintain model compatibility. This class provides the interface for implementing different padding strategies for data tokens """ @abstractmethod def pad_input_tokens( self, input_ids: List[int], mm_inputs: MultimodalInputs ) -> List[int]: """ Pad the input ids sequence containing data tokens, and replace them with pad_values """ pass class MultiModalityDataPaddingPatternTokenPairs(MultiModalityDataPaddingPattern): """In this pattern, data tokens should be enclosed by special token pairs (e.g. ..., data_token_pairs) This strategy should be applied when data content is marked by start/end token pairs in the input sequence. """ def __init__(self, data_token_pairs: Optional[List[Tuple[int, int]]]) -> None: self.data_token_id_pairs = data_token_pairs def pad_input_tokens( self, input_ids: List[int], mm_inputs: MultimodalInputs ) -> List[int]: """ This function will replace the data-tokens inbetween with pad_values accordingly """ pad_values = [item.pad_value for item in mm_inputs.mm_items] data_token_pairs = self.data_token_id_pairs mm_inputs.data_offsets = [] if data_token_pairs is None: data_token_pairs = [mm_inputs.im_start_id, mm_inputs.im_end_id] if data_token_pairs is None: print_warning_once( "No data_token_pairs provided, RadixAttention might be influenced." ) return input_ids start_token_ids = [s for s, _e in data_token_pairs] end_tokens_ids = [e for _s, e in data_token_pairs] padded_ids = [] last_idx = 0 data_idx = -1 start_indices = [i for i, x in enumerate(input_ids) if x in start_token_ids] end_indices = [i for i, x in enumerate(input_ids) if x in end_tokens_ids] if len(start_indices) != len(end_indices): return input_ids for start_idx, end_idx in zip(start_indices, end_indices): padded_ids.extend(input_ids[last_idx : start_idx + 1]) if input_ids[start_idx] in start_token_ids: data_idx += 1 mm_inputs.data_offsets += [start_idx] if data_idx >= len(pad_values): data_idx = len(pad_values) - 1 num_tokens = end_idx - start_idx - 1 pad_value = pad_values[data_idx] padded_ids.extend([pad_value] * num_tokens) last_idx = end_idx padded_ids.extend(input_ids[last_idx:]) assert len(input_ids) == len(padded_ids), "Length validation fails" return padded_ids class MultiModalityDataPaddingPatternImageTokens(MultiModalityDataPaddingPattern): """In this pattern, data tokens should be represented as repetitions of a single token e.g. ...., or