Fix flash_attn fallback

This commit is contained in:
kijai
2024-06-17 19:47:34 +03:00
parent 2c5cc1e27f
commit 15cdc3eab3
+4 -2
View File
@@ -26,7 +26,6 @@ import torch.nn.functional as F
from .components import RMSNorm
import comfy.model_management
import comfy.ops
ops = comfy.ops.manual_cast
@@ -365,7 +364,10 @@ class Attention(nn.Module):
# end var_len_flash_attn
else:
raise Exception("Flash attention (flash_attn) is not available and currently required for Lumina-next -models.")
n_rep = self.n_local_heads // self.n_local_kv_heads
if n_rep >= 1:
xk = xk.unsqueeze(3).repeat(1, 1, 1, n_rep, 1).flatten(2, 3)
xv = xv.unsqueeze(3).repeat(1, 1, 1, n_rep, 1).flatten(2, 3)
output = (
F.scaled_dot_product_attention(
xq.permute(0, 2, 1, 3),