From ad672c897e310b8893e1fd41bb687800e7ace516 Mon Sep 17 00:00:00 2001 From: "Zheng, Beilei" Date: Mon, 22 Sep 2025 01:23:56 -0700 Subject: [PATCH] Revert "Update block mask construction based on the latest pytorch" This reverts commit c067edfee53452cbf2cc247d010ace785ad47379. --- src/transformers/cache_utils.py | 13 +++++++++++-- 1 file changed, 11 insertions(+), 2 deletions(-) diff --git a/src/transformers/cache_utils.py b/src/transformers/cache_utils.py index 89f9e2e04edb..ca9e3ec6e110 100644 --- a/src/transformers/cache_utils.py +++ b/src/transformers/cache_utils.py @@ -2153,12 +2153,21 @@ def __init__( self.value_cache.append(torch.zeros(1, KV_H, max_cached_seq_len, V_D, device=device, dtype=dtype)) self.batch_reserve(self.paged_attentions[i], torch.tensor([max_cache_len for _ in range(batch_size)])) + def generate_causal_offset(offset: torch.Tensor): + def causal_offset_mask(b, h, q_idx, kv_idx): + return (offset + q_idx) >= kv_idx + + return causal_offset_mask + self.batch_size = batch_size self.max_cache_len = max_cache_len self.block_masks = [] - block_mask = create_block_mask(noop_mask, batch_size, 1, 1, max_cache_len, device=device, BLOCK_SIZE=page_size) for i in range(max_cache_len): - self.block_masks.append(self.paged_attentions[0].convert_logical_block_mask(block_mask, kv_len=torch.tensor([i]*batch_size))) + mod = generate_causal_offset( + torch.tensor(i, device=device, dtype=torch.int32) + ) + block_mask = create_block_mask(mod, batch_size, 1, 1, max_cache_len, device=device, BLOCK_SIZE=page_size) + self.block_masks.append(self.paged_attentions[0].convert_logical_block_mask(block_mask)) self.score_mods = [] self.score_mods.append(None) self.score_mods.append(None)