Compare commits
32
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
a25930febb | ||
|
|
5ebcb25e1a | ||
|
|
f683bbf3fa | ||
|
|
0fb05601d3 | ||
|
|
c5f40ef1eb | ||
|
|
f994b3ad93 | ||
|
|
fefe22c53b | ||
|
|
091d58b418 | ||
|
|
5183349eae | ||
|
|
7754387cd8 | ||
|
|
79c3f050a7 | ||
|
|
a8a711b136 | ||
|
|
d1bf08309d | ||
|
|
7e996b8383 | ||
|
|
6b34969257 | ||
|
|
cdf6ab1254 | ||
|
|
8bfe08b826 | ||
|
|
3fe810464a | ||
|
|
5a55a6d10c | ||
|
|
f70ee71b58 | ||
|
|
29cac0d5c8 | ||
|
|
6fe702e4b7 | ||
|
|
2afbcabd21 | ||
|
|
3084207579 | ||
|
|
ad8cd4cde5 | ||
|
|
34916a6adc | ||
|
|
78a50149b6 | ||
|
|
c76b8614fb | ||
|
|
d0787398bc | ||
|
|
9b9b05fb1a | ||
|
|
90c1a6710c | ||
|
|
a468e82fd6 |
@@ -83,16 +83,16 @@ def parallel_attention(q, k, v, img_q_len, img_kv_len, text_mask, mask_strategy=
|
||||
encoder_sequence_length = encoder_query.size(1)
|
||||
|
||||
if mask_strategy[0] is not None:
|
||||
query = torch.cat([tile(query, nccl_info.sp_size), encoder_query], dim=1).transpose(1, 2)
|
||||
key = torch.cat([tile(key, nccl_info.sp_size), encoder_key], dim=1).transpose(1, 2)
|
||||
value = torch.cat([tile(value, nccl_info.sp_size), encoder_value], dim=1).transpose(1, 2)
|
||||
query = torch.cat([tile(query, nccl_info.sp_size), encoder_query], dim=1).transpose(1, 2).contiguous()
|
||||
key = torch.cat([tile(key, nccl_info.sp_size), encoder_key], dim=1).transpose(1, 2).contiguous()
|
||||
value = torch.cat([tile(value, nccl_info.sp_size), encoder_value], dim=1).transpose(1, 2).contiguous()
|
||||
|
||||
head_num = query.size(1)
|
||||
current_rank = nccl_info.rank_within_group
|
||||
start_head = current_rank * head_num
|
||||
windows = [mask_strategy[head_idx + start_head] for head_idx in range(head_num)]
|
||||
|
||||
hidden_states = sliding_tile_attention(query, key, value, windows, text_length).transpose(1, 2)
|
||||
hidden_states = sliding_tile_attention(query, key, value, windows, text_length).transpose(1, 2).contiguous()
|
||||
else:
|
||||
query = torch.cat([query, encoder_query], dim=1)
|
||||
key = torch.cat([key, encoder_key], dim=1)
|
||||
|
||||
Reference in New Issue
Block a user