model: qwen3-omni (thinker-only) (#10911)

Co-authored-by: Xinyuan Tong <xinyuantong.cs@gmail.com>
This commit is contained in:
Mick
2025-10-16 13:20:38 -07:00
committed by GitHub
co-authored by Xinyuan Tong
parent 85ebeecf06
commit 86b04d25b3
16 changed files with 1954 additions and 335 deletions
+357 -2
View File
@@ -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."""