[bugfix]: normalize flash attention metadata masks

This commit is contained in:
Aryan Kumar
2026-07-06 17:15:41 -07:00
parent 38d1d7bdc5
commit 98116070c5
2 changed files with 7 additions and 2 deletions
+5 -1
View File
@@ -338,7 +338,11 @@ class FlashAttentionImpl(AttentionImpl):
)
qkv = torch.stack([query, key, value], dim=2)
attn_mask_padded = F.pad(attn_mask, (qkv.shape[1] - attn_mask.shape[1], 0), value=True)
key_padding_mask = _key_padding_mask_from_attn_mask(attn_mask, attn_mask.shape[-1]).to(device=query.device)
if key_padding_mask.shape[-1] > qkv.shape[1]:
raise ValueError("Invalid key padding mask length for FLASH_ATTN: "
f"expected at most {qkv.shape[1]}, got {key_padding_mask.shape[-1]}")
attn_mask_padded = F.pad(key_padding_mask, (qkv.shape[1] - key_padding_mask.shape[-1], 0), value=True)
output = flash_attn_no_pad(qkv, attn_mask_padded, causal=False, dropout_p=0, softmax_scale=None)
elif self.nvfp4_fa4:
output = self._forward_nvfp4(query, key, value)
+2 -1
View File
@@ -6,7 +6,6 @@ from fastvideo.logger import init_logger
import numpy as np
import torch
from pytorch_msssim import ms_ssim, ssim
logger = init_logger(__name__)
@@ -87,6 +86,8 @@ def compute_video_ssim_torchvision(video1_path, video2_path, use_ms_ssim=True):
video2_path: Path to the second video.
use_ms_ssim: Whether to use Multi-Scale Structural Similarity(MS-SSIM) instead of SSIM.
"""
from pytorch_msssim import ms_ssim, ssim
print(f"Computing SSIM between {video1_path} and {video2_path}...")
if not os.path.exists(video1_path):
raise FileNotFoundError(f"Video1 not found: {video1_path}")