attention_mask

This commit is contained in:
Daniel Han-Chen 2024-02-04 02:19:42 +11:00
commit 711e5c0922
2 changed files with 3 additions and 13 deletions

View file

@ -309,7 +309,6 @@ def LlamaAttention_fast_forward(
V = V.transpose(1, 2)
A = flash_attn_func(Q, K, V, causal = True)
else:
print("0", end = "")
# Grouped query attention
if n_groups != 1:
K = K[:, :, None, :, :].expand(bsz, n_kv_heads, n_groups, kv_seq_len, head_dim)
@ -515,6 +514,7 @@ def LlamaModel_fast_forward(
# if 0 in attention_mask:
# padding_mask = attention_mask
# else:
print(attention_mask)
padding_mask = None
attention_mask = _prepare_4d_causal_attention_mask_for_sdpa(
@ -1208,11 +1208,6 @@ class FastLlamaModel:
@staticmethod
def for_inference(model):
if not hasattr(model, "_original_forward"):
model._original_forward = model.forward
pass
model.forward = torch.inference_mode(model._original_forward)
internal_model = model
internal_model.gradient_checkpointing = False
internal_model.training = False
@ -1227,10 +1222,6 @@ class FastLlamaModel:
@staticmethod
def for_training(model, use_gradient_checkpointing = True):
if hasattr(model, "_original_forward"):
model.forward = model._original_forward
pass
internal_model = model
internal_model.gradient_checkpointing = use_gradient_checkpointing
internal_model.training = True

View file

@ -128,7 +128,7 @@ def MistralAttention_fast_forward(
A = xformers_attention(Q, K, V, attn_bias = causal_mask)
A = A.view(bsz, q_len, n_heads, head_dim)
elif (HAS_FLASH_ATTENTION and attention_mask is None):
elif HAS_FLASH_ATTENTION and attention_mask is None:
Q = Q.transpose(1, 2)
K = K.transpose(1, 2)
V = V.transpose(1, 2)
@ -137,7 +137,6 @@ def MistralAttention_fast_forward(
window = (-1, -1) if (kv_seq_len <= sw) else (sw, sw)
A = flash_attn_func(Q, K, V, causal = True, window_size = window)
else:
print("0", end = "")
# Grouped query attention
# if n_groups != 1:
K = K[:, :, None, :, :].expand(bsz, n_kv_heads, n_groups, kv_seq_len, head_dim)
@ -178,7 +177,7 @@ def MistralForCausalLM_fast_forward(
*args, **kwargs,
) -> Union[Tuple, CausalLMOutputWithPast]:
if causal_mask is None:
if causal_mask is None and past_key_values is None:
bsz, q_len = input_ids.shape
sliding_window = getattr(self.config, "sliding_window", None)
if sliding_window is None or sliding_window == "null" or sliding_window <= 0: