model: support parakeet nemotron encoder (#23568)

Co-authored-by: trangdough <trangtdo22@gmail.com>
This commit is contained in:
Yuhao Yang
2026-04-25 11:00:23 +08:00
committed by GitHub
co-authored by trangdough
parent 465abadd3c
commit 4a3fe2a091
15 changed files with 1289 additions and 116 deletions
+3 -1
View File
@@ -200,11 +200,13 @@ class ModelConfig:
# Models expose audio_config at different nesting levels:
# - top-level audio_config: e.g. Qwen2Audio
# - thinker_config.audio_config: Qwen3-Omni, Qwen3-ASR (nested thinker arch)
# - is_audio_model(): Whisper, Qwen3-ASR (architecture-based fallback)\
# - sound_config: Nemotron AVLM with Parakeet audio encoder
# - is_audio_model(): Whisper, Qwen3-ASR (architecture-based fallback)
# TODO: Handle this more robustly by standardizing the config structure in the future
self.is_audio_understandable_model = enable_multimodal and (
hasattr(self.hf_config, "audio_config")
or hasattr(getattr(self.hf_config, "thinker_config", None), "audio_config")
or getattr(self.hf_config, "sound_config", None) is not None
or is_audio_model(self.hf_config.architectures)
)
@@ -38,6 +38,7 @@ class NemotronH_Nano_VL_V2_Config(PretrainedConfig):
self,
vision_config=None,
llm_config=None,
sound_config=None,
force_image_size: int = 512,
patch_size: int = 16,
downsample_ratio=0.5,
@@ -51,6 +52,9 @@ class NemotronH_Nano_VL_V2_Config(PretrainedConfig):
img_context_token: str = "<image>",
img_start_token: str = "<img>",
img_end_token: str = "</img>",
audio_context_token: str = "<so_embedding>",
audio_start_token: str = "<so_start>",
audio_end_token: str = "<so_end>",
norm_mean: tuple[float, float, float] | list[float] = IMAGENET_MEAN,
norm_std: tuple[float, float, float] | list[float] = IMAGENET_STD,
use_thumbnail: bool = True,
@@ -68,6 +72,12 @@ class NemotronH_Nano_VL_V2_Config(PretrainedConfig):
self.llm_config = NemotronHConfig()
self.raw_vision_config = {}
# Audio (Parakeet) config: stored as a PretrainedConfig sub-object
if sound_config is not None and isinstance(sound_config, dict):
self.sound_config = PretrainedConfig.from_dict(sound_config)
else:
self.sound_config = sound_config
# Assign configuration values
vision_image_size = self.raw_vision_config.get("image_size", force_image_size)
vision_patch_size = self.raw_vision_config.get("patch_size", patch_size)
@@ -97,6 +107,28 @@ class NemotronH_Nano_VL_V2_Config(PretrainedConfig):
self.use_thumbnail = use_thumbnail
self.img_start_token = img_start_token
self.img_end_token = img_end_token
self.audio_context_token = audio_context_token
self.audio_start_token = audio_start_token
self.audio_end_token = audio_end_token
# Dynamic resolution: from vision_config top-level
self.min_num_patches = self.raw_vision_config.get("min_num_patches", 0)
self.max_num_patches = self.raw_vision_config.get("max_num_patches", 0)
self.dynamic_resolution = self.min_num_patches > 0
# Video temporal compression: from vision_config top-level
self.video_temporal_patch_size = self.raw_vision_config.get(
"video_temporal_patch_size", 1
)
self.separate_video_embedder = self.raw_vision_config.get(
"separate_video_embedder", True
)
self.video_target_num_patches = self.raw_vision_config.get(
"video_target_num_patches", 0
)
self.video_maintain_aspect_ratio = self.raw_vision_config.get(
"video_maintain_aspect_ratio", True
)
def create_radio_config(self):
config = self.raw_vision_config
@@ -110,5 +142,11 @@ class NemotronH_Nano_VL_V2_Config(PretrainedConfig):
model_name=model_name,
reg_tokens=reg_tokens,
image_size=image_size,
min_num_patches=self.min_num_patches,
max_num_patches=self.max_num_patches,
video_temporal_patch_size=self.video_temporal_patch_size,
separate_video_embedder=self.separate_video_embedder,
video_target_num_patches=self.video_target_num_patches,
video_maintain_aspect_ratio=self.video_maintain_aspect_ratio,
)
return radio_config
+74
View File
@@ -0,0 +1,74 @@
# Copyright 2026 SGLang Team
# 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.
# ==============================================================================
# Adapted from https://github.com/vllm-project/vllm/blob/main/vllm/transformers_utils/configs/parakeet.py
from dataclasses import dataclass
from transformers import ParakeetEncoderConfig, PretrainedConfig
class ParakeetConfig(ParakeetEncoderConfig):
def __init__(
self,
llm_hidden_size: int,
projection_hidden_size: int,
projection_bias: bool,
sampling_rate: int,
projection_eps: float = 1e-5,
**kwargs,
):
super().__init__(**kwargs)
self.llm_hidden_size = llm_hidden_size
self.projection_hidden_size = projection_hidden_size
self.projection_bias = projection_bias
self.sampling_rate = sampling_rate
self.projection_eps = projection_eps
@staticmethod
def from_hf_config(
config: PretrainedConfig, *, llm_hidden_size: int, max_model_len: int
) -> "ParakeetConfig":
assert isinstance(config, PretrainedConfig)
return ParakeetConfig(
**config.to_dict(),
scale_input=False,
attention_bias=False,
llm_hidden_size=llm_hidden_size,
max_position_embeddings=max_model_len + 1,
)
@dataclass(kw_only=True, frozen=True)
class ExtractorConfig:
feature_size: int
sampling_rate: int
subsampling_factor: int
subsampling_conv_kernel_size: int
subsampling_conv_stride: int
hop_length: int = 160
clip_duration_s: int = 30
clip_min_duration_s: float = 0.1
@staticmethod
def from_hf_config(config: PretrainedConfig) -> "ExtractorConfig":
assert isinstance(config, PretrainedConfig)
hop_length = int(getattr(config, "hop_length", ExtractorConfig.hop_length))
return ExtractorConfig(
feature_size=config.num_mel_bins,
sampling_rate=config.sampling_rate,
hop_length=hop_length,
subsampling_factor=config.subsampling_factor,
subsampling_conv_kernel_size=config.subsampling_conv_kernel_size,
subsampling_conv_stride=config.subsampling_conv_stride,
)
+12
View File
@@ -74,6 +74,12 @@ class RadioConfig(PretrainedConfig):
norm_mean: tuple[float, float, float] | list = OPENAI_CLIP_MEAN,
norm_std: tuple[float, float, float] | list = OPENAI_CLIP_STD,
reg_tokens: int | None = None,
min_num_patches: int = 0,
max_num_patches: int = 0,
video_temporal_patch_size: int = 1,
separate_video_embedder: bool = True,
video_target_num_patches: int = 0,
video_maintain_aspect_ratio: bool = True,
drop_path_rate: float = 0.0,
dropout: float = 0.0,
**kwargs,
@@ -101,6 +107,12 @@ class RadioConfig(PretrainedConfig):
list(norm_std) if isinstance(norm_std, (tuple, list)) else norm_std
)
self.reg_tokens = reg_tokens
self.min_num_patches = min_num_patches
self.max_num_patches = max_num_patches
self.video_temporal_patch_size = video_temporal_patch_size
self.separate_video_embedder = separate_video_embedder
self.video_target_num_patches = video_target_num_patches
self.video_maintain_aspect_ratio = video_maintain_aspect_ratio
self.drop_path_rate = drop_path_rate
self.dropout = dropout
super().__init__(**kwargs)
@@ -645,9 +645,10 @@ class ChatCompletionRequest(BaseModel):
stream_reasoning: bool = True
chat_template_kwargs: Optional[Dict] = None
# SGLang multimodal tiling controls (extensions)
# SGLang multimodal controls (extensions)
max_dynamic_patch: Optional[int] = None
min_dynamic_patch: Optional[int] = None
use_audio_in_video: bool = False
# Custom logit processor for advanced sampling control
custom_logit_processor: Optional[Union[List[Optional[str]], str]] = None
@@ -419,6 +419,7 @@ class OpenAIServingChat(OpenAIServingBase):
image_max_dynamic_patch=img_max_dynamic_patch,
video_max_dynamic_patch=vid_max_dynamic_patch,
max_dynamic_patch=getattr(request, "max_dynamic_patch", None),
use_audio_in_video=getattr(request, "use_audio_in_video", False),
)
return adapted_request, request
+2
View File
@@ -153,6 +153,8 @@ class GenerateReqInput(BaseReq):
video_data: Optional[MultimodalDataInputFormat] = None
# The audio input. Like image data, it can be a file name, a url, or base64 encoded string.
audio_data: Optional[MultimodalDataInputFormat] = None
# Whether to extract and process audio from video inputs.
use_audio_in_video: bool = False
# The sampling_params. See descriptions below.
sampling_params: Optional[Union[List[Dict], Dict]] = None
# Whether to return logprobs.
+3 -1
View File
@@ -333,6 +333,7 @@ class InternVisionEncoder(nn.Module):
def forward(
self,
inputs_embeds,
cu_seqlens=None,
output_hidden_states: Optional[bool] = None,
return_dict: Optional[bool] = None,
) -> Union[Tuple, BaseModelOutput]:
@@ -366,7 +367,8 @@ class InternVisionEncoder(nn.Module):
encoder_states = () if output_hidden_states else None
hidden_states = inputs_embeds
cu_seqlens = SingletonCache()
if cu_seqlens is None:
cu_seqlens = SingletonCache()
for idx, encoder_layer in enumerate(self.layers):
if output_hidden_states:
+171 -20
View File
@@ -35,8 +35,10 @@ from sglang.srt.managers.schedule_batch import (
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.models.nemotron_h import NemotronHForCausalLM
from sglang.srt.models.parakeet import ProjectedParakeet
from sglang.srt.models.radio import RadioModel
from sglang.srt.multimodal.evs import EVS, EVSConfig
from sglang.srt.multimodal.evs.evs_module import VideoEVSDataItem
from sglang.srt.utils import add_prefix
logger = logging.getLogger(__name__)
@@ -66,9 +68,13 @@ class NemotronH_Nano_VL_V2(EVS):
)
vit_hidden_size = config.vit_hidden_size
self.rmsnorm_hidden_size = vit_hidden_size * int(1 / self.downsample_ratio) ** 2
self.rmsnorm_hidden_size = (
vit_hidden_size * int(round(1 / self.downsample_ratio)) ** 2
)
vision_projection_hidden_size = config.projector_hidden_size
llm_hidden_size = config.llm_config.hidden_size
self.llm_hidden_size = llm_hidden_size
self.model_dtype = self.language_model.config.torch_dtype
self.mlp1 = nn.Sequential(
RMSNorm(
@@ -82,18 +88,58 @@ class NemotronH_Nano_VL_V2(EVS):
),
ReLU2(),
nn.Linear(vision_projection_hidden_size, llm_hidden_size, bias=False),
).to(self.language_model.config.torch_dtype)
).to(self.model_dtype)
self.sound_encoder: ProjectedParakeet | None = None
if getattr(config, "sound_config", None) is not None:
logger.info(
"Found sound config, initializing sound encoder for Nemotron AVLM"
)
self.sound_encoder = ProjectedParakeet(
config.sound_config,
dtype=self.language_model.config.torch_dtype,
llm_hidden_size=llm_hidden_size,
max_model_len=getattr(config, "max_model_len", 8192),
)
self.config = config
def pad_input_ids(self, input_ids: list[int], mm_inputs: MultimodalInputs):
# Get all special token IDs
im_start_id: int = mm_inputs.im_start_id
im_end_id: int = mm_inputs.im_end_id
media_token_pairs = [(im_start_id, im_end_id)]
helper = MultiModalityDataPaddingPatternTokenPairs(media_token_pairs)
visual_items = [item for item in mm_inputs.mm_items if not item.is_audio()]
audio_items = [item for item in mm_inputs.mm_items if item.is_audio()]
return helper.pad_input_tokens(input_ids, mm_inputs)
all_data_offsets = []
if visual_items:
mm_inputs.mm_items = visual_items
helper = MultiModalityDataPaddingPatternTokenPairs(
[(im_start_id, im_end_id)]
)
input_ids = helper.pad_input_tokens(input_ids, mm_inputs)
all_data_offsets.extend(mm_inputs.data_offsets)
audio_start_id = getattr(mm_inputs, "audio_start_id", None)
audio_end_id = getattr(mm_inputs, "audio_end_id", None)
if audio_items and audio_start_id is not None and audio_end_id is not None:
mm_inputs.mm_items = audio_items
helper = MultiModalityDataPaddingPatternTokenPairs(
[(audio_start_id, audio_end_id)]
)
input_ids = helper.pad_input_tokens(input_ids, mm_inputs)
all_data_offsets.extend(mm_inputs.data_offsets)
mm_inputs.mm_items = visual_items + audio_items
mm_inputs.data_offsets = all_data_offsets
if audio_items:
for item in visual_items:
if isinstance(item, VideoEVSDataItem):
item.pre_chunked_input_ids = input_ids
return input_ids
def pixel_shuffle(self, x: torch.Tensor, scale_factor: float = 0.5) -> torch.Tensor:
n, w, h, c = x.size()
@@ -118,28 +164,64 @@ class NemotronH_Nano_VL_V2(EVS):
x = x.permute(0, 2, 1, 3).contiguous()
return x
def extract_feature_dynamic(self, pixel_values_list: list[torch.Tensor]):
"""Extract features from variable-size images (dynamic resolution).
Each image has different spatial dimensions. They are passed as a list
to RADIO which handles ragged packing with cu_seqlens internally.
"""
features, num_patches_list = self.vision_model(pixel_values_list)
patch_size = self.config.patch_size
results = []
offset = 0
for i, num_patches in enumerate(num_patches_list):
img_feats = features[0, offset : offset + num_patches]
h_patches = pixel_values_list[i].shape[-2] // patch_size
w_patches = pixel_values_list[i].shape[-1] // patch_size
img_feats = img_feats.reshape(1, h_patches, w_patches, -1)
img_feats = self.pixel_shuffle(img_feats, self.downsample_ratio)
img_feats = img_feats.view(-1, self.rmsnorm_hidden_size)
img_feats = self.mlp1(img_feats)
results.append(img_feats)
offset += num_patches
return torch.cat(results, dim=0)
def extract_video_feature_temporal(self, pixel_values, num_frames):
"""Extract video features with temporal compression (tubelet grouping)."""
vit_embeds = self.vision_model(pixel_values, num_frames=num_frames)
num_tubelets = vit_embeds.shape[0]
patch_size = self.config.patch_size
h_patches = pixel_values.shape[-2] // patch_size
w_patches = pixel_values.shape[-1] // patch_size
vit_embeds = vit_embeds.reshape(num_tubelets, h_patches, w_patches, -1)
vit_embeds = self.pixel_shuffle(vit_embeds, self.downsample_ratio)
vit_embeds = vit_embeds.view(-1, self.rmsnorm_hidden_size)
vit_embeds = self.mlp1(vit_embeds)
vit_embeds = vit_embeds.view(num_tubelets, -1, self.llm_hidden_size)
return vit_embeds
def get_input_embeddings(self):
return self.language_model.get_input_embeddings()
def extract_feature(self, pixel_values):
# Process images in a micro-batch of at most 128 frames per call
# This is done on purpose to ensure peak GPU ram usage of huge batch
# (namely for really long videos with EVS ON) won't cause any problems
# as we don't support chunked prefill for video media
micro_batch_size = 128
n = pixel_values.shape[0]
patch_size = self.config.patch_size
h_patches = pixel_values.shape[-2] // patch_size
w_patches = pixel_values.shape[-1] // patch_size
vit_embeds_list = []
for i in range(0, n, micro_batch_size):
vit_embeds = self.vision_model(pixel_values[i : i + micro_batch_size])
vit_embeds = vit_embeds.to(dtype=torch.bfloat16)
h = w = int(vit_embeds.shape[1] ** 0.5)
vit_embeds = vit_embeds.reshape(vit_embeds.shape[0], h, w, -1)
chunk = pixel_values[i : i + micro_batch_size]
batch_size = chunk.shape[0]
vit_embeds = self.vision_model(chunk)
vit_embeds = vit_embeds.to(dtype=self.model_dtype)
vit_embeds = vit_embeds.reshape(batch_size, h_patches, w_patches, -1)
vit_embeds = self.pixel_shuffle(
vit_embeds, scale_factor=self.downsample_ratio
)
vit_embeds = vit_embeds.view(-1, self.rmsnorm_hidden_size)
vit_embeds = self.mlp1(vit_embeds)
vit_embeds = vit_embeds.view(n, -1, self.rmsnorm_hidden_size)
vit_embeds = vit_embeds.view(batch_size, -1, self.llm_hidden_size)
vit_embeds_list.append(vit_embeds)
vit_embeds = torch.cat(vit_embeds_list, dim=0)
return vit_embeds
@@ -151,6 +233,11 @@ class NemotronH_Nano_VL_V2(EVS):
Returns:
image_features (`torch.Tensor`): Image feature tensor of shape `(num_images, image_length, embed_dim)`).
"""
is_dynamic = any(getattr(item, "is_dynamic", False) for item in items)
if is_dynamic:
pixel_values_list = [item.feature for item in items]
return self.extract_feature_dynamic(pixel_values_list)
pixel_values = torch.cat([item.feature for item in items])
image_features = self.extract_feature(pixel_values)
return image_features
@@ -163,9 +250,60 @@ class NemotronH_Nano_VL_V2(EVS):
video_features (`torch.Tensor`): Video feature tensor of shape `(num_videos, video_length, embed_dim)`).
"""
pixel_values = torch.cat([item.feature for item in items])
if getattr(self.config, "video_temporal_patch_size", 1) > 1:
num_frames = pixel_values.shape[0]
return self.extract_video_feature_temporal(pixel_values, num_frames)
video_features = self.extract_feature(pixel_values)
return video_features
def get_audio_feature(self, items: list[MultimodalDataItem]):
"""
Encode audio features through the Parakeet sound encoder.
Each item carries mel spectrogram features, an attention mask, and a
clip count. Multiple clips per audio item are grouped and concatenated
(trimmed to valid output lengths) to form a single embedding per item.
"""
assert self.sound_encoder is not None
all_features = []
all_masks = []
all_num_clips = []
for item in items:
all_features.append(item.feature)
all_masks.append(item.feature_attention_mask)
all_num_clips.append(item.audio_num_clips)
input_audio_features = torch.cat(all_features, dim=0)
feature_attention_mask = torch.cat(all_masks, dim=0)
target_device = next(self.sound_encoder.parameters()).device
input_audio_features = input_audio_features.to(
dtype=self.language_model.config.torch_dtype, device=target_device
)
feature_attention_mask = feature_attention_mask.to(device=target_device)
sound_embeds = self.sound_encoder(input_audio_features, feature_attention_mask)
valid_input_lens = feature_attention_mask.sum(dim=1)
valid_output_lens = (
self.sound_encoder.encoder._get_subsampling_output_length(valid_input_lens)
.long()
.tolist()
)
grouped_embeds = []
clip_offset = 0
for num_clips in all_num_clips:
embeds = []
for clip_idx in range(clip_offset, clip_offset + num_clips):
valid_len = valid_output_lens[clip_idx]
embeds.append(sound_embeds[clip_idx, :valid_len])
grouped_embeds.append(torch.cat(embeds, dim=0))
clip_offset += num_clips
return torch.cat(grouped_embeds, dim=0)
@torch.no_grad()
def forward(
self,
@@ -174,15 +312,19 @@ class NemotronH_Nano_VL_V2(EVS):
forward_batch: ForwardBatch,
get_embedding: bool = False,
):
data_embedding_funcs = {
Modality.IMAGE: self.get_image_feature,
Modality.VIDEO: self.get_video_feature,
}
if self.sound_encoder is not None:
data_embedding_funcs[Modality.AUDIO] = self.get_audio_feature
hidden_states = general_mm_embed_routine(
input_ids=input_ids,
forward_batch=forward_batch,
language_model=self.language_model,
multimodal_model=self,
data_embedding_funcs={
Modality.IMAGE: self.get_image_feature,
Modality.VIDEO: self.get_video_feature,
},
data_embedding_funcs=data_embedding_funcs,
positions=positions,
)
return hidden_states
@@ -199,9 +341,13 @@ class NemotronH_Nano_VL_V2(EVS):
def is_vision_weights(name: str) -> bool:
return name.startswith("vision_model.radio_model.")
def is_sound_weights(name: str) -> bool:
return name.startswith("sound")
# Separate weights by component
llm_weights = []
vision_weights = []
sound_weights = []
for name, w in weights:
if is_llm(name):
@@ -215,10 +361,15 @@ class NemotronH_Nano_VL_V2(EVS):
default_weight_loader(param, w)
elif is_vision_weights(name):
# Convert: vision_model.radio_model.* → radio_model.*
hf_key = name[len("vision_model.") :] # Remove "vision_model." prefix
hf_key = name[len("vision_model.") :]
vision_weights.append((hf_key, w))
elif is_sound_weights(name):
sound_weights.append((name, w))
self.language_model.load_weights(llm_weights)
self.vision_model.load_weights(vision_weights)
if self.sound_encoder is not None and len(sound_weights) > 0:
self.sound_encoder.load_weights(sound_weights)
EntryClass = [NemotronH_Nano_VL_V2]
+182
View File
@@ -0,0 +1,182 @@
# Copyright 2026 SGLang Team
# 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.
# ==============================================================================
# Adapted from https://github.com/vllm-project/vllm/blob/main/vllm/model_executor/models/parakeet.py
#
# Audio encoder component used by models/nano_nemotron_vl.py
from collections.abc import Iterable
from dataclasses import asdict
import numpy as np
import torch
import torch.nn as nn
from transformers import ParakeetEncoder as HFParakeetEncoder
from transformers import ParakeetFeatureExtractor, PretrainedConfig
from sglang.srt.configs.parakeet import ExtractorConfig, ParakeetConfig
from sglang.srt.layers.activation import ReLU2
from sglang.srt.layers.layernorm import RMSNorm
from sglang.srt.model_loader.weight_utils import default_weight_loader
class ParakeetProjection(nn.Module):
def __init__(self, config: ParakeetConfig) -> None:
super().__init__()
sound_hidden_size = config.hidden_size
proj_hidden_size = config.projection_hidden_size
llm_hidden_size = config.llm_hidden_size
bias = config.projection_bias
self.norm = RMSNorm(sound_hidden_size, eps=config.projection_eps)
self.linear1 = nn.Linear(sound_hidden_size, proj_hidden_size, bias=bias)
self.activation = ReLU2()
self.linear2 = nn.Linear(proj_hidden_size, llm_hidden_size, bias=bias)
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
hidden_states = self.norm(hidden_states)
hidden_states = self.linear1(hidden_states)
hidden_states = self.activation(hidden_states)
hidden_states = self.linear2(hidden_states)
return hidden_states
class ProjectedParakeet(nn.Module):
def __init__(
self,
config: PretrainedConfig,
*,
dtype: torch.dtype,
llm_hidden_size: int,
max_model_len: int,
) -> None:
super().__init__()
self.config = ParakeetConfig.from_hf_config(
config, llm_hidden_size=llm_hidden_size, max_model_len=max_model_len
)
self.encoder = HFParakeetEncoder(self.config)
self.encoder = self.encoder.to(dtype)
self.projection = ParakeetProjection(self.config)
self.projection = self.projection.to(dtype)
def forward(
self, input_features: torch.Tensor, attention_mask: torch.Tensor | None = None
) -> torch.Tensor:
outputs = self.encoder(
input_features=input_features, attention_mask=attention_mask
)
outputs = outputs.last_hidden_state
outputs = self.projection(outputs)
return outputs
def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]:
loaded_params: set[str] = set()
params_dict = dict(self.named_parameters())
buffers_dict = dict(self.named_buffers())
if isinstance(weights, dict):
weights_list = list(weights.items())
else:
weights_list = list(weights)
for name, weight in weights_list:
if name.startswith("sound_encoder.encoder.feature_extractor."):
continue
if name.startswith("sound_encoder."):
target_name = name[len("sound_encoder.") :]
elif name.startswith("sound_projection."):
target_name = f"projection.{name[len('sound_projection.'):]}"
else:
continue
target = params_dict.get(target_name)
if target is None:
target = buffers_dict.get(target_name)
if target is None:
continue
weight_loader = getattr(target, "weight_loader", default_weight_loader)
with torch.no_grad():
weight_loader(target, weight)
loaded_params.add(target_name)
return loaded_params
class ParakeetExtractor(ParakeetFeatureExtractor):
def __init__(self, config: PretrainedConfig) -> None:
self.config = ExtractorConfig.from_hf_config(config)
super().__init__(**asdict(self.config))
self._clip_target_samples = int(
round(self.config.clip_duration_s * self.sampling_rate)
)
self._tail_min_samples = int(
round(self.config.clip_min_duration_s * self.sampling_rate)
)
def _clip_sizes(self, audio_len: int) -> list[int]:
audio_len = max(audio_len, self._tail_min_samples)
num_full_clips, remainder = divmod(audio_len, self._clip_target_samples)
clip_sizes = [self._clip_target_samples] * num_full_clips
if remainder > 0:
clip_sizes.append(max(remainder, self._tail_min_samples))
return clip_sizes
def _subsampling_output_length(self, length: int) -> int:
import math
kernel_size = self.config.subsampling_conv_kernel_size
stride = self.config.subsampling_conv_stride
num_layers = int(math.log2(self.config.subsampling_factor))
add_pad = (kernel_size - 1) // 2 * 2 - kernel_size
for _ in range(num_layers):
length = int(math.floor((length + add_pad) / stride + 1.0))
return max(1, length)
def audio_token_count(self, audio_len: int) -> int:
total_tokens = 0
for clip_size in self._clip_sizes(audio_len):
num_frames = clip_size // self.hop_length
total_tokens += self._subsampling_output_length(num_frames)
return max(1, total_tokens)
def split_audio_into_clips(self, audio: np.ndarray) -> list[np.ndarray]:
assert audio.ndim == 1
audio_len = int(audio.shape[0])
clip_sizes = self._clip_sizes(audio_len)
target_len = sum(clip_sizes)
if audio_len < target_len:
audio = np.pad(audio, (0, target_len - audio_len))
clips = list[np.ndarray]()
offset = 0
for clip_size in clip_sizes:
clips.append(audio[offset : offset + clip_size])
offset += clip_size
return clips
def __call__(self, raw_speech: list[np.ndarray], *args, **kwargs):
audio_clips = list[np.ndarray]()
audio_num_clips = list[int]()
for audio in raw_speech:
clips = self.split_audio_into_clips(audio)
audio_clips.extend(clips)
audio_num_clips.append(len(clips))
outputs = super().__call__(audio_clips, *args, **kwargs)
outputs["audio_num_clips"] = audio_num_clips
return outputs
@staticmethod
def audio_length(raw_config: PretrainedConfig, audio_tokens: int) -> int:
config = ExtractorConfig.from_hf_config(raw_config)
return int(audio_tokens * config.subsampling_factor * config.hop_length)
+135 -57
View File
@@ -13,6 +13,7 @@
# ==============================================================================
# Adapted from https://github.com/vllm-project/vllm/blob/main/vllm/model_executor/models/radio.py
import logging
import math
from collections.abc import Iterable
from itertools import repeat
@@ -33,6 +34,8 @@ from sglang.srt.model_loader.weight_utils import (
)
from sglang.srt.models.internvl import InternVisionEncoder
logger = logging.getLogger(__name__)
input_dim_t: TypeAlias = int | tuple[int, int]
norm_t: TypeAlias = tuple[float, float, float] | torch.Tensor
@@ -105,7 +108,6 @@ class ClsToken(nn.Module):
class ViTPatchGenerator(nn.Module):
def __init__(
self,
# config: PretrainedConfig,
patch_size: int,
embed_dim: int,
input_dims: input_dim_t,
@@ -119,6 +121,8 @@ class ViTPatchGenerator(nn.Module):
register_multiple: int | None = None,
num_registers: int | None = None,
patch_bias: bool = False,
video_temporal_patch_size: int = 1,
separate_video_embedder: bool = True,
device=None,
dtype=None,
):
@@ -174,6 +178,17 @@ class ViTPatchGenerator(nn.Module):
nn.LayerNorm(embed_dim) if normalize_patches else nn.Identity()
)
self.video_temporal_patch_size = video_temporal_patch_size
self.video_embedder = None
self._video_embedder_loaded = False
if video_temporal_patch_size > 1 and separate_video_embedder:
self.video_embedder = nn.Linear(
3 * video_temporal_patch_size * patch_size * patch_size,
embed_dim,
bias=False,
**factory,
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
patches = self.embed_patches(x)
patches, pos_enc = self.apply_pos_enc(patches, input_size=x.shape[2:])
@@ -183,6 +198,40 @@ class ViTPatchGenerator(nn.Module):
return patches, pos_enc
return patches
def forward_video(self, x: torch.Tensor, temporal_patch_size: int) -> torch.Tensor:
"""Embed video frames with temporal compression via tubelet grouping."""
assert (
self.video_embedder is not None
), "video_embedder is required for temporal compression"
T = temporal_patch_size
num_frames = x.shape[0]
if num_frames % T != 0:
pad = T - (num_frames % T)
x = torch.cat(
[x, x[-1:].expand(pad, -1, -1, -1)],
dim=0,
)
padded_frames = x.shape[0]
num_tubelets = padded_frames // T
patches = self.im_to_patches(x)
num_spatial = patches.shape[1]
feat_dim = patches.shape[2]
patches = patches.reshape(num_tubelets, T, num_spatial, feat_dim)
patches = patches.permute(0, 2, 1, 3).reshape(
num_tubelets, num_spatial, T * feat_dim
)
patches = self.video_embedder(patches)
patches, _ = self.apply_pos_enc(patches, input_size=x.shape[2:])
patches = self.cls_token(patches)
patches = self.patch_normalizer(patches)
return patches
@property
def apply_cls_token(self):
return self.cls_token.enabled
@@ -319,66 +368,21 @@ class ViTPatchGenerator(nn.Module):
return pos_embed
if self.cpe_mode:
if self.training:
min_scale = math.sqrt(0.1)
scale = (
torch.rand(batch_size, 1, 1, device=pos_embed.device)
* (1 - min_scale)
+ min_scale
)
aspect_min = math.log(3 / 4)
aspect_max = -aspect_min
aspect = torch.exp(
torch.rand(batch_size, 1, 1, device=pos_embed.device)
* (aspect_max - aspect_min)
+ aspect_min
)
max_dim = max(input_dims)
pos_embed = F.interpolate(
pos_embed.float(),
size=(max_dim, max_dim),
align_corners=False,
mode="bilinear",
).to(pos_embed.dtype)
scale_x = scale * aspect
scale_y = scale * (1 / aspect)
scale_xy = torch.stack([scale_x, scale_y], dim=-1).clamp_(0, 1)
pos_xy = torch.rand(batch_size, 1, 1, 2, device=pos_embed.device) * (
1 - scale_xy
)
lin_x = torch.linspace(
0, 1, steps=input_dims[1], device=pos_embed.device
)[None, None].expand(batch_size, input_dims[0], -1)
lin_y = torch.linspace(
0, 1, steps=input_dims[0], device=pos_embed.device
)[None, :, None].expand(batch_size, -1, input_dims[1])
lin_xy = torch.stack([lin_x, lin_y], dim=-1)
grid_xy = lin_xy * scale_xy + pos_xy
# Convert to [-1, 1] range
grid_xy.mul_(2).sub_(1)
pos_embed = F.grid_sample(
pos_embed.float().expand(batch_size, -1, -1, -1),
grid=grid_xy,
mode="bilinear",
padding_mode="zeros",
align_corners=True,
).to(pos_embed.dtype)
else:
max_dim = max(input_dims)
pos_embed = F.interpolate(
pos_embed.float(),
size=(max_dim, max_dim),
align_corners=True,
mode="bilinear",
).to(pos_embed.dtype)
pos_embed = window_select(pos_embed)
pos_embed = window_select(pos_embed)
else:
pos_embed = window_select(pos_embed)
if pos_embed.shape[-2:] != input_dims:
pos_embed = F.interpolate(
pos_embed.float(), size=input_dims, align_corners=True, mode="bilinear"
pos_embed.float(), size=input_dims, align_corners=False, mode="bilinear"
).to(pos_embed.dtype)
pos_embed = pos_embed.flatten(2).permute(0, 2, 1)
@@ -435,6 +439,9 @@ class RadioInternVisionModel(nn.Module):
max_img_size = int(
round(config.max_img_size / config.patch_size) * config.patch_size
)
video_temporal_patch_size = getattr(config, "video_temporal_patch_size", 1)
separate_video_embedder = getattr(config, "separate_video_embedder", True)
self.patch_generator = ViTPatchGenerator(
config.patch_size,
config.hidden_size,
@@ -442,6 +449,8 @@ class RadioInternVisionModel(nn.Module):
max_input_dims=max_img_size,
cls_token=True,
register_multiple=config.reg_tokens,
video_temporal_patch_size=video_temporal_patch_size,
separate_video_embedder=separate_video_embedder,
)
self.encoder = InternVisionEncoder(config=config, quant_config=quant_config)
@@ -485,12 +494,79 @@ class RadioModel(nn.Module):
def forward(
self,
pixel_values: torch.Tensor | None = None,
pixel_embeds: torch.Tensor | None = None,
pixel_values: torch.Tensor | list[torch.Tensor] | None = None,
num_frames: int | None = None,
) -> torch.FloatTensor:
if (
num_frames is not None
and getattr(self.config, "video_temporal_patch_size", 1) > 1
):
return self._forward_video_temporal(pixel_values, num_frames)
if isinstance(pixel_values, list):
return self._forward_dynamic(pixel_values)
y = self.model(pixel_values)
return self._extract_final(y)
def _forward_dynamic(
self, images: list[torch.Tensor]
) -> tuple[torch.Tensor, list[int]]:
"""Process variable-size images with ragged packing via cu_seqlens."""
patch_gen = self.model.patch_generator
all_patches = []
seqlens = [0]
for img in images:
patches = patch_gen(img)
seq_len = patches.shape[1]
all_patches.append(patches.squeeze(0))
seqlens.append(seqlens[-1] + seq_len)
hidden = torch.cat(all_patches, dim=0).unsqueeze(0)
cu_seqlens = torch.tensor(seqlens, dtype=torch.int32, device=hidden.device)
out = self.model.encoder.forward(inputs_embeds=hidden, cu_seqlens=cu_seqlens)
features = out.last_hidden_state
num_skip = patch_gen.num_skip
per_image_features = []
num_patches_list = []
for i in range(len(images)):
start = seqlens[i] + num_skip
end = seqlens[i + 1]
per_image_features.append(features[0, start:end])
num_patches_list.append(end - start)
return (
torch.cat(per_image_features, dim=0).unsqueeze(0),
num_patches_list,
)
def _forward_video_temporal(
self, pixel_values: torch.Tensor, num_frames: int
) -> torch.Tensor:
"""Process video frames with temporal compression (tubelet grouping)."""
T = self.config.video_temporal_patch_size
patch_gen = self.model.patch_generator
patches = patch_gen.forward_video(pixel_values, T)
num_tubelets = patches.shape[0]
seq_per_tubelet = patches.shape[1]
cu_seqlens = torch.arange(
0,
(num_tubelets + 1) * seq_per_tubelet,
seq_per_tubelet,
dtype=torch.int32,
device=patches.device,
)
packed = patches.reshape(1, -1, patches.shape[-1])
out = self.model.encoder.forward(inputs_embeds=packed, cu_seqlens=cu_seqlens)
features = out.last_hidden_state.reshape(num_tubelets, seq_per_tubelet, -1)
num_skip = patch_gen.num_skip
return features[:, num_skip:]
def load_weights(self, weights) -> set[str]:
remap_substrings = {
"attn": "attn.attn",
@@ -520,6 +596,8 @@ class RadioModel(nn.Module):
weight_loader = getattr(param, "weight_loader", default_weight_loader)
weight_loader(param, weight)
loaded_params.add(name)
if "video_embedder" in name:
self.model.patch_generator._video_embedder_loaded = True
return loaded_params
@@ -0,0 +1,89 @@
# Copyright 2025 SGLang Team
# 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.
# ==============================================================================
"""Extract audio from video bytes using PyAV (in-process, CUDA-safe).
PyAV wraps FFmpeg's C libraries in-process, avoiding subprocess forks which
would crash CUDA-active workers.
"""
import io
import logging
import numpy as np
logger = logging.getLogger(__name__)
def extract_audio_from_video_bytes(
video_bytes: bytes,
target_sr: int = 16000,
) -> np.ndarray | None:
"""Extract mono audio from video bytes at the target sample rate.
Args:
video_bytes: Raw video file bytes (e.g. MP4).
target_sr: Target sample rate for the output waveform.
Returns:
1-D float32 numpy array of audio samples, or None if the video
has no audio track.
"""
try:
import av
except ImportError:
logger.warning(
"PyAV (av) is not installed. Cannot extract audio from video. "
"Install with: pip install av"
)
return None
try:
container = av.open(io.BytesIO(video_bytes))
except Exception:
logger.warning("Failed to open video bytes for audio extraction")
return None
if not container.streams.audio:
container.close()
return None
try:
audio_stream = container.streams.audio[0]
native_sr = audio_stream.rate or target_sr
resampler = av.audio.resampler.AudioResampler(
format="flt",
layout="mono",
rate=target_sr,
)
chunks = []
for frame in container.decode(audio=0):
resampled = resampler.resample(frame)
for rf in resampled:
arr = rf.to_ndarray().flatten()
chunks.append(arr)
container.close()
if not chunks:
return None
waveform = np.concatenate(chunks).astype(np.float32)
return waveform
except Exception:
logger.warning("Error extracting audio from video", exc_info=True)
container.close()
return None
@@ -1,4 +1,6 @@
# copy from https://huggingface.co/OpenGVLab/InternVL3-1B
import math
import torch
import torchvision.transforms as T
from PIL import Image
@@ -113,3 +115,240 @@ def image_to_pixel_values(
pixel_values = [transform(image) for image in images]
pixel_values = torch.stack(pixel_values)
return pixel_values
def compute_dynamic_image_size(
orig_w: int,
orig_h: int,
patch_size: int,
downsample_ratio: float,
min_num_patches: int,
max_num_patches: int,
) -> tuple[int, int, int]:
"""Compute optimal resize dimensions for dynamic resolution.
The image is resized (not tiled) to a variable size that respects the
aspect ratio while staying within the patch budget. Dimensions are
snapped to multiples of ``patch_size * ds`` so that pixel-shuffle
downsampling produces integer grid sizes.
Returns:
(target_w, target_h, num_tokens) where num_tokens is the
post-pixel-shuffle token count.
"""
ds = int(1 / downsample_ratio)
snap = patch_size * ds
pw = max(1, round(orig_w / patch_size))
ph = max(1, round(orig_h / patch_size))
native_patches = pw * ph
budget = min(native_patches, max_num_patches)
budget = max(budget, min_num_patches)
factor = math.sqrt(budget / max(native_patches, 1))
factor = min(factor, 1.0)
target_pw = max(ds, int(round(pw * factor / ds)) * ds)
target_ph = max(ds, int(round(ph * factor / ds)) * ds)
if target_pw * target_ph < min_num_patches:
up = math.sqrt(min_num_patches / (target_pw * target_ph))
target_pw = max(ds, int(math.ceil(target_pw * up / ds)) * ds)
target_ph = max(ds, int(math.ceil(target_ph * up / ds)) * ds)
if target_pw * target_ph > max_num_patches:
down = math.sqrt(max_num_patches / (target_pw * target_ph))
target_pw = max(ds, int(math.floor(target_pw * down / ds)) * ds)
target_ph = max(ds, int(math.floor(target_ph * down / ds)) * ds)
target_w = target_pw * patch_size
target_h = target_ph * patch_size
num_tokens = (target_pw * target_ph) // (ds * ds)
return target_w, target_h, num_tokens
def dynamic_resize_image(
image: Image.Image,
patch_size: int,
downsample_ratio: float,
min_num_patches: int,
max_num_patches: int,
mean: tuple[float, float, float] = IMAGENET_MEAN,
std: tuple[float, float, float] = IMAGENET_STD,
) -> tuple[torch.Tensor, int]:
"""Resize image for dynamic resolution and return pixel tensor + token count.
Returns:
(pixel_values [1, 3, H, W], num_tokens)
"""
orig_w, orig_h = image.size
target_w, target_h, num_tokens = compute_dynamic_image_size(
orig_w,
orig_h,
patch_size,
downsample_ratio,
min_num_patches,
max_num_patches,
)
image = image.convert("RGB")
image = image.resize((target_w, target_h), Image.BICUBIC)
transform = T.Compose(
[
T.ToTensor(),
T.Normalize(mean=mean, std=std),
]
)
pixel_values = transform(image).unsqueeze(0)
return pixel_values, num_tokens
def resize_image_to_pixels(
image: Image.Image,
target_w: int,
target_h: int,
mean: tuple[float, float, float] = IMAGENET_MEAN,
std: tuple[float, float, float] = IMAGENET_STD,
) -> torch.Tensor:
"""Resize image to exact target dimensions and return normalized tensor.
Returns:
pixel_values tensor of shape [1, 3, target_h, target_w].
"""
image = image.convert("RGB")
image = image.resize((target_w, target_h), Image.BICUBIC)
transform = T.Compose(
[
T.ToTensor(),
T.Normalize(mean=mean, std=std),
]
)
return transform(image).unsqueeze(0)
def compute_budgeted_image_sizes(
image_sizes: list[tuple[int, int]],
total_token_budget: int,
patch_size: int,
downsample_ratio: float,
min_num_patches: int,
max_num_patches: int,
max_iterations: int = 10,
) -> list[tuple[int, int, int]]:
"""Compute per-image sizes that fit within a total token budget.
When multiple images share a prompt, their combined post-pixel-shuffle
tokens must not exceed ``total_token_budget``. This function iteratively
reduces per-image patch limits until the total fits.
Returns:
List of (target_w, target_h, num_tokens) per image.
"""
n = len(image_sizes)
if n == 0:
return []
ds = int(round(1 / downsample_ratio))
per_image_max = [max_num_patches] * n
results: list[tuple[int, int, int]] = []
for _ in range(max_iterations):
results = [
compute_dynamic_image_size(
orig_w,
orig_h,
patch_size,
downsample_ratio,
min_num_patches,
per_image_max[i],
)
for i, (orig_w, orig_h) in enumerate(image_sizes)
]
total_tokens = sum(num_tokens for _, _, num_tokens in results)
if total_tokens <= total_token_budget:
return results
scale = total_token_budget / total_tokens
for i in range(n):
current_patches = results[i][2] * ds * ds
per_image_max[i] = max(min_num_patches, int(current_patches * scale))
return results
def get_video_target_size_and_feature_size(
orig_w: int,
orig_h: int,
target_num_patches: int,
maintain_aspect_ratio: bool,
patch_size: int,
downsample_ratio: float,
) -> tuple[int, int, int]:
"""Compute target resize dimensions and post-downsample token count for video.
Single source of truth for video spatial dimensions — used by both
video_to_pixel_values (resize) and the processor (token counting).
Returns:
(target_w, target_h, feature_size) where feature_size is the
post-pixel-shuffle token count.
"""
ds = int(1 / downsample_ratio)
if target_num_patches > 0 and maintain_aspect_ratio:
aspect = orig_w / max(orig_h, 1)
ph = math.sqrt(target_num_patches / max(aspect, 1e-6))
pw = ph * aspect
target_pw = max(ds, int(round(pw / ds)) * ds)
target_ph = max(ds, int(round(ph / ds)) * ds)
elif target_num_patches > 0:
side = int(math.sqrt(target_num_patches))
target_pw = max(ds, int(round(side / ds)) * ds)
target_ph = target_pw
else:
target_pw = max(ds, round(orig_w / patch_size / ds) * ds)
target_ph = max(ds, round(orig_h / patch_size / ds) * ds)
target_w = target_pw * patch_size
target_h = target_ph * patch_size
feature_size = (target_pw // ds) * (target_ph // ds)
return target_w, target_h, feature_size
def video_to_pixel_values(
frame: Image.Image,
patch_size: int,
downsample_ratio: float,
target_num_patches: int,
maintain_aspect_ratio: bool,
mean: tuple[float, float, float] = IMAGENET_MEAN,
std: tuple[float, float, float] = IMAGENET_STD,
) -> tuple[torch.Tensor, int]:
"""Resize a single video frame for temporal compression pipeline.
Returns:
(pixel_values [1, 3, H, W], feature_size) where feature_size is
the post-pixel-shuffle token count.
"""
orig_w, orig_h = frame.size
target_w, target_h, feature_size = get_video_target_size_and_feature_size(
orig_w,
orig_h,
target_num_patches,
maintain_aspect_ratio,
patch_size,
downsample_ratio,
)
frame = frame.convert("RGB")
frame = frame.resize((target_w, target_h), Image.BICUBIC)
transform = T.Compose(
[
T.ToTensor(),
T.Normalize(mean=mean, std=std),
]
)
pixel_values = transform(frame).unsqueeze(0)
return pixel_values, feature_size
@@ -11,6 +11,8 @@
# See the License for the specific language governing permissions and
# limitations under the License.
import logging
import math
from math import sqrt
import numpy as np
@@ -18,16 +20,30 @@ import torch
from PIL import Image
from sglang.srt.configs.nano_nemotron_vl import NemotronH_Nano_VL_V2_Config
from sglang.srt.managers.schedule_batch import MultimodalProcessorOutput
from sglang.srt.managers.schedule_batch import (
Modality,
MultimodalDataItem,
MultimodalProcessorOutput,
)
from sglang.srt.models.nano_nemotron_vl import NemotronH_Nano_VL_V2
from sglang.srt.models.parakeet import ParakeetExtractor
from sglang.srt.multimodal.audio_from_video import extract_audio_from_video_bytes
from sglang.srt.multimodal.evs import EVSProcessor
from sglang.srt.multimodal.internvl_utils import image_to_pixel_values
from sglang.srt.multimodal.internvl_utils import (
compute_budgeted_image_sizes,
get_video_target_size_and_feature_size,
image_to_pixel_values,
resize_image_to_pixels,
video_to_pixel_values,
)
from sglang.srt.multimodal.processors.base_processor import (
BaseMultimodalProcessor,
MultimodalSpecialTokens,
)
from sglang.srt.utils.common import sample_video_frames
logger = logging.getLogger(__name__)
DEFAULT_NUM_TILES = 12
NUM_VIDEO_TILES = 1
DESIRED_FPS = 2 # TODO: allow desired fps/num frames to be configurable
@@ -63,11 +79,35 @@ class NanoNemotronVLImageProcessor(BaseMultimodalProcessor):
self.img_start_token_id = tokenizer.convert_tokens_to_ids(self.IMG_START_TOKEN)
self.img_end_token_id = tokenizer.convert_tokens_to_ids(self.IMG_END_TOKEN)
# Audio support: initialize Parakeet extractor if sound_config is present
self.audio_extractor: ParakeetExtractor | None = None
self.AUDIO_CONTEXT_TOKEN = getattr(
hf_config, "audio_context_token", "<so_embedding>"
)
self.AUDIO_START_TOKEN = getattr(hf_config, "audio_start_token", "<so_start>")
self.AUDIO_END_TOKEN = getattr(hf_config, "audio_end_token", "<so_end>")
audio_token_str = None
audio_token_id = None
if getattr(hf_config, "sound_config", None) is not None:
self.audio_extractor = ParakeetExtractor(hf_config.sound_config)
audio_token_str = self.AUDIO_CONTEXT_TOKEN
audio_token_id = tokenizer.convert_tokens_to_ids(self.AUDIO_CONTEXT_TOKEN)
self.audio_start_token_id = tokenizer.convert_tokens_to_ids(
self.AUDIO_START_TOKEN
)
self.audio_end_token_id = tokenizer.convert_tokens_to_ids(
self.AUDIO_END_TOKEN
)
self.mm_tokens = MultimodalSpecialTokens(
image_token=self.IMG_CONTEXT_TOKEN,
image_token_id=tokenizer.convert_tokens_to_ids(self.IMG_CONTEXT_TOKEN),
video_token=self.VIDEO_CONTEXT_TOKEN,
video_token_id=tokenizer.convert_tokens_to_ids(self.VIDEO_CONTEXT_TOKEN),
audio_token=audio_token_str,
audio_token_id=audio_token_id,
).build(_image_processor)
# Normalization config (mean/std) and tiling behavior
@@ -75,6 +115,26 @@ class NanoNemotronVLImageProcessor(BaseMultimodalProcessor):
self.norm_std = hf_config.norm_std
self.use_thumbnail = hf_config.use_thumbnail
# Dynamic resolution config
self.dynamic_resolution = getattr(hf_config, "dynamic_resolution", False)
self.min_num_patches = getattr(hf_config, "min_num_patches", 0)
self.max_num_patches = getattr(hf_config, "max_num_patches", 0)
self.patch_size = hf_config.patch_size
self.downsample_ratio = hf_config.downsample_ratio
# Video temporal compression config
self.video_temporal_patch_size = getattr(
hf_config, "video_temporal_patch_size", 1
)
self.video_target_num_patches = getattr(
hf_config, "video_target_num_patches", 0
)
self.video_maintain_aspect_ratio = getattr(
hf_config, "video_maintain_aspect_ratio", True
)
self.max_model_len = getattr(server_args, "context_length", None) or 8192
self.PLACEHOLDER = self.tokenizer.unk_token
assert isinstance(self.PLACEHOLDER, str)
self.PLACEHOLDER_ID = tokenizer.convert_tokens_to_ids(self.PLACEHOLDER)
@@ -95,6 +155,27 @@ class NanoNemotronVLImageProcessor(BaseMultimodalProcessor):
def render_image(self, *, num_tiles: int):
return f"{self.IMG_START_TOKEN}{self.IMG_CONTEXT_TOKEN * self.num_image_token * num_tiles}{self.IMG_END_TOKEN}"
def render_image_dynamic(self, *, num_tokens: int):
return f"{self.IMG_START_TOKEN}{self.IMG_CONTEXT_TOKEN * num_tokens}{self.IMG_END_TOKEN}"
def render_tubelet(
self,
tubelet_index: int,
frame_indices: list[int],
timestamps: list[float],
num_tokens: int,
):
"""Render a tubelet (group of T frames) for temporal compression."""
if len(frame_indices) == 1:
return self.render_frame(
frame_indices[0], timestamp=timestamps[0], num_tokens=num_tokens
)
parts = " and ".join(
f"frame {fi + 1} sampled at {ts:.2f} seconds"
for fi, ts in zip(frame_indices, timestamps)
)
return f"{parts}: {self.PLACEHOLDER}{self.IMG_CONTEXT_TOKEN * num_tokens}{self.IMG_END_TOKEN}"
def render_frame(self, frame_index: int, *, timestamp: float, num_tokens: int):
return f"Frame {frame_index + 1} sampled at {timestamp:.2f} seconds: {self.PLACEHOLDER}{self.IMG_CONTEXT_TOKEN * num_tokens}{self.IMG_END_TOKEN}"
@@ -112,30 +193,106 @@ class NanoNemotronVLImageProcessor(BaseMultimodalProcessor):
timestamps = [i * frame_duration_ms / 1000.0 for i in frames]
return video_array, timestamps
def render_audio(self, *, num_tokens: int):
return (
f"{self.AUDIO_START_TOKEN}"
f"{self.AUDIO_CONTEXT_TOKEN * num_tokens}"
f"{self.AUDIO_END_TOKEN}"
)
async def process_mm_data_async(
self, image_data, input_text, request_obj, **kwargs
self, image_data, audio_data, input_text, request_obj, **kwargs
):
base_output = self.load_mm_data(
prompt=input_text,
image_data=image_data,
video_data=request_obj.video_data,
audio_data=audio_data if self.audio_extractor else None,
multimodal_tokens=self.mm_tokens,
discard_alpha_channel=True,
audio_sample_rate=(
self.audio_extractor.sampling_rate if self.audio_extractor else None
),
)
videos = [self.parse_video(video) for video in base_output.videos]
rows = cols = int(sqrt(self.num_image_token))
create_data_items, tokens_per_frame = self.evs.static_size_data_items(
frames_per_video=[len(frames) for frames, _ in videos],
num_images=len(base_output.images),
rows=rows,
cols=cols,
)
T = self.video_temporal_patch_size
if T > 1:
tubelets_per_video = [math.ceil(len(frames) / T) for frames, _ in videos]
if self.video_target_num_patches > 0 and videos:
frame_h, frame_w = videos[0][0][0].shape[:2]
target_w, target_h, tokens_per_tubelet = (
get_video_target_size_and_feature_size(
frame_w,
frame_h,
self.video_target_num_patches,
self.video_maintain_aspect_ratio,
self.patch_size,
self.downsample_ratio,
)
)
ds = int(1 / self.downsample_ratio)
rows = target_h // self.patch_size // ds
cols = target_w // self.patch_size // ds
else:
tokens_per_tubelet = self.num_image_token
rows = cols = int(sqrt(tokens_per_tubelet))
create_data_items, tokens_per_frame = self.evs.static_size_data_items(
frames_per_video=tubelets_per_video,
num_images=len(base_output.images),
rows=rows,
cols=cols,
)
else:
rows = cols = int(sqrt(self.num_image_token))
create_data_items, tokens_per_frame = self.evs.static_size_data_items(
frames_per_video=[len(frames) for frames, _ in videos],
num_images=len(base_output.images),
rows=rows,
cols=cols,
)
prompt = input_text
image_is_dynamic = False
num_tokens_per_image = []
image_feature = None
if base_output.images:
if base_output.images and self.dynamic_resolution:
image_is_dynamic = True
image_sizes = [(img.width, img.height) for img in base_output.images]
text_only = input_text.replace(self.IMG_CONTEXT_TOKEN, "")
text_tokens = len(
self.tokenizer(text_only, add_special_tokens=False)["input_ids"]
)
total_token_budget = self.max_model_len - text_tokens
budgeted_sizes = compute_budgeted_image_sizes(
image_sizes,
total_token_budget,
self.patch_size,
self.downsample_ratio,
self.min_num_patches,
self.max_num_patches,
)
preprocessed_images = []
for image, (target_w, target_h, n_tokens) in zip(
base_output.images, budgeted_sizes
):
pv = resize_image_to_pixels(
image,
target_w,
target_h,
mean=self.norm_mean,
std=self.norm_std,
)
preprocessed_images.append(pv.to(dtype=torch.bfloat16))
num_tokens_per_image.append(n_tokens)
rendered_images = [
self.render_image_dynamic(num_tokens=nt) for nt in num_tokens_per_image
]
prompt = prompt.replace(self.IMG_CONTEXT_TOKEN, "".join(rendered_images), 1)
image_feature = preprocessed_images
elif base_output.images:
preprocessed_images = [
self.preprocess_image(image) for image in base_output.images
]
@@ -147,35 +304,130 @@ class NanoNemotronVLImageProcessor(BaseMultimodalProcessor):
image_feature = torch.cat(preprocessed_images, dim=0)
video_feature = None
T = self.video_temporal_patch_size
if base_output.videos:
preprocessed_videos = []
for (video_array, timestamps), tpf in zip(
videos, tokens_per_frame, strict=True
):
frames_tensors = [
self.preprocess_image(
Image.fromarray(frame, mode="RGB"),
max_num_tiles=NUM_VIDEO_TILES,
)
for frame in video_array
]
if self.video_target_num_patches > 0:
frames_tensors = []
for frame in video_array:
pv, _ = video_to_pixel_values(
Image.fromarray(frame, mode="RGB"),
patch_size=self.patch_size,
downsample_ratio=self.downsample_ratio,
target_num_patches=self.video_target_num_patches,
maintain_aspect_ratio=self.video_maintain_aspect_ratio,
mean=self.norm_mean,
std=self.norm_std,
)
frames_tensors.append(pv.to(dtype=torch.bfloat16))
else:
frames_tensors = [
self.preprocess_image(
Image.fromarray(frame, mode="RGB"),
max_num_tiles=NUM_VIDEO_TILES,
)
for frame in video_array
]
preprocessed_video = torch.cat(frames_tensors, dim=0)
preprocessed_videos.append(preprocessed_video)
rendered_frames = [
self.render_frame(
i,
timestamp=timestamp,
num_tokens=num_tokens,
if T > 1:
num_frames = len(video_array)
num_tubelets = math.ceil(num_frames / T)
rendered_parts = []
for ti in range(num_tubelets):
start_fi = ti * T
end_fi = min(start_fi + T, num_frames)
fi_list = list(range(start_fi, end_fi))
ts_list = [timestamps[fi] for fi in fi_list]
rendered_parts.append(
self.render_tubelet(
ti, fi_list, ts_list, num_tokens=tpf[ti]
)
)
prompt = prompt.replace(
self.VIDEO_CONTEXT_TOKEN, "\n".join(rendered_parts), 1
)
for i, (timestamp, num_tokens) in enumerate(
zip(timestamps, tpf, strict=True)
else:
rendered_frames = [
self.render_frame(
i,
timestamp=timestamp,
num_tokens=num_tokens,
)
for i, (timestamp, num_tokens) in enumerate(
zip(timestamps, tpf, strict=True)
)
]
prompt = prompt.replace(
self.VIDEO_CONTEXT_TOKEN, "".join(rendered_frames), 1
)
]
prompt = prompt.replace(
self.VIDEO_CONTEXT_TOKEN, "".join(rendered_frames), 1
)
video_feature = torch.cat(preprocessed_videos, dim=0)
# Extract audio from video if requested and no explicit audio provided
use_audio_in_video = getattr(request_obj, "use_audio_in_video", False)
extracted_audios: list[np.ndarray] = []
if (
use_audio_in_video
and base_output.videos
and not base_output.audios
and self.audio_extractor is not None
):
for video_wrapper in base_output.videos:
video_bytes = video_wrapper.source_bytes
if video_bytes is not None:
audio_array = extract_audio_from_video_bytes(
video_bytes,
target_sr=self.audio_extractor.sampling_rate,
)
if audio_array is not None:
extracted_audios.append(audio_array)
all_audios: list[np.ndarray] = (
list(base_output.audios) if base_output.audios else []
)
all_audios.extend(extracted_audios)
# Process audio data through the Parakeet feature extractor
audio_items: list[MultimodalDataItem] = []
if all_audios and self.audio_extractor is not None:
extractor = self.audio_extractor
for audio in all_audios:
num_tokens = extractor.audio_token_count(len(audio))
rendered = self.render_audio(num_tokens=num_tokens)
if self.AUDIO_CONTEXT_TOKEN in prompt:
prompt = prompt.replace(self.AUDIO_CONTEXT_TOKEN, rendered, 1)
else:
prompt = prompt + rendered
extracted = extractor(
all_audios,
sampling_rate=extractor.sampling_rate,
return_tensors="pt",
)
input_features = extracted.input_features
attention_mask = extracted.attention_mask
clip_counts = extracted.audio_num_clips
clip_offset = 0
for audio_idx, num_clips in enumerate(clip_counts):
audio_features = input_features[clip_offset : clip_offset + num_clips]
audio_mask = attention_mask[clip_offset : clip_offset + num_clips]
clip_offset += num_clips
audio_items.append(
MultimodalDataItem(
modality=Modality.AUDIO,
feature=audio_features,
model_specific_data={
"feature_attention_mask": audio_mask,
"audio_num_clips": num_clips,
},
)
)
prompt_ids = self.tokenizer(
prompt, add_special_tokens=False, return_tensors="pt"
)["input_ids"].flatten()
@@ -193,15 +445,46 @@ class NanoNemotronVLImageProcessor(BaseMultimodalProcessor):
# Cleanup:
prompt_ids[prompt_ids == self.PLACEHOLDER_ID] = self.img_start_token_id
# Compute audio offsets
if audio_items:
audio_token_id = self.mm_tokens.audio_token_id
audio_offsets_list = self.get_mm_items_offset(prompt_ids, audio_token_id)
for item, offset in zip(audio_items, audio_offsets_list):
item.offsets = [offset]
prompt_ids_list = prompt_ids.tolist()
items = create_data_items(
image=image_feature,
image_offsets=img_offsets,
video=video_feature,
video_offsets=video_offsets,
input_ids_list=prompt_ids_list,
)
if image_is_dynamic and image_feature is not None:
items = []
for i, (pv, offset) in enumerate(zip(image_feature, img_offsets)):
items.append(
MultimodalDataItem(
modality=Modality.IMAGE,
feature=pv,
offsets=[offset],
model_specific_data={
"num_tokens": num_tokens_per_image[i],
"is_dynamic": True,
},
)
)
if video_feature is not None:
items.append(
MultimodalDataItem(
modality=Modality.VIDEO,
feature=video_feature,
offsets=video_offsets,
)
)
else:
items = create_data_items(
image=image_feature,
image_offsets=img_offsets,
video=video_feature,
video_offsets=video_offsets,
input_ids_list=prompt_ids_list,
)
items.extend(audio_items)
return MultimodalProcessorOutput(
input_ids=prompt_ids_list,
@@ -210,4 +493,7 @@ class NanoNemotronVLImageProcessor(BaseMultimodalProcessor):
im_end_id=self.img_end_token_id,
im_token_id=self.mm_tokens.image_token_id,
video_token_id=self.mm_tokens.image_token_id,
audio_token_id=self.mm_tokens.audio_token_id if audio_items else None,
audio_start_id=(self.audio_start_token_id if audio_items else None),
audio_end_id=(self.audio_end_token_id if audio_items else None),
)
+16
View File
@@ -42,6 +42,8 @@ class VideoDecoderWrapper:
"""source: file path (str) or video bytes.
device: "cpu" or "cuda". GPU decoding only supported with torchcodec.
"""
self._source_bytes = source if isinstance(source, bytes) else None
self._source_path = source if isinstance(source, str) else None
self._tmp_path = None
if _BACKEND == "torchcodec":
kwargs = {"dimension_order": "NHWC"}
@@ -110,6 +112,20 @@ class VideoDecoderWrapper:
arr = self._decoder.get_batch(indices).asnumpy()
return torch.from_numpy(arr).pin_memory()
@property
def source_bytes(self) -> bytes | None:
"""Return raw video bytes if available (needed for audio extraction)."""
if self._source_bytes is not None:
return self._source_bytes
path = self._tmp_path or self._source_path
if path is not None:
import os
if os.path.isfile(path):
with open(path, "rb") as f:
return f.read()
return None
def close(self):
"""Explicitly clean up temporary files."""
if self._tmp_path is not None: