Spaces:
Running on Zero
Running on Zero
| # Copyright 2025 Bytedance Ltd. and/or its affiliates | |
| # | |
| # Licensed under the Apache License, Version 2.0 (the "License"); | |
| # you may not use this file except in compliance with the License. | |
| # You may obtain a copy of the License at | |
| # | |
| # http://www.apache.org/licenses/LICENSE-2.0 | |
| # | |
| # Unless required by applicable law or agreed to in writing, software | |
| # distributed under the License is distributed on an "AS IS" BASIS, | |
| # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. | |
| # See the License for the specific language governing permissions and | |
| # limitations under the License. | |
| import json | |
| from collections import defaultdict | |
| from typing import TYPE_CHECKING, Any, Callable, Dict, List | |
| import torch | |
| from ...utils.constants import TYPE2INDEX | |
| from ...utils.import_utils import is_video_audio_available | |
| from .image_utils import fetch_images | |
| from .preprocess import conv_preprocess | |
| if is_video_audio_available(): | |
| from .audio_utils import fetch_audios | |
| from .video_utils import fetch_videos | |
| else: | |
| def fetch_videos(*args, **kwargs): | |
| return [], [] | |
| def fetch_audios(*args, **kwargs): | |
| return [] | |
| if TYPE_CHECKING: | |
| from ...models.seed_omni import SeedOmniProcessor | |
| from .multimodal_chat_template import MultimodalChatTemplate | |
| def mask_before_position_id_func(input_ids: torch.Tensor): | |
| """Mask special multimodal tokens in input_ids to input_mm_token for position_id. | |
| Only supports special image tokens now. (input_image_id=-200, output_image_id=-201->-200) | |
| Similar to veomni.module.seed_omni.modeling_seed_omni.mask_before_text_encoder | |
| Args: | |
| input_ids (torch.Tensor) | |
| Returns: | |
| input_ids (torch.Tensor) | |
| """ | |
| for modality in ["image", "video", "audio"]: | |
| output_mask = input_ids == TYPE2INDEX["output"][modality] | |
| input_mask = input_ids == TYPE2INDEX["input"][modality] | |
| input_ids = torch.where(output_mask | input_mask, TYPE2INDEX["input"][modality], input_ids) | |
| return input_ids | |
| def mask_input_ids(modality_info: Dict, input_ids: torch.Tensor): | |
| """Mask special multimodal tokens in input_ids to 0 for text_encoder.word_embedding. | |
| And return masks including: image_input_mask, image_output_mask, etc | |
| For example: | |
| input_ids: torch.tensor([-200, -200, 2, -200, -200, 4, 5, 6, -201, -201]) | |
| Returns: | |
| input_ids: torch.tensor([0, 0, 2, 0, 0, 4, 5, 6, 0, 0 ]) | |
| image_input_mask: torch.tensor([1, 1, 0, 1, 1, 0, 0, 0, 0, 0 ]) | |
| image_output_mask: torch.tensor([0, 0, 0, 0, 0, 0, 0, 0, 1, 1 ]) | |
| Args: | |
| input_ids (torch.Tensor) | |
| Returns: | |
| input_ids (torch.Tensor) | |
| mask_dict (Dict) : {modal}_[input/output]_mask. | |
| """ | |
| mask_dict = {} | |
| for data_type in modality_info.keys(): | |
| for modal in modality_info[data_type]: | |
| mask = input_ids == TYPE2INDEX[data_type][modal] | |
| mask_dict[f"{modal}_{data_type}_mask"] = mask | |
| input_ids = torch.where(mask, 0, input_ids) | |
| return input_ids, mask_dict | |
| def process_mm_data( | |
| conversations, images: List[Any], videos: List[Any], video_audios: List[Any], audio_audios: List[Any] | |
| ): | |
| """ | |
| Processes multi-modal conversation data and aligns images, videos, and audio | |
| with a corresponding output mask indicating whether the data was produced by the assistant. | |
| Parameters: | |
| ---------- | |
| conversations : List[List] | |
| images : List[Any, List of image data in order. | |
| videos : List[Any], List of video data in order. | |
| video_audios : List[Any], List of audio tracks corresponding to the videos. | |
| audio_audios : List[Any], List of standalone audio samples. | |
| Returns: | |
| ------- | |
| conv_images : List[Any], List of images in the order they appeared in conversations. | |
| conv_videos : List[Any], List of videos in the order they appeared in conversations. | |
| conv_audios : List[Any], List of all audio data, including both video audio and standalone audio. | |
| mask : Dict[str, torch.BoolTensor] | |
| A dictionary with modality names as keys ("image", "video", "audio"), and boolean tensors | |
| indicating whether each sample was produced by the assistant (True) or the user (False). | |
| Example: | |
| -------- | |
| Input: | |
| conversations = [ | |
| ["user", ["video"], ["audio"], ["video"], ["text"]], | |
| ["assistant", ["audio"]] | |
| ] | |
| videos = ["video1", "video2"] | |
| video_audios = ["v_audio1", "v_audio2"] | |
| audio_audios = ["audio1", "audio2"] | |
| Output: | |
| conv_videos = ["video1", "video2"] | |
| conv_audios = ["v_audio1", "audio1", "v_audio2", "audio2"] | |
| mask["video"] = tensor([False, False]) # user videos | |
| mask["audio"] = tensor([False, False, False, True]) # user+assistant audios | |
| """ | |
| images, videos, video_audios, audio_audios = iter(images), iter(videos), iter(video_audios), iter(audio_audios) | |
| conv_images, conv_videos, conv_audios = [], [], [] | |
| mask = defaultdict(list) | |
| for conversation in conversations: | |
| role = conversation[0] | |
| is_output = role == "assistant" | |
| for message in conversation[1:]: | |
| data_type = message[0] | |
| if data_type == "text": | |
| continue | |
| elif data_type == "image": | |
| conv_images.append(next(images)) | |
| mask["image"].append(is_output) | |
| elif data_type == "video": | |
| conv_videos.append(next(videos)) | |
| conv_audios.append(next(video_audios)) | |
| mask["video"].append(is_output) | |
| mask["audio"].append(is_output) | |
| elif data_type == "audio": | |
| conv_audios.append(next(audio_audios)) | |
| mask["audio"].append(is_output) | |
| else: | |
| raise ValueError(f"Unknown data type: {data_type}") | |
| mask = {key: torch.tensor(value).type(torch.bool) for key, value in mask.items()} | |
| return conv_images, conv_videos, conv_audios, mask | |
| def get_multimodal_configs(modality_input: Dict, multimodal_output_mask: Dict): | |
| multimodal_configs, config_repr = {}, {} | |
| for key in modality_input.keys(): | |
| config_key = key.split("_", 2)[-1] | |
| if config_key != "features": | |
| config_repr[config_key] = modality_input[key] | |
| for config_key, repr in config_repr.items(): | |
| multimodal_configs[config_key] = {} | |
| for modal, mm_mask in multimodal_output_mask.items(): | |
| if ( | |
| f"{modal}_input_{config_key}" not in modality_input | |
| and f"{modal}_output_{config_key}" not in modality_input | |
| ): | |
| continue | |
| input_config = modality_input.get(f"{modal}_input_{config_key}", torch.empty_like(repr)) | |
| output_config = modality_input.get(f"{modal}_output_{config_key}", torch.empty_like(repr)) | |
| config = torch.zeros_like(repr) | |
| config = config.repeat_interleave(mm_mask.shape[0], dim=0) | |
| config[mm_mask] = output_config | |
| config[~mm_mask] = input_config | |
| multimodal_configs[config_key][modal] = config | |
| return multimodal_configs | |
| def keep_input_only(multimodal_config: Dict, multimodal_output_mask: Dict): | |
| """Only keep the input data in multimodal_config. Used when use_special_rope=False. | |
| When use_special_rope=False, only do special_rope on input_multimodal_data. | |
| For example: 2d_rope on input_image_token, but 1d_rope on output_image_token. | |
| """ | |
| for config in multimodal_config.keys(): | |
| for modal in multimodal_config[config].keys(): | |
| multimodal_config[config][modal] = multimodal_config[config][modal][~multimodal_output_mask[modal]] | |
| def encode_multimodal_sample( | |
| sample: Dict[str, Any], | |
| processor: "SeedOmniProcessor", | |
| chat_template: "MultimodalChatTemplate", | |
| position_id_func: "Callable", | |
| modality_info: Dict, | |
| use_special_rope=False, # 2d rope position id for image generation | |
| **kwargs, | |
| ) -> Dict[str, List[int]]: | |
| model_inputs = {} | |
| source = sample.pop("source_name") if "source_name" in sample else kwargs["source_name"] | |
| modality = set(modality_info["input"] + modality_info["output"]) | |
| conversations = sample["conversations"] if ("conversations" in sample and sample["conversations"]) else sample | |
| if isinstance(conversations, bytes): | |
| conversations = json.loads(conversations.decode("utf-8")) | |
| conversations = conv_preprocess(source, conversations, **kwargs) | |
| processor_input = {} | |
| if "image" in modality: | |
| images = fetch_images(sample.get("images", []), **kwargs) | |
| else: | |
| images = [] | |
| if "video" in modality: | |
| videos, video_audios = fetch_videos(sample.get("videos", []), **kwargs) | |
| if "audio" not in modality: | |
| video_audios = [None] * len(videos) | |
| else: | |
| videos, video_audios = [], [] | |
| if "audio" in modality: | |
| audio_audios = fetch_audios(sample.get("audios", []), **kwargs) | |
| else: | |
| audio_audios = [] | |
| images, videos, audios, multimodal_output_mask = process_mm_data( | |
| conversations, images, videos, video_audios, audio_audios | |
| ) | |
| if images: | |
| processor_input.update( | |
| { | |
| "input_images": [img for img, mask in zip(images, multimodal_output_mask["image"]) if not mask], | |
| "output_images": [img for img, mask in zip(images, multimodal_output_mask["image"]) if mask], | |
| } | |
| ) | |
| if videos: | |
| processor_input.update( | |
| { | |
| "input_videos": [vid for vid, mask in zip(videos, multimodal_output_mask["video"]) if not mask], | |
| "output_videos": [img for img, mask in zip(videos, multimodal_output_mask["video"]) if mask], | |
| } | |
| ) | |
| if audios and "audio" in modality: | |
| processor_input.update( | |
| { | |
| "input_audios": [aud for aud, mask in zip(audios, multimodal_output_mask["audio"]) if not mask], | |
| "output_audios": [aud for aud, mask in zip(audios, multimodal_output_mask["audio"]) if mask], | |
| } | |
| ) | |
| modality_input = processor(return_tensors="pt", **processor_input) | |
| multimodal_config = get_multimodal_configs(modality_input, multimodal_output_mask) | |
| text_inputs = chat_template.encode_messages(conversations, **multimodal_config) | |
| model_inputs.update(modality_input) | |
| model_inputs.update(text_inputs) | |
| # position_ids (dim, len) | |
| if position_id_func is None: # default position_ids | |
| position_ids = torch.arange(0, len(text_inputs["input_ids"])).unsqueeze(0) | |
| else: # customized position_ids | |
| input_ids = text_inputs["input_ids"].clone() | |
| attention_mask = text_inputs["attention_mask"].clone() | |
| if use_special_rope: | |
| input_ids = mask_before_position_id_func(input_ids) | |
| else: | |
| keep_input_only(multimodal_config, multimodal_output_mask) | |
| position_ids = position_id_func( | |
| input_ids=input_ids.unsqueeze(0), attention_mask=attention_mask.unsqueeze(0), **multimodal_config | |
| )["position_ids"] | |
| model_inputs["position_ids"] = position_ids | |
| input_ids, mask_dict = mask_input_ids(modality_info, model_inputs["input_ids"]) | |
| model_inputs["input_ids"] = input_ids | |
| model_inputs.update(mask_dict) | |
| return [model_inputs] | |
| def encode_multimodal_sample_inference( | |
| sample: Dict[str, Any], | |
| processor: "SeedOmniProcessor", | |
| chat_template: "MultimodalChatTemplate", | |
| position_id_func: "Callable", | |
| modality_info: Dict, | |
| force_image_gen: bool, | |
| **kwargs, | |
| ): | |
| model_inputs = {} | |
| modality = set(modality_info["input"] + modality_info["output"]) | |
| conversations = sample["conversations"] | |
| processor_input = {} | |
| if "image" in modality: | |
| images = fetch_images(sample.get("images", []), **kwargs) | |
| else: | |
| images = [] | |
| if "video" in modality: | |
| videos, video_audios = fetch_videos(sample.get("videos", []), **kwargs) | |
| if "audio" not in modality: | |
| video_audios = [None] * len(videos) | |
| else: | |
| videos, video_audios = [], [] | |
| if "audio" in modality: | |
| audio_audios = fetch_audios(sample.get("audios", []), **kwargs) | |
| else: | |
| audio_audios = [] | |
| images, videos, audios, multimodal_output_mask = process_mm_data( | |
| conversations, images, videos, video_audios, audio_audios | |
| ) | |
| if images: | |
| processor_input["input_images"] = images | |
| if videos: | |
| processor_input["input_videos"] = videos | |
| if audios and "audio" in modality: | |
| processor_input["input_audios"] = audios | |
| modality_input = processor(return_tensors="pt", **processor_input) | |
| multimodal_config = get_multimodal_configs(modality_input, multimodal_output_mask) | |
| text_inputs = chat_template.encode_messages(conversations, **multimodal_config) | |
| if force_image_gen: | |
| text_inputs["input_ids"] = torch.cat( | |
| [text_inputs["input_ids"], torch.tensor([chat_template.image_start_id])], | |
| dim=-1, | |
| ) | |
| text_inputs["attention_mask"] = torch.cat([text_inputs["attention_mask"], torch.tensor([1])], dim=-1) | |
| model_inputs.update(modality_input) | |
| model_inputs.update(text_inputs) | |
| # position_ids (dim, len) | |
| if position_id_func is None: # default position_ids | |
| position_id_returns = {"position_ids": torch.arange(0, len(text_inputs["input_ids"])).unsqueeze(0)} | |
| else: # customized position_ids | |
| input_ids = text_inputs["input_ids"].clone() | |
| attention_mask = text_inputs["attention_mask"].clone() | |
| position_id_returns = position_id_func( | |
| input_ids=input_ids.unsqueeze(0), attention_mask=attention_mask.unsqueeze(0), **multimodal_config | |
| ) | |
| model_inputs.update(position_id_returns) | |
| input_ids, mask_dict = mask_input_ids(modality_info, model_inputs["input_ids"]) | |
| model_inputs["input_ids"] = input_ids | |
| model_inputs.update(mask_dict) | |
| return [model_inputs] | |