model: qwen3-omni (thinker-only) (#10911)
Co-authored-by: Xinyuan Tong <xinyuantong.cs@gmail.com>
This commit is contained in:
@@ -1156,6 +1156,20 @@ class MRotaryEmbedding(RotaryEmbedding):
|
||||
second_per_grid_ts: Optional[torch.Tensor] = None,
|
||||
**kwargs,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
if model_type == "qwen3_omni_moe":
|
||||
# For qwen3-omni
|
||||
return MRotaryEmbedding.get_rope_index_qwen3_omni(
|
||||
spatial_merge_size,
|
||||
image_token_id,
|
||||
video_token_id,
|
||||
vision_start_token_id,
|
||||
tokens_per_second,
|
||||
input_ids,
|
||||
image_grid_thw,
|
||||
video_grid_thw,
|
||||
second_per_grid_ts,
|
||||
**kwargs,
|
||||
)
|
||||
if (
|
||||
model_type.startswith("qwen3_vl") or model_type.startswith("qwen3_vl_moe")
|
||||
) and video_grid_thw is not None:
|
||||
@@ -1163,6 +1177,7 @@ class MRotaryEmbedding(RotaryEmbedding):
|
||||
video_grid_thw, video_grid_thw[:, 0], dim=0
|
||||
)
|
||||
video_grid_thw[:, 0] = 1
|
||||
|
||||
mrope_position_deltas = []
|
||||
if input_ids is not None and (
|
||||
image_grid_thw is not None or video_grid_thw is not None
|
||||
@@ -1248,7 +1263,11 @@ class MRotaryEmbedding(RotaryEmbedding):
|
||||
|
||||
time_tensor_long = time_tensor.long()
|
||||
t_index = time_tensor_long.flatten()
|
||||
elif model_type in ("qwen2_vl", "qwen3_vl", "qwen3_vl_moe"):
|
||||
elif model_type in (
|
||||
"qwen2_vl",
|
||||
"qwen3_vl",
|
||||
"qwen3_vl_moe",
|
||||
):
|
||||
t_index = (
|
||||
torch.arange(llm_grid_t)
|
||||
.view(-1, 1)
|
||||
@@ -1256,7 +1275,7 @@ class MRotaryEmbedding(RotaryEmbedding):
|
||||
.flatten()
|
||||
)
|
||||
else:
|
||||
raise RuntimeError("Unimplemented")
|
||||
raise RuntimeError(f"Unimplemented model type: {model_type}")
|
||||
h_index = (
|
||||
torch.arange(llm_grid_h)
|
||||
.view(1, -1, 1)
|
||||
@@ -1306,6 +1325,304 @@ class MRotaryEmbedding(RotaryEmbedding):
|
||||
mrope_position_deltas = max_position_ids + 1 - s
|
||||
return position_ids, mrope_position_deltas
|
||||
|
||||
@staticmethod
|
||||
def get_rope_index_qwen3_omni(
|
||||
spatial_merge_size: int,
|
||||
image_token_id: int,
|
||||
video_token_id: int,
|
||||
vision_start_token_id: int,
|
||||
tokens_per_second: Optional[int] = None,
|
||||
input_ids: Optional[torch.LongTensor] = None,
|
||||
image_grid_thw: Optional[torch.LongTensor] = None,
|
||||
video_grid_thw: Optional[torch.LongTensor] = None,
|
||||
second_per_grid_ts: Optional[torch.Tensor] = None,
|
||||
**kwargs,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
# For qwen3-omni
|
||||
audio_token_id = kwargs["audio_token_id"]
|
||||
audio_start_token_id = kwargs["audio_start_token_id"]
|
||||
position_id_per_seconds = kwargs["position_id_per_seconds"]
|
||||
use_audio_in_video = kwargs.get("use_audio_in_video", False)
|
||||
audio_seqlens = kwargs.get("audio_seqlens", None)
|
||||
second_per_grids = second_per_grid_ts
|
||||
|
||||
mrope_position_deltas = []
|
||||
if input_ids is not None and (
|
||||
image_grid_thw is not None or video_grid_thw is not None
|
||||
):
|
||||
total_input_ids = input_ids
|
||||
position_ids = torch.zeros(
|
||||
3,
|
||||
input_ids.shape[0],
|
||||
input_ids.shape[1],
|
||||
dtype=torch.float,
|
||||
device=input_ids.device,
|
||||
)
|
||||
image_idx, video_idx, audio_idx = 0, 0, 0
|
||||
for i, current_input_ids in enumerate(total_input_ids):
|
||||
image_nums, video_nums, audio_nums = 0, 0, 0
|
||||
vision_start_indices = torch.argwhere(
|
||||
current_input_ids == vision_start_token_id
|
||||
).squeeze(1)
|
||||
if vision_start_indices.numel() > 0:
|
||||
vision_tokens = current_input_ids[vision_start_indices + 1]
|
||||
image_nums = (vision_tokens == image_token_id).sum()
|
||||
video_nums = (
|
||||
(vision_tokens == audio_start_token_id).sum()
|
||||
if use_audio_in_video
|
||||
else (vision_tokens == video_token_id).sum()
|
||||
)
|
||||
audio_nums = torch.sum(current_input_ids == audio_start_token_id)
|
||||
input_tokens = current_input_ids.tolist()
|
||||
llm_pos_ids_list: list = []
|
||||
st = 0
|
||||
remain_images, remain_videos, remain_audios = (
|
||||
image_nums,
|
||||
video_nums,
|
||||
audio_nums,
|
||||
)
|
||||
multimodal_nums = (
|
||||
image_nums + audio_nums
|
||||
if use_audio_in_video
|
||||
else image_nums + video_nums + audio_nums
|
||||
)
|
||||
for _ in range(multimodal_nums):
|
||||
st_idx = (
|
||||
llm_pos_ids_list[-1].max() + 1
|
||||
if len(llm_pos_ids_list) > 0
|
||||
else 0
|
||||
)
|
||||
ed_vision_start = (
|
||||
input_tokens.index(vision_start_token_id, st)
|
||||
if (
|
||||
(
|
||||
image_token_id in input_tokens
|
||||
or video_token_id in input_tokens
|
||||
)
|
||||
and (remain_videos > 0 or remain_images > 0)
|
||||
)
|
||||
else len(input_tokens) + 1
|
||||
)
|
||||
ed_audio_start = (
|
||||
input_tokens.index(audio_start_token_id, st)
|
||||
if (audio_token_id in input_tokens and remain_audios > 0)
|
||||
else len(input_tokens) + 1
|
||||
)
|
||||
min_ed = min(ed_vision_start, ed_audio_start)
|
||||
|
||||
text_len = min_ed - st
|
||||
if text_len != 0:
|
||||
llm_pos_ids_list.append(
|
||||
torch.arange(text_len).view(1, -1).expand(3, -1) + st_idx
|
||||
)
|
||||
st_idx += text_len
|
||||
# Audio in Video
|
||||
if (
|
||||
min_ed == ed_vision_start
|
||||
and ed_vision_start + 1 == ed_audio_start
|
||||
):
|
||||
bos_len, eos_len = 2, 2
|
||||
else:
|
||||
bos_len, eos_len = 1, 1
|
||||
llm_pos_ids_list.append(
|
||||
torch.arange(bos_len).view(1, -1).expand(3, -1) + st_idx
|
||||
)
|
||||
st_idx += bos_len
|
||||
# Audio Only
|
||||
if min_ed == ed_audio_start:
|
||||
audio_len = MRotaryEmbedding._get_feat_extract_output_lengths(
|
||||
audio_seqlens[audio_idx]
|
||||
)
|
||||
llm_pos_ids = (
|
||||
torch.arange(audio_len).view(1, -1).expand(3, -1) + st_idx
|
||||
)
|
||||
llm_pos_ids_list.append(llm_pos_ids)
|
||||
|
||||
st += int(text_len + bos_len + audio_len + eos_len)
|
||||
audio_idx += 1
|
||||
remain_audios -= 1
|
||||
|
||||
# Image Only
|
||||
elif (
|
||||
min_ed == ed_vision_start
|
||||
and current_input_ids[ed_vision_start + 1] == image_token_id
|
||||
):
|
||||
grid_t = image_grid_thw[image_idx][0]
|
||||
grid_hs = image_grid_thw[:, 1]
|
||||
grid_ws = image_grid_thw[:, 2]
|
||||
t_index = (
|
||||
torch.arange(grid_t) * 1 * position_id_per_seconds
|
||||
).float()
|
||||
llm_pos_ids = MRotaryEmbedding._get_llm_pos_ids_for_vision(
|
||||
st_idx,
|
||||
image_idx,
|
||||
spatial_merge_size,
|
||||
t_index,
|
||||
grid_hs,
|
||||
grid_ws,
|
||||
input_ids.device,
|
||||
)
|
||||
image_len = image_grid_thw[image_idx].prod() // (
|
||||
spatial_merge_size**2
|
||||
)
|
||||
llm_pos_ids_list.append(llm_pos_ids)
|
||||
|
||||
st += int(text_len + bos_len + image_len + eos_len)
|
||||
image_idx += 1
|
||||
remain_images -= 1
|
||||
|
||||
# Video Only
|
||||
elif (
|
||||
min_ed == ed_vision_start
|
||||
and current_input_ids[ed_vision_start + 1] == video_token_id
|
||||
):
|
||||
grid_t = video_grid_thw[video_idx][0]
|
||||
grid_hs = video_grid_thw[:, 1]
|
||||
grid_ws = video_grid_thw[:, 2]
|
||||
t_index = (
|
||||
torch.arange(grid_t)
|
||||
* second_per_grids[video_idx].cpu().float()
|
||||
* position_id_per_seconds
|
||||
).float()
|
||||
llm_pos_ids = MRotaryEmbedding._get_llm_pos_ids_for_vision(
|
||||
st_idx,
|
||||
video_idx,
|
||||
spatial_merge_size,
|
||||
t_index,
|
||||
grid_hs,
|
||||
grid_ws,
|
||||
input_ids.device,
|
||||
)
|
||||
video_len = video_grid_thw[video_idx].prod() // (
|
||||
spatial_merge_size**2
|
||||
)
|
||||
llm_pos_ids_list.append(llm_pos_ids)
|
||||
|
||||
st += int(text_len + bos_len + video_len + eos_len)
|
||||
video_idx += 1
|
||||
remain_videos -= 1
|
||||
|
||||
# Audio in Video
|
||||
elif (
|
||||
min_ed == ed_vision_start
|
||||
and ed_vision_start + 1 == ed_audio_start
|
||||
):
|
||||
audio_len = MRotaryEmbedding._get_feat_extract_output_lengths(
|
||||
audio_seqlens[audio_idx]
|
||||
)
|
||||
audio_llm_pos_ids = (
|
||||
torch.arange(audio_len).view(1, -1).expand(3, -1) + st_idx
|
||||
)
|
||||
grid_t = video_grid_thw[video_idx][0]
|
||||
grid_hs = video_grid_thw[:, 1]
|
||||
grid_ws = video_grid_thw[:, 2]
|
||||
|
||||
t_index = (
|
||||
torch.arange(grid_t)
|
||||
* second_per_grids[video_idx].cpu().float()
|
||||
* position_id_per_seconds
|
||||
).float()
|
||||
video_llm_pos_ids = (
|
||||
MRotaryEmbedding._get_llm_pos_ids_for_vision(
|
||||
st_idx,
|
||||
video_idx,
|
||||
spatial_merge_size,
|
||||
t_index,
|
||||
grid_hs,
|
||||
grid_ws,
|
||||
input_ids.device,
|
||||
)
|
||||
)
|
||||
video_data_index, audio_data_index = 0, 0
|
||||
while (
|
||||
video_data_index < video_llm_pos_ids.shape[-1]
|
||||
and audio_data_index < audio_llm_pos_ids.shape[-1]
|
||||
):
|
||||
if (
|
||||
video_llm_pos_ids[0][video_data_index]
|
||||
<= audio_llm_pos_ids[0][audio_data_index]
|
||||
):
|
||||
llm_pos_ids_list.append(
|
||||
video_llm_pos_ids[
|
||||
:, video_data_index : video_data_index + 1
|
||||
]
|
||||
)
|
||||
video_data_index += 1
|
||||
else:
|
||||
llm_pos_ids_list.append(
|
||||
audio_llm_pos_ids[
|
||||
:, audio_data_index : audio_data_index + 1
|
||||
]
|
||||
)
|
||||
audio_data_index += 1
|
||||
if video_data_index < video_llm_pos_ids.shape[-1]:
|
||||
llm_pos_ids_list.append(
|
||||
video_llm_pos_ids[
|
||||
:, video_data_index : video_llm_pos_ids.shape[-1]
|
||||
]
|
||||
)
|
||||
if audio_data_index < audio_llm_pos_ids.shape[-1]:
|
||||
llm_pos_ids_list.append(
|
||||
audio_llm_pos_ids[
|
||||
:, audio_data_index : audio_llm_pos_ids.shape[-1]
|
||||
]
|
||||
)
|
||||
video_len = video_grid_thw[video_idx].prod() // (
|
||||
spatial_merge_size**2
|
||||
)
|
||||
|
||||
st += int(text_len + bos_len + audio_len + video_len + eos_len)
|
||||
|
||||
audio_idx += 1
|
||||
video_idx += 1
|
||||
remain_videos -= 1
|
||||
remain_audios -= 1
|
||||
st_idx = (
|
||||
llm_pos_ids_list[-1].max() + 1
|
||||
if len(llm_pos_ids_list) > 0
|
||||
else 0
|
||||
)
|
||||
llm_pos_ids_list.append(
|
||||
torch.arange(eos_len).view(1, -1).expand(3, -1) + st_idx
|
||||
)
|
||||
|
||||
if st < len(input_tokens):
|
||||
st_idx = (
|
||||
llm_pos_ids_list[-1].max() + 1
|
||||
if len(llm_pos_ids_list) > 0
|
||||
else 0
|
||||
)
|
||||
text_len = len(input_tokens) - st
|
||||
llm_pos_ids_list.append(
|
||||
torch.arange(text_len).view(1, -1).expand(3, -1) + st_idx
|
||||
)
|
||||
|
||||
llm_positions = torch.cat(
|
||||
[item.float() for item in llm_pos_ids_list], dim=1
|
||||
).reshape(3, -1)
|
||||
|
||||
position_ids[..., i, :] = llm_positions.to(position_ids.device)
|
||||
mrope_position_deltas.append(
|
||||
llm_positions.max() + 1 - len(current_input_ids)
|
||||
)
|
||||
mrope_position_deltas = torch.tensor(
|
||||
mrope_position_deltas, device=input_ids.device
|
||||
).unsqueeze(1)
|
||||
|
||||
return position_ids, mrope_position_deltas
|
||||
else:
|
||||
s = input_ids.shape[1]
|
||||
position_ids = torch.arange(s)
|
||||
position_ids = (
|
||||
position_ids.unsqueeze(0).expand(3, -1, -1).to(input_ids.device)
|
||||
)
|
||||
max_position_ids = position_ids.max(0, keepdim=False)[0].max(
|
||||
-1, keepdim=True
|
||||
)[0]
|
||||
mrope_position_deltas = max_position_ids + 1 - s
|
||||
|
||||
return position_ids, mrope_position_deltas
|
||||
|
||||
# Adapted from https://github.com/vllm-project/vllm/blob/3779eb8c81449b924a23457fc77e45a0e6171178/vllm/model_executor/layers/rotary_embedding.py#L1120
|
||||
@staticmethod
|
||||
def get_rope_index_glm4v(
|
||||
@@ -1504,6 +1821,44 @@ class MRotaryEmbedding(RotaryEmbedding):
|
||||
|
||||
return position_ids, mrope_position_deltas
|
||||
|
||||
# For qwen3-omni
|
||||
@staticmethod
|
||||
def _get_feat_extract_output_lengths(input_lengths):
|
||||
"""
|
||||
Computes the output length of the convolutional layers and the output length of the audio encoder
|
||||
"""
|
||||
input_lengths_leave = input_lengths % 100
|
||||
feat_lengths = (input_lengths_leave - 1) // 2 + 1
|
||||
output_lengths = (
|
||||
((feat_lengths - 1) // 2 + 1 - 1) // 2 + 1 + (input_lengths // 100) * 13
|
||||
)
|
||||
return output_lengths
|
||||
|
||||
# For qwen3-omni
|
||||
@staticmethod
|
||||
def _get_llm_pos_ids_for_vision(
|
||||
st_idx, vision_idx, spatial_merge_size, t_index, grid_hs, grid_ws, device
|
||||
):
|
||||
grid_h = grid_hs[vision_idx] // spatial_merge_size
|
||||
grid_w = grid_ws[vision_idx] // spatial_merge_size
|
||||
|
||||
h_index = (
|
||||
torch.arange(grid_h, device=device)
|
||||
.view(1, -1, 1)
|
||||
.expand(len(t_index), -1, grid_w)
|
||||
.flatten()
|
||||
)
|
||||
w_index = (
|
||||
torch.arange(grid_w, device=device)
|
||||
.view(1, 1, -1)
|
||||
.expand(len(t_index), grid_h, -1)
|
||||
.flatten()
|
||||
)
|
||||
t_index = t_index.view(-1, 1).expand(-1, grid_h * grid_w).flatten()
|
||||
|
||||
llm_pos_ids = torch.stack([t_index, h_index, w_index], dim=0) + st_idx
|
||||
return llm_pos_ids
|
||||
|
||||
|
||||
class DualChunkRotaryEmbedding(CustomOp):
|
||||
"""Rotary positional embedding for Dual Chunk Attention."""
|
||||
|
||||
Reference in New Issue
Block a user