bernini-diffusers-v2-demo / veomni /data /multimodal /multimodal_transform.py
multimodalart's picture
multimodalart HF Staff
Bernini-Diffusers-v2 r2v demo
fed6c68 verified
Raw
History Blame Contribute Delete
14.1 kB
# 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]