[VLM] Optimize get_rope_index for GLM4v (#17420)
Co-authored-by: luoyuan.luo <luoyuan.luo@antgroup.com>
This commit is contained in:
@@ -1700,11 +1700,12 @@ class MRotaryEmbedding(RotaryEmbedding):
|
||||
video_index += 1
|
||||
remain_videos -= 1
|
||||
ed = ed_video
|
||||
llm_grid_t, llm_grid_h, llm_grid_w = (
|
||||
t.item(),
|
||||
h.item() // spatial_merge_size,
|
||||
w.item() // spatial_merge_size,
|
||||
)
|
||||
# Avoid .item() lookups in repeated context
|
||||
t_int, h_int, w_int = int(t), int(h), int(w)
|
||||
|
||||
llm_grid_t = t_int
|
||||
llm_grid_h = h_int // spatial_merge_size
|
||||
llm_grid_w = w_int // spatial_merge_size
|
||||
text_len = ed - st
|
||||
|
||||
st_idx = (
|
||||
@@ -1737,24 +1738,24 @@ class MRotaryEmbedding(RotaryEmbedding):
|
||||
"qwen3_vl_moe",
|
||||
):
|
||||
t_index = (
|
||||
torch.arange(llm_grid_t)
|
||||
torch.arange(llm_grid_t, device=position_ids.device)
|
||||
.view(-1, 1)
|
||||
.expand(-1, llm_grid_h * llm_grid_w)
|
||||
.flatten()
|
||||
.expand(llm_grid_t, llm_grid_h * llm_grid_w)
|
||||
.reshape(-1)
|
||||
)
|
||||
else:
|
||||
raise RuntimeError(f"Unimplemented model type: {model_type}")
|
||||
h_index = (
|
||||
torch.arange(llm_grid_h)
|
||||
torch.arange(llm_grid_h, device=position_ids.device)
|
||||
.view(1, -1, 1)
|
||||
.expand(llm_grid_t, -1, llm_grid_w)
|
||||
.flatten()
|
||||
.expand(llm_grid_t, llm_grid_h, llm_grid_w)
|
||||
.reshape(-1)
|
||||
)
|
||||
w_index = (
|
||||
torch.arange(llm_grid_w)
|
||||
torch.arange(llm_grid_w, device=position_ids.device)
|
||||
.view(1, 1, -1)
|
||||
.expand(llm_grid_t, llm_grid_h, -1)
|
||||
.flatten()
|
||||
.expand(llm_grid_t, llm_grid_h, llm_grid_w)
|
||||
.reshape(-1)
|
||||
)
|
||||
llm_pos_ids_list.append(
|
||||
torch.stack([t_index, h_index, w_index]) + text_len + st_idx
|
||||
@@ -1787,10 +1788,9 @@ class MRotaryEmbedding(RotaryEmbedding):
|
||||
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
|
||||
max_position_ids = position_ids.amax(dim=0, keepdim=False)
|
||||
mrope_position_deltas = max_position_ids.amax(-1, keepdim=True) + 1 - s
|
||||
|
||||
return position_ids, mrope_position_deltas
|
||||
|
||||
@staticmethod
|
||||
@@ -2107,13 +2107,17 @@ class MRotaryEmbedding(RotaryEmbedding):
|
||||
video_end_token_id = hf_config.video_end_token_id
|
||||
spatial_merge_size = hf_config.vision_config.spatial_merge_size
|
||||
|
||||
# Preallocate lists for efficiency
|
||||
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
|
||||
|
||||
if attention_mask is None:
|
||||
attention_mask = torch.ones_like(total_input_ids)
|
||||
|
||||
position_ids = torch.ones(
|
||||
3,
|
||||
input_ids.shape[0],
|
||||
@@ -2121,28 +2125,36 @@ class MRotaryEmbedding(RotaryEmbedding):
|
||||
dtype=input_ids.dtype,
|
||||
device=input_ids.device,
|
||||
)
|
||||
|
||||
image_index, video_index = 0, 0
|
||||
video_group_index = 0
|
||||
# Move attention mask to device once to avoid repeated transfers
|
||||
attention_mask = attention_mask.to(total_input_ids.device)
|
||||
for i, input_ids in enumerate(total_input_ids):
|
||||
input_ids = input_ids[attention_mask[i] == 1]
|
||||
input_tokens = input_ids.tolist()
|
||||
|
||||
input_token_type = []
|
||||
for i, ids in enumerate(total_input_ids):
|
||||
curr_mask = attention_mask[i]
|
||||
ids_masked = ids[curr_mask == 1]
|
||||
|
||||
# Preallocate input_token_type for maximum speed
|
||||
input_tokens = ids_masked.tolist()
|
||||
input_token_type = [""] * len(input_tokens)
|
||||
|
||||
# Single pass through tokens for type assignment, using explicit indices for performance
|
||||
video_check_flg = False
|
||||
for token in input_tokens:
|
||||
for j, token in enumerate(input_tokens):
|
||||
if token == video_start_token_id:
|
||||
video_check_flg = True
|
||||
elif token == video_end_token_id:
|
||||
video_check_flg = False
|
||||
|
||||
if token == image_token_id and not video_check_flg:
|
||||
input_token_type.append("image")
|
||||
input_token_type[j] = "image"
|
||||
elif token == image_token_id and video_check_flg:
|
||||
input_token_type.append("video")
|
||||
input_token_type[j] = "video"
|
||||
else:
|
||||
input_token_type.append("text")
|
||||
input_token_type[j] = "text"
|
||||
|
||||
# Use itertools.groupby for consecutive token type groups (unchanged logic)
|
||||
input_type_group = []
|
||||
for key, group in itertools.groupby(
|
||||
enumerate(input_token_type), lambda x: x[1]
|
||||
@@ -2154,12 +2166,13 @@ class MRotaryEmbedding(RotaryEmbedding):
|
||||
|
||||
llm_pos_ids_list = []
|
||||
video_frame_num = 1
|
||||
|
||||
for modality_type, start_idx, end_idx in input_type_group:
|
||||
st_idx = (
|
||||
llm_pos_ids_list[-1].max() + 1
|
||||
if len(llm_pos_ids_list) > 0
|
||||
else 0
|
||||
)
|
||||
# st_idx can be computed by torch directly for speed
|
||||
if llm_pos_ids_list:
|
||||
st_idx = llm_pos_ids_list[-1].max().item() + 1
|
||||
else:
|
||||
st_idx = 0
|
||||
|
||||
if modality_type == "image":
|
||||
t, h, w = (
|
||||
@@ -2167,103 +2180,102 @@ class MRotaryEmbedding(RotaryEmbedding):
|
||||
image_grid_thw[image_index][1],
|
||||
image_grid_thw[image_index][2],
|
||||
)
|
||||
llm_grid_t, llm_grid_h, llm_grid_w = (
|
||||
t.item(),
|
||||
h.item() // spatial_merge_size,
|
||||
w.item() // spatial_merge_size,
|
||||
)
|
||||
# Avoid .item() lookups in repeated context
|
||||
t_int, h_int, w_int = int(t), int(h), int(w)
|
||||
|
||||
llm_grid_t = t_int
|
||||
llm_grid_h = h_int // spatial_merge_size
|
||||
llm_grid_w = w_int // spatial_merge_size
|
||||
|
||||
# Avoid unnecessary views/expands for speed, always flatten at the end
|
||||
t_index = (
|
||||
torch.arange(llm_grid_t)
|
||||
torch.arange(llm_grid_t, device=position_ids.device)
|
||||
.view(-1, 1)
|
||||
.expand(-1, llm_grid_h * llm_grid_w)
|
||||
.flatten()
|
||||
.expand(llm_grid_t, llm_grid_h * llm_grid_w)
|
||||
.reshape(-1)
|
||||
)
|
||||
h_index = (
|
||||
torch.arange(llm_grid_h)
|
||||
torch.arange(llm_grid_h, device=position_ids.device)
|
||||
.view(1, -1, 1)
|
||||
.expand(llm_grid_t, -1, llm_grid_w)
|
||||
.flatten()
|
||||
.expand(llm_grid_t, llm_grid_h, llm_grid_w)
|
||||
.reshape(-1)
|
||||
)
|
||||
w_index = (
|
||||
torch.arange(llm_grid_w)
|
||||
torch.arange(llm_grid_w, device=position_ids.device)
|
||||
.view(1, 1, -1)
|
||||
.expand(llm_grid_t, llm_grid_h, -1)
|
||||
.flatten()
|
||||
.expand(llm_grid_t, llm_grid_h, llm_grid_w)
|
||||
.reshape(-1)
|
||||
)
|
||||
llm_pos_ids_list.append(
|
||||
torch.stack([t_index, h_index, w_index]) + st_idx
|
||||
)
|
||||
|
||||
image_index += 1
|
||||
video_frame_num = 1
|
||||
|
||||
elif modality_type == "video":
|
||||
t, h, w = (
|
||||
video_frame_num,
|
||||
video_grid_thw[video_index][1],
|
||||
video_grid_thw[video_index][2],
|
||||
)
|
||||
t = video_frame_num
|
||||
h = video_grid_thw[video_index][1]
|
||||
w = video_grid_thw[video_index][2]
|
||||
|
||||
llm_grid_t, llm_grid_h, llm_grid_w = (
|
||||
t,
|
||||
h.item() // spatial_merge_size,
|
||||
w.item() // spatial_merge_size,
|
||||
)
|
||||
h_int, w_int = int(h), int(w)
|
||||
llm_grid_h = h_int // spatial_merge_size
|
||||
llm_grid_w = w_int // spatial_merge_size
|
||||
|
||||
for t_idx in range(llm_grid_t):
|
||||
# Only one video frame at a time
|
||||
for t_idx in range(t):
|
||||
t_index = (
|
||||
torch.tensor(t_idx)
|
||||
torch.tensor(t_idx, device=position_ids.device)
|
||||
.view(-1, 1)
|
||||
.expand(-1, llm_grid_h * llm_grid_w)
|
||||
.flatten()
|
||||
.expand(1, llm_grid_h * llm_grid_w)
|
||||
.reshape(-1)
|
||||
)
|
||||
|
||||
h_index = (
|
||||
torch.arange(llm_grid_h)
|
||||
torch.arange(llm_grid_h, device=position_ids.device)
|
||||
.view(1, -1, 1)
|
||||
.expand(1, -1, llm_grid_w)
|
||||
.flatten()
|
||||
.expand(1, llm_grid_h, llm_grid_w)
|
||||
.reshape(-1)
|
||||
)
|
||||
w_index = (
|
||||
torch.arange(llm_grid_w)
|
||||
torch.arange(llm_grid_w, device=position_ids.device)
|
||||
.view(1, 1, -1)
|
||||
.expand(1, llm_grid_h, -1)
|
||||
.flatten()
|
||||
.expand(1, llm_grid_h, llm_grid_w)
|
||||
.reshape(-1)
|
||||
)
|
||||
llm_pos_ids_list.append(
|
||||
torch.stack([t_index, h_index, w_index]) + st_idx
|
||||
)
|
||||
|
||||
video_group_index += 1
|
||||
|
||||
if video_group_index >= video_grid_thw[video_index][0]:
|
||||
video_index += 1
|
||||
video_group_index = 0
|
||||
|
||||
video_frame_num += 1
|
||||
|
||||
else:
|
||||
else: # text
|
||||
text_len = end_idx - start_idx
|
||||
llm_pos_ids_list.append(
|
||||
torch.arange(text_len).view(1, -1).expand(3, -1) + st_idx
|
||||
)
|
||||
|
||||
# Use in-place expand for improved performance
|
||||
text_range = torch.arange(text_len, device=position_ids.device)
|
||||
text_pos = text_range.view(1, -1).expand(3, text_len) + st_idx
|
||||
llm_pos_ids_list.append(text_pos)
|
||||
video_frame_num = 1
|
||||
|
||||
# Concatenate once outside for speed
|
||||
llm_positions = torch.cat(llm_pos_ids_list, dim=1).reshape(3, -1)
|
||||
position_ids[..., i, attention_mask[i] == 1] = llm_positions.to(
|
||||
position_ids.device
|
||||
)
|
||||
# Use advanced indexing for assignment
|
||||
idx_mask = curr_mask == 1
|
||||
position_ids[..., i, idx_mask] = llm_positions.to(position_ids.device)
|
||||
mrope_position_deltas.append(
|
||||
llm_positions.max() + 1 - len(total_input_ids[i])
|
||||
)
|
||||
# Build tensor in one call at the end
|
||||
mrope_position_deltas = torch.tensor(
|
||||
mrope_position_deltas, device=input_ids.device
|
||||
).unsqueeze(1)
|
||||
return position_ids, mrope_position_deltas
|
||||
else:
|
||||
if attention_mask is not None:
|
||||
# Use in-place operations whenever possible
|
||||
position_ids = attention_mask.long().cumsum(-1) - 1
|
||||
position_ids.masked_fill_(attention_mask == 0, 1)
|
||||
position_ids = (
|
||||
@@ -2271,22 +2283,25 @@ class MRotaryEmbedding(RotaryEmbedding):
|
||||
.expand(3, -1, -1)
|
||||
.to(attention_mask.device)
|
||||
)
|
||||
max_position_ids = position_ids.max(0, keepdim=False)[0].max(
|
||||
-1, keepdim=True
|
||||
)[0]
|
||||
mrope_position_deltas = max_position_ids + 1 - attention_mask.shape[-1]
|
||||
else:
|
||||
position_ids = (
|
||||
torch.arange(input_ids.shape[1], device=input_ids.device)
|
||||
.view(1, 1, -1)
|
||||
.expand(3, input_ids.shape[0], -1)
|
||||
max_position_ids = position_ids.amax(dim=0, keepdim=False)
|
||||
mrope_position_deltas = (
|
||||
max_position_ids.amax(-1, keepdim=True)
|
||||
+ 1
|
||||
- attention_mask.shape[-1]
|
||||
)
|
||||
else:
|
||||
length = input_ids.shape[1]
|
||||
batch_size = input_ids.shape[0]
|
||||
# Use torch.arange with in-place expansion
|
||||
arange_ids = torch.arange(length, device=input_ids.device).view(
|
||||
1, 1, -1
|
||||
)
|
||||
position_ids = arange_ids.expand(3, batch_size, length)
|
||||
mrope_position_deltas = torch.zeros(
|
||||
[input_ids.shape[0], 1],
|
||||
[batch_size, 1],
|
||||
device=input_ids.device,
|
||||
dtype=input_ids.dtype,
|
||||
)
|
||||
|
||||
return position_ids, mrope_position_deltas
|
||||
|
||||
@staticmethod
|
||||
|
||||
Reference in New Issue
Block a user