[bugfix]: normalize flash attention metadata masks
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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}")
|
||||
|
||||
Reference in New Issue
Block a user