[CPU] Fix issues when running llama3.2-11B vision model with image tasks (#8666)

Co-authored-by: JieXin Liang <Alcanderian@users.noreply.github.com>
Co-authored-by: Yineng Zhang <me@zhyncs.com>
Co-authored-by: jianan-gu <jianan.gu@intel.com>
This commit is contained in:
blzheng
2026-05-21 13:09:18 +08:00
committed by GitHub
co-authored by JieXin Liang Yineng Zhang jianan-gu
parent 79b937aefb
commit 84ea47eb22
15 changed files with 481 additions and 231 deletions
+48 -17
View File
@@ -21,9 +21,11 @@ class TestExtendAttention(CustomTestCase):
seq_lens: torch.Tensor,
extend_prefix_lens: torch.Tensor,
extend_seq_lens: torch.Tensor,
encoder_lens=None,
scaling=None,
enable_gqa=False,
causal=False,
is_cross_attn=False,
):
assert seq_lens.shape[0] == extend_prefix_lens.shape[0]
@@ -40,7 +42,14 @@ class TestExtendAttention(CustomTestCase):
seq_len_kv = seq_lens[seq_idx]
end_q = start_q + extend_seq_len_q
end_kv = start_kv + seq_len_kv
if encoder_lens is not None:
start_kv = 0 if is_cross_attn else encoder_lens[seq_idx]
end_kv = (
encoder_lens[seq_idx] if is_cross_attn else start_kv + seq_len_kv
)
else:
start_kv = 0
end_kv = start_kv + seq_len_kv
per_req_query = query[:, start_q:end_q, :]
per_req_query_redudant = torch.empty(
@@ -54,7 +63,7 @@ class TestExtendAttention(CustomTestCase):
# get key and value from cache. per_req_tokens contains the kv cache
# index for each token in the sequence.
req_pool_idx = req_pool_indices[seq_idx]
per_req_tokens = req_to_token[req_pool_idx, :seq_len_kv]
per_req_tokens = req_to_token[req_pool_idx, start_kv:end_kv]
per_req_key = k_cache[per_req_tokens].movedim(0, query.dim() - 2)
per_req_value = v_cache[per_req_tokens].movedim(0, query.dim() - 2)
@@ -83,6 +92,7 @@ class TestExtendAttention(CustomTestCase):
D,
DV,
mla=False,
is_cross_attn=False,
*,
b_seq_len_prefix=None,
b_seq_len_extend=None,
@@ -91,32 +101,36 @@ class TestExtendAttention(CustomTestCase):
if b_seq_len_prefix is None:
b_seq_len_prefix = torch.randint(1, N_CTX // 2, (B,), dtype=torch.int32)
if mla:
b_seq_len_prefix.zero_()
else:
b_seq_len_prefix = torch.as_tensor(b_seq_len_prefix, dtype=torch.int32)
encoder_lens = torch.randint(1, N_CTX // 2, (B,), dtype=torch.int64)
if mla:
b_seq_len_prefix.zero_()
encoder_lens.zero_()
if b_seq_len_extend is None:
b_seq_len_extend = torch.randint(1, N_CTX // 2, (B,), dtype=torch.int32)
else:
b_seq_len_extend = torch.as_tensor(b_seq_len_extend, dtype=torch.int32)
b_seq_len = b_seq_len_prefix + b_seq_len_extend
max_len_in_batch = torch.max(b_seq_len, 0)[0].item()
max_len_in_batch = (
torch.max(b_seq_len, 0)[0].item() + torch.max(encoder_lens, 0)[0].item()
)
b_req_idx = torch.arange(B, dtype=torch.int32)
req_to_tokens = torch.empty((B, max_len_in_batch), dtype=torch.int32)
b_start_loc = torch.zeros((B,), dtype=torch.int32)
b_start_loc[1:] = torch.cumsum(b_seq_len[:-1], 0)
b_start_loc[1:] = torch.cumsum(b_seq_len[:-1] + encoder_lens[:-1], 0)
b_start_loc_extend = torch.zeros((B,), dtype=torch.int32)
b_start_loc_extend[1:] = torch.cumsum(b_seq_len_extend[:-1], 0)
for i in range(B):
req_to_tokens[i, : b_seq_len[i]] = torch.arange(
b_start_loc[i], b_start_loc[i] + b_seq_len[i]
req_to_tokens[i, : b_seq_len[i] + encoder_lens[i]] = torch.arange(
b_start_loc[i], b_start_loc[i] + b_seq_len[i] + encoder_lens[i]
)
total_token_num = torch.sum(b_seq_len).item()
total_token_num = torch.sum(b_seq_len).item() + torch.sum(encoder_lens).item()
extend_token_num = torch.sum(b_seq_len_extend).item()
H_BUF = 1 if mla else H_KV
@@ -128,8 +142,10 @@ class TestExtendAttention(CustomTestCase):
q_extend = torch.empty((extend_token_num, H_Q, D), dtype=dtype)
for i in range(B):
extend_start_in_buffer = b_start_loc[i] + b_seq_len_prefix[i]
extend_end_in_buffer = b_start_loc[i] + b_seq_len[i]
extend_start_in_buffer = (
b_start_loc[i] + b_seq_len_prefix[i] + encoder_lens[i]
)
extend_end_in_buffer = b_start_loc[i] + b_seq_len[i] + encoder_lens[i]
extend_start = b_start_loc_extend[i]
extend_end = b_start_loc_extend[i] + b_seq_len_extend[i]
k_extend[extend_start:extend_end] = k_buffer[
@@ -175,7 +191,9 @@ class TestExtendAttention(CustomTestCase):
b_seq_len_extend,
scaling=sm_scale,
enable_gqa=enable_gqa,
causal=True,
causal=not is_cross_attn,
is_cross_attn=is_cross_attn,
encoder_lens=encoder_lens,
)
o_extend = torch.empty((extend_token_num, H_Q, DV), dtype=dtype)
@@ -194,16 +212,29 @@ class TestExtendAttention(CustomTestCase):
max_len_extend,
sm_scale,
logit_cap,
is_cross_attn,
encoder_lens,
)
torch.testing.assert_close(o_ref, o_extend, atol=1e-2, rtol=1e-2)
def test_extend_attention(self):
for is_mla in [True, False]:
self._test_extend_attention_once(1, 123, 1, 1, 128, 96, is_mla)
self._test_extend_attention_once(1, 123, 16, 1, 128, 96, is_mla)
self._test_extend_attention_once(4, 1230, 16, 4, 128, 96, is_mla)
self._test_extend_attention_once(1, 9000, 16, 1, 32, 32, is_mla)
for is_cross_attn in [True, False]:
if is_mla and is_cross_attn:
continue
self._test_extend_attention_once(
1, 123, 1, 1, 128, 96, is_mla, is_cross_attn
)
self._test_extend_attention_once(
1, 123, 16, 1, 128, 96, is_mla, is_cross_attn
)
self._test_extend_attention_once(
4, 1230, 16, 4, 128, 96, is_mla, is_cross_attn
)
self._test_extend_attention_once(
1, 9000, 16, 1, 32, 32, is_mla, is_cross_attn
)
def test_extend_attention_large_seq_causal_mask(self):
self._test_extend_attention_once(