Compare commits

...
Author SHA1 Message Date
William Lin 193ce7efa1 Revert "[chore] release v0.1.7 (#955)"
This reverts commit 9cd6a86b95.
2025-12-27 05:05:42 -06:00
11 changed files with 130 additions and 717 deletions
+39 -15
View File
@@ -22,7 +22,7 @@ steps:
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 20m .buildkite/scripts/pr_test.sh"
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Encoder Tests"
env:
- TEST_TYPE=encoder
@@ -35,7 +35,7 @@ steps:
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 20m .buildkite/scripts/pr_test.sh"
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "VAE Tests"
env:
- TEST_TYPE=vae
@@ -157,9 +157,33 @@ steps:
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Kernel Tests"
label: "Precision Tests STA"
env:
- TEST_TYPE=kernel_tests
- TEST_TYPE=precision_sta
agents:
queue: "default"
- path:
- "fastvideo-kernel/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Precision Tests VSA"
env:
- TEST_TYPE=precision_vsa
agents:
queue: "default"
- path:
- "fastvideo-kernel/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Precision Tests VMoBA"
env:
- TEST_TYPE=precision_vmoba
agents:
queue: "default"
- path:
- "fastvideo-kernel/**"
- "fastvideo/attention/backends/vmoba.py"
@@ -183,14 +207,14 @@ steps:
- TEST_TYPE=unit_test
agents:
queue: "default"
# - path:
# - "scripts/lora_extraction/**"
# - "pyproject.toml"
# - "docker/Dockerfile.python3.12"
# config:
# command: "timeout 90m .buildkite/scripts/pr_test.sh"
# label: "LoRA Extraction Tests"
# env:
# - TEST_TYPE=lora_extraction
# agents:
# queue: "default"
- path:
- "scripts/lora_extraction/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 90m .buildkite/scripts/pr_test.sh"
label: "LoRA Extraction Tests"
env:
- TEST_TYPE=lora_extraction
agents:
queue: "default"
+11 -3
View File
@@ -93,9 +93,13 @@ case "$TEST_TYPE" in
log "Running inference STA tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_inference_tests_STA"
;;
"kernel_tests")
log "Running kernel tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_kernel_tests"
"precision_sta")
log "Running precision STA tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_precision_tests_STA"
;;
"precision_vsa")
log "Running precision VSA tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_precision_tests_VSA"
;;
"inference_lora")
log "Running LoRA tests..."
@@ -114,6 +118,10 @@ case "$TEST_TYPE" in
log "Running V-MoBA inference tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_inference_tests_vmoba"
;;
"precision_vmoba")
log "Running V-MoBA precision tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_precision_tests_vmoba"
;;
"unit_test")
log "Running unit tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_unit_test"
@@ -0,0 +1,63 @@
import torch
import sys
import os
from tqdm import tqdm
# Local support import
from .support_flex_sta import get_sliding_tile_attention_mask
# USE OUR NEW PACKAGE!
from fastvideo_kernel import sliding_tile_attention
from torch.nn.attention.flex_attention import flex_attention
flex_attention = torch.compile(flex_attention, dynamic=False)
def flex_test(Q, K, V, kernel_size):
mask = get_sliding_tile_attention_mask(kernel_size, (6, 8, 8), (18, 48, 80), 0, 'cuda', 0)
output = flex_attention(Q, K, V, block_mask=mask)
return output
def h100_fwd_kernel_test(Q, K, V, kernel_size):
# Using the same parameters as the original test
o = sliding_tile_attention(Q, K, V, [kernel_size] * 24, 0, False, '18x48x80')
return o
def generate_tensor(shape, mean, std, dtype, device):
tensor = torch.randn(shape, dtype=dtype, device=device)
magnitude = torch.norm(tensor, dim=-1, keepdim=True)
scaled_tensor = tensor * (torch.randn(magnitude.shape, dtype=dtype, device=device) * std + mean) / magnitude
return scaled_tensor.contiguous()
def check_correctness(b, h, n, d, causal, mean, std, num_iterations=2):
print(f"Running correctness check: batch={b}, heads={h}, seq_len={n}, dim={d}")
kernel_size_ls = [(3, 3, 5), (3, 1, 10)]
for kernel_size in kernel_size_ls:
print(f"Testing kernel_size: {kernel_size}")
for xi in tqdm(range(num_iterations)):
torch.manual_seed(xi)
Q = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
K = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
V = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
tk_o = h100_fwd_kernel_test(Q, K, V, kernel_size)
pt_o = flex_test(Q, K, V, kernel_size)
diff = pt_o - tk_o
abs_diff = torch.abs(diff)
max_d = torch.max(abs_diff).item()
avg_d = torch.sum(abs_diff).item() / (b * h * n * d)
if max_d > 0.1:
print(f"Warning: Large diff detected! max={max_d}, avg={avg_d}")
print("\n✅ TEST COMPLETE: New package matches FlexAttention behavior.")
if __name__ == "__main__":
b, h, d = 2, 24, 128
n = 69120
causal = False
mean = 1e-1
std = 10
check_correctness(b, h, n, d, causal, mean, std, num_iterations=2)
-95
View File
@@ -1,95 +0,0 @@
import torch
from .support_flex_sta import get_sliding_tile_attention_mask
from fastvideo_kernel import sliding_tile_attention
from torch.nn.attention.flex_attention import flex_attention
# from flash_attn_interface import flash_attn_func
from tqdm import tqdm
flex_attention = torch.compile(flex_attention, dynamic=False)
def flex_test(Q, K, V, kernel_size):
mask = get_sliding_tile_attention_mask(kernel_size, (6, 8, 8), (18, 48, 80), 0, 'cuda', 0)
output = flex_attention(Q, K, V, block_mask=mask)
return output
def h100_fwd_kernel_test(Q, K, V, kernel_size):
o = sliding_tile_attention(Q, K, V, [kernel_size] * 24, 0, False, '18x48x80')
return o
def generate_tensor(shape, mean, std, dtype, device):
tensor = torch.randn(shape, dtype=dtype, device=device)
magnitude = torch.norm(tensor, dim=-1, keepdim=True)
scaled_tensor = tensor * (torch.randn(magnitude.shape, dtype=dtype, device=device) * std + mean) / magnitude
return scaled_tensor.contiguous()
def check_correctness(b, h, n, d, causal, mean, std, num_iterations=50, error_mode='all'):
results = {
'TK vs FLEX': {
'sum_diff': 0,
'sum_abs': 0,
'max_diff': 0
},
}
kernel_size_ls = [(3, 3, 5), (3, 1, 10)]
from tqdm import tqdm
for kernel_size in tqdm(kernel_size_ls):
for _ in range(num_iterations):
torch.manual_seed(0)
Q = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
K = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
V = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
tk_o = h100_fwd_kernel_test(Q, K, V, kernel_size)
pt_o = flex_test(Q, K, V, kernel_size)
diff = pt_o - tk_o
abs_diff = torch.abs(diff)
results['TK vs FLEX']['sum_diff'] += torch.sum(abs_diff).item()
results['TK vs FLEX']['max_diff'] = max(results['TK vs FLEX']['max_diff'], torch.max(abs_diff).item())
torch.cuda.empty_cache()
print("kernel_size", kernel_size)
print("max_diff", torch.max(abs_diff).item())
print(
"avg_diff",
torch.sum(abs_diff).item() / (b * h * n * d *
(1 if error_mode == 'output' else 3 if error_mode == 'backward' else 4)))
total_elements = b * h * n * d * num_iterations * (1 if error_mode == 'output' else
3 if error_mode == 'backward' else 4) * len(kernel_size_ls)
for name, data in results.items():
avg_diff = data['sum_diff'] / total_elements
max_diff = data['max_diff']
results[name] = {'avg_diff': avg_diff, 'max_diff': max_diff}
return results
# Example usage
def test_sliding_tile_attention():
if not torch.cuda.is_available():
return
b, h, d = 2, 24, 128
n = 69120 # Sequence length
causal = False
mean = 1e-1
std = 10
# Run correctness check directly
results = check_correctness(b, h, n, d, causal, mean, std, error_mode='output')
assert results['TK vs FLEX']['avg_diff'] < 3e-6, f"Average difference: {results['TK vs FLEX']['avg_diff']} is too large"
assert results['TK vs FLEX']['max_diff'] < 4e-2, f"Maximum difference: {results['TK vs FLEX']['max_diff']} is too large"
print(f"Average difference: {results['TK vs FLEX']['avg_diff']}")
print(f"Maximum difference: {results['TK vs FLEX']['max_diff']}")
if __name__ == "__main__":
test_sliding_tile_attention()
-282
View File
@@ -1,282 +0,0 @@
import torch
import pytest
import sys
import os
import numpy as np
from tqdm import tqdm
from .utils import generate_block_sparse_mask_for_function, create_full_mask_from_block_mask
# Use installed package
# from fastvideo_kernel import video_sparse_attn as block_sparse_attn
BLOCK_M = 64
BLOCK_N = 64
def pytorch_test(Q, K, V, block_sparse_mask, dO):
q_ = Q.clone().float().requires_grad_()
k_ = K.clone().float().requires_grad_()
v_ = V.clone().float().requires_grad_()
QK = torch.matmul(q_, k_.transpose(-2, -1))
QK /= (q_.size(-1) ** 0.5)
QK = QK.masked_fill(~block_sparse_mask.unsqueeze(0), float('-inf'))
QK = torch.nn.functional.softmax(QK, dim=-1)
output = torch.matmul(QK, v_)
dO_ = dO
output.backward(dO_)
return (
output.to(torch.bfloat16),
q_.grad.to(torch.bfloat16),
k_.grad.to(torch.bfloat16),
v_.grad.to(torch.bfloat16),
)
def block_sparse_kernel_test(Q, K, V, block_sparse_mask, variable_block_sizes, q_non_pad_index, kv_non_pad_index, q_num_blocks, kv_num_blocks, dO):
Q = Q.detach().requires_grad_()
K = K.detach().requires_grad_()
V = V.detach().requires_grad_()
q_padded = vsa_pad(Q, q_non_pad_index, q_num_blocks, BLOCK_M)
k_padded = vsa_pad(K, kv_non_pad_index, kv_num_blocks, BLOCK_M)
v_padded = vsa_pad(V, kv_non_pad_index, kv_num_blocks, BLOCK_M)
# Use raw kernel or triton
try:
from fastvideo_kernel._C import fastvideo_kernel_ops
raw_kernel = getattr(fastvideo_kernel_ops, "block_sparse_fwd", None)
except ImportError:
raw_kernel = None
from fastvideo_kernel.triton_kernels.index import map_to_index
# Convert mask to indices
# block_sparse_mask is [H, M, N] bool
# We need to map it to index.
# block_sparse_mask needs to be expanded/reshaped?
# generate_block_sparse_mask_for_function returns [H, NumBlocksQ, NumBlocksKV]
# Ops.py logic:
# mask = torch.zeros_like(scores, dtype=torch.bool).scatter_(-1, topk_idx, True)
# idx, num = map_to_index(mask)
idx, num = map_to_index(block_sparse_mask.unsqueeze(0)) # Add batch dim [1, H, M, N]
if raw_kernel:
out_s = raw_kernel(q_padded, k_padded, v_padded, idx, num, variable_block_sizes.int())
output = out_s[0]
else:
# Fallback to triton testing if C++ not available
from fastvideo_kernel.triton_kernels.block_sparse_attn_triton import triton_block_sparse_attn_forward
output, _ = triton_block_sparse_attn_forward(q_padded, k_padded, v_padded, idx, num, variable_block_sizes)
output = output[:, :, q_non_pad_index, :]
output.backward(dO)
return output, Q.grad, K.grad, V.grad
def get_non_pad_index(
vid_len: torch.LongTensor,
n_win: int,
win_size: int,
):
device = vid_len.device
starts_pad = torch.arange(n_win, device=device) * win_size
index_pad = starts_pad[:, None] + torch.arange(win_size, device=device)[None, :]
index_mask = torch.arange(win_size, device=device)[None, :] < vid_len[:, None]
return index_pad[index_mask]
def generate_tensor(shape, dtype, device):
tensor = torch.randn(shape, dtype=dtype, device=device)
return tensor
def generate_variable_block_sizes(num_blocks, min_size=16, max_size=64, device="cuda"):
return torch.randint(min_size, max_size + 1, (num_blocks,), device=device, dtype=torch.int32)
def vsa_pad(x, non_pad_index, num_blocks, block_size):
padded_x = torch.zeros((1, x.shape[1], num_blocks * BLOCK_M, x.shape[3]), device=x.device, dtype=x.dtype)
padded_x[:, :, non_pad_index, :] = x
return padded_x
def check_correctness(h, d, num_blocks, k, num_iterations=20, error_mode='all'):
results = {
'gO': {'sum_diff': 0.0, 'sum_abs': 0.0, 'max_diff': 0.0},
'gQ': {'sum_diff': 0.0, 'sum_abs': 0.0, 'max_diff': 0.0},
'gK': {'sum_diff': 0.0, 'sum_abs': 0.0, 'max_diff': 0.0},
'gV': {'sum_diff': 0.0, 'sum_abs': 0.0, 'max_diff': 0.0},
}
device = "cuda" if torch.cuda.is_available() else "cpu"
variable_block_sizes = generate_variable_block_sizes(num_blocks, device=device)
S = int(variable_block_sizes.sum().item())
padded_S = num_blocks * BLOCK_M
non_pad_index = get_non_pad_index(variable_block_sizes, num_blocks, BLOCK_M)
block_mask = generate_block_sparse_mask_for_function(h, num_blocks, num_blocks, k, device)
full_mask = create_full_mask_from_block_mask(block_mask, variable_block_sizes, variable_block_sizes, device)
for _ in range(num_iterations):
Q = generate_tensor((1, h, S, d), torch.bfloat16, device)
K = generate_tensor((1, h, S, d), torch.bfloat16, device)
V = generate_tensor((1, h, S, d), torch.bfloat16, device)
dO = generate_tensor((1, h, S, d), torch.bfloat16, device)
# print(Q.shape, K.shape, V.shape, dO.shape)
# dO_padded = torch.zeros_like(dO_padded)
# dO_padded[:, :, non_pad_index, :] = dO
pt_o, pt_qg, pt_kg, pt_vg = pytorch_test(Q, K, V, full_mask, dO)
bs_o, bs_qg, bs_kg, bs_vg = block_sparse_kernel_test(Q, K, V, block_mask.unsqueeze(0), variable_block_sizes, non_pad_index, non_pad_index, num_blocks, num_blocks, dO)
for name, (pt, bs) in zip(['gQ', 'gK', 'gV', 'gO'], [(pt_qg, bs_qg), (pt_kg, bs_kg), (pt_vg, bs_vg), (pt_o, bs_o)]):
if bs is not None:
diff = pt - bs
abs_diff = torch.abs(diff)
results[name]['sum_diff'] += torch.sum(abs_diff).item()
results[name]['sum_abs'] += torch.sum(torch.abs(pt)).item()
rel_max_diff = torch.max(abs_diff) / torch.mean(torch.abs(pt))
results[name]['max_diff'] = max(results[name]['max_diff'], rel_max_diff.item())
if torch.cuda.is_available():
torch.cuda.empty_cache()
total_elements = h * S * d * num_iterations
for name, data in results.items():
avg_diff = data['sum_diff'] / total_elements
max_diff = data['max_diff']
results[name] = {'avg_diff': avg_diff, 'max_diff': max_diff}
return results
def check_correctness_qkdiff(h, d, num_q_blocks, num_kv_blocks, k, num_iterations=20, error_mode='all'):
results = {
'gO': {'sum_diff': 0.0, 'sum_abs': 0.0, 'max_diff': 0.0},
'gQ': {'sum_diff': 0.0, 'sum_abs': 0.0, 'max_diff': 0.0},
'gK': {'sum_diff': 0.0, 'sum_abs': 0.0, 'max_diff': 0.0},
'gV': {'sum_diff': 0.0, 'sum_abs': 0.0, 'max_diff': 0.0},
}
device = "cuda" if torch.cuda.is_available() else "cpu"
q_variable_block_sizes = generate_variable_block_sizes(num_q_blocks, device=device)
kv_variable_block_sizes = generate_variable_block_sizes(num_kv_blocks, device=device)
S_q = int(q_variable_block_sizes.sum().item())
S_kv = int(kv_variable_block_sizes.sum().item())
q_non_pad_index = get_non_pad_index(q_variable_block_sizes, num_q_blocks, BLOCK_M)
kv_non_pad_index = get_non_pad_index(kv_variable_block_sizes, num_kv_blocks, BLOCK_M)
block_mask = generate_block_sparse_mask_for_function(h, num_q_blocks, num_kv_blocks, k, device)
full_mask = create_full_mask_from_block_mask(block_mask, q_variable_block_sizes, kv_variable_block_sizes, device)
for _ in range(num_iterations):
Q = generate_tensor((1, h, S_q, d), torch.bfloat16, device)
K = generate_tensor((1, h, S_kv, d), torch.bfloat16, device)
V = generate_tensor((1, h, S_kv, d), torch.bfloat16, device)
dO = generate_tensor((1, h, S_q, d), torch.bfloat16, device)
# print(Q.shape, K.shape, V.shape, dO.shape)
pt_o, pt_qg, pt_kg, pt_vg = pytorch_test(Q, K, V, full_mask, dO)
bs_o, bs_qg, bs_kg, bs_vg = block_sparse_kernel_test(Q, K, V, block_mask.unsqueeze(0), kv_variable_block_sizes, q_non_pad_index, kv_non_pad_index, num_q_blocks, num_kv_blocks, dO)
for name, (pt, bs) in zip(['gQ', 'gK', 'gV', 'gO'], [(pt_qg, bs_qg), (pt_kg, bs_kg), (pt_vg, bs_vg), (pt_o, bs_o)]):
if bs is not None:
diff = pt - bs
abs_diff = torch.abs(diff)
results[name]['sum_diff'] += torch.sum(abs_diff).item()
results[name]['sum_abs'] += torch.sum(torch.abs(pt)).item()
rel_max_diff = torch.max(abs_diff) / torch.mean(torch.abs(pt))
results[name]['max_diff'] = max(results[name]['max_diff'], rel_max_diff.item())
if torch.cuda.is_available():
torch.cuda.empty_cache()
total_elements_q = h * S_q * d * num_iterations
total_elements_kv = h * S_kv * d * num_iterations
for name, data in results.items():
total_elements = total_elements_q if name in ['gQ', 'gO'] else total_elements_kv
avg_diff = data['sum_diff'] / total_elements
max_diff = data['max_diff']
results[name] = {'avg_diff': avg_diff, 'max_diff': max_diff}
return results
def generate_error_graphs(h, d, error_mode='all'):
test_configs = [
{"num_blocks": 16, "k": 2, "description": "Small sequence"},
{"num_blocks": 32, "k": 4, "description": "Medium sequence"},
{"num_blocks": 53, "k": 6, "description": "Large sequence"},
]
print(f"\nError Analysis for h={h}, d={d}, mode={error_mode}")
print("=" * 150)
print(f"{'Config':<20} {'Blocks':<8} {'K':<4} "
f"{'gQ Avg':<12} {'Rel gQ Max':<12} "
f"{'gK Avg':<12} {'Rel gK Max':<12} "
f"{'gV Avg':<12} {'Rel gV Max':<12} "
f"{'gO Avg':<12} {'Rel gO Max':<12}")
print("-" * 150)
for config in test_configs:
num_blocks = config["num_blocks"]
k = config["k"]
description = config["description"]
results = check_correctness(h, d, num_blocks, k, error_mode=error_mode)
print(f"{description:<20} {num_blocks:<8} {k:<4} "
f"{results['gQ']['avg_diff']:<12.6e} {results['gQ']['max_diff']:<12.6e} "
f"{results['gK']['avg_diff']:<12.6e} {results['gK']['max_diff']:<12.6e} "
f"{results['gV']['avg_diff']:<12.6e} {results['gV']['max_diff']:<12.6e} "
f"{results['gO']['avg_diff']:<12.6e} {results['gO']['max_diff']:<12.6e}")
print("-" * 150)
def generate_error_graphs_qkdiff(h, d, error_mode='all'):
test_configs = [
{"num_q_blocks": 16, "num_kv_blocks": 32, "k": 2, "description": "Small Q, Med KV"},
{"num_q_blocks": 32, "num_kv_blocks": 16, "k": 4, "description": "Med Q, Small KV"},
{"num_q_blocks": 53, "num_kv_blocks": 32, "k": 6, "description": "Large Q, Med KV"},
{"num_q_blocks": 16, "num_kv_blocks": 48, "k": 2, "description": "Small Q, Large KV"},
{"num_q_blocks": 48, "num_kv_blocks": 16, "k": 2, "description": "Large Q, Small KV"},
]
print(f"\nError Analysis (QK Diff) for h={h}, d={d}, mode={error_mode}")
print("=" * 150)
print(f"{'Config':<20} {'Q Blks':<8} {'KV Blks':<8} {'K':<4} "
f"{'gQ Avg':<12} {'Rel gQ Max':<12} "
f"{'gK Avg':<12} {'Rel gK Max':<12} "
f"{'gV Avg':<12} {'Rel gV Max':<12} "
f"{'gO Avg':<12} {'Rel gO Max':<12}")
print("-" * 150)
for config in test_configs:
num_q_blocks = config["num_q_blocks"]
num_kv_blocks = config["num_kv_blocks"]
k = config["k"]
description = config["description"]
results = check_correctness_qkdiff(h, d, num_q_blocks, num_kv_blocks, k, error_mode=error_mode)
print(f"{description:<20} {num_q_blocks:<8} {num_kv_blocks:<8} {k:<4} "
f"{results['gQ']['avg_diff']:<12.6e} {results['gQ']['max_diff']:<12.6e} "
f"{results['gK']['avg_diff']:<12.6e} {results['gK']['max_diff']:<12.6e} "
f"{results['gV']['avg_diff']:<12.6e} {results['gV']['max_diff']:<12.6e} "
f"{results['gO']['avg_diff']:<12.6e} {results['gO']['max_diff']:<12.6e}")
print("-" * 150)
@pytest.mark.skip()
def test_video_sparse_attention_backward():
if not torch.cuda.is_available():
return
h, d = 16, 128
print("Block Sparse Attention with Variable Block Sizes Analysis")
print("=" * 60)
for mode in ['backward']:
generate_error_graphs(h, d, error_mode=mode)
generate_error_graphs_qkdiff(h, d, error_mode=mode)
print("\nAnalysis completed for all modes.")
if __name__ == "__main__":
test_video_sparse_attention_backward()
-247
View File
@@ -1,247 +0,0 @@
import os
import sys
from typing import Tuple
import torch
from .utils import (
generate_block_sparse_mask_for_function,
create_full_mask_from_block_mask,
)
from .test_vsa import BLOCK_M # Import from local test_vsa
from . import test_vsa as ref
def pytorch_forward(
Q: torch.Tensor,
K: torch.Tensor,
V: torch.Tensor,
block_sparse_mask: torch.Tensor,
) -> torch.Tensor:
"""
Dense PyTorch reference forward:
- Q: [1, h, S_q, d]
- K,V: [1, h, S_kv, d]
- block_sparse_mask: [h, S_q, S_kv] bool
"""
q = Q.clone().float()
k = K.clone().float()
v = V.clone().float()
attn = torch.matmul(q, k.transpose(-2, -1)) # [1, h, S_q, S_kv]
attn = attn / (q.size(-1) ** 0.5)
attn = attn.masked_fill(~block_sparse_mask.unsqueeze(0), float("-inf"))
attn = torch.nn.functional.softmax(attn, dim=-1)
out = torch.matmul(attn, v) # [1, h, S_q, d]
return out.to(torch.bfloat16)
def block_sparse_forward_test(
Q: torch.Tensor,
K: torch.Tensor,
V: torch.Tensor,
block_sparse_mask: torch.Tensor,
variable_block_sizes: torch.Tensor,
q_non_pad_index: torch.Tensor,
kv_non_pad_index: torch.Tensor,
q_num_blocks: int,
kv_num_blocks: int,
) -> torch.Tensor:
"""
Forward-only wrapper
"""
Q = Q.detach()
K = K.detach()
V = V.detach()
q_padded = ref.vsa_pad(Q, q_non_pad_index, q_num_blocks, BLOCK_M)
k_padded = ref.vsa_pad(K, kv_non_pad_index, kv_num_blocks, BLOCK_M)
v_padded = ref.vsa_pad(V, kv_non_pad_index, kv_num_blocks, BLOCK_M)
# Use raw kernel or triton
try:
from fastvideo_kernel._C import fastvideo_kernel_ops
raw_kernel = getattr(fastvideo_kernel_ops, "block_sparse_fwd", None)
except ImportError:
raw_kernel = None
from fastvideo_kernel.triton_kernels.index import map_to_index
idx, num = map_to_index(block_sparse_mask)
if raw_kernel:
out_padded = raw_kernel(q_padded, k_padded, v_padded, idx, num, variable_block_sizes.int())[0]
else:
from fastvideo_kernel.triton_kernels.block_sparse_attn_triton import triton_block_sparse_attn_forward
out_padded, _ = triton_block_sparse_attn_forward(
q_padded, k_padded, v_padded, idx, num, variable_block_sizes
)
# Remove padding on the query side
out = out_padded[:, :, q_non_pad_index, :]
return out
def run_forward_equal_qk(
h: int = 16,
d: int = 128,
num_blocks: int = 16,
k: int = 2,
num_iterations: int = 5,
) -> Tuple[float, float]:
"""
Forward-only correctness test for the case S_q == S_kv.
Mirrors `check_correctness` but only compares forward outputs.
"""
assert torch.cuda.is_available(), "VSA kernels require CUDA"
device = "cuda"
variable_block_sizes = ref.generate_variable_block_sizes(
num_blocks, device=device
)
S = int(variable_block_sizes.sum().item())
non_pad_index = ref.get_non_pad_index(
variable_block_sizes, num_blocks, BLOCK_M
)
block_mask = generate_block_sparse_mask_for_function(
h, num_blocks, num_blocks, k, device
)
full_mask = create_full_mask_from_block_mask(
block_mask, variable_block_sizes, variable_block_sizes, device
)
print(f"[qkequal] h: {h}, d: {d}, num_blocks: {num_blocks}, k: {k}")
print(f"[qkequal] variable_block_sizes: {variable_block_sizes}, non_pad_index: {non_pad_index.shape}, block_mask: {block_mask.shape}, full_mask: {full_mask.shape}")
sum_diff = 0.0
sum_abs = 0.0
max_rel_diff = 0.0
for i in range(num_iterations):
Q = ref.generate_tensor((1, h, S, d), torch.bfloat16, device)
K = ref.generate_tensor((1, h, S, d), torch.bfloat16, device)
V = ref.generate_tensor((1, h, S, d), torch.bfloat16, device)
if i == 0: print(f"[qkequal] Q: {Q.shape}, K: {K.shape}, V: {V.shape}, full_mask: {full_mask.shape}")
if i == 0: print(f"[qkequal] block_mask: {block_mask.shape}")
pt_o = pytorch_forward(Q, K, V, full_mask)
bs_o = block_sparse_forward_test(
Q,
K,
V,
block_mask.unsqueeze(0),
variable_block_sizes,
non_pad_index,
non_pad_index,
num_blocks,
num_blocks,
)
diff = (pt_o - bs_o).abs()
sum_diff += diff.sum().item()
sum_abs += pt_o.abs().sum().item()
rel_max = diff.max() / (pt_o.abs().mean() + 1e-6)
max_rel_diff = max(max_rel_diff, rel_max.item())
total_elems = h * S * d * num_iterations
avg_abs_err = sum_diff / total_elems
return avg_abs_err, max_rel_diff
def run_forward_qk_diff(
h: int = 16,
d: int = 128,
num_q_blocks: int = 16,
num_kv_blocks: int = 32,
k: int = 2,
num_iterations: int = 5,
) -> Tuple[float, float]:
"""
Forward-only correctness test for the case S_q != S_kv.
NOTE:
- The Triton backend supports different Q/KV logical lengths via padding.
- The SM90 (H100) CUDA backend currently assumes the same number of blocks
for Q and KV, so we skip this test there.
"""
assert torch.cuda.is_available(), "VSA kernels require CUDA"
device = "cuda"
q_variable_block_sizes = ref.generate_variable_block_sizes(
num_q_blocks, device=device
)
kv_variable_block_sizes = ref.generate_variable_block_sizes(
num_kv_blocks, device=device
)
S_q = int(q_variable_block_sizes.sum().item())
S_kv = int(kv_variable_block_sizes.sum().item())
q_non_pad_index = ref.get_non_pad_index(
q_variable_block_sizes, num_q_blocks, BLOCK_M
)
kv_non_pad_index = ref.get_non_pad_index(
kv_variable_block_sizes, num_kv_blocks, BLOCK_M
)
block_mask = generate_block_sparse_mask_for_function(
h, num_q_blocks, num_kv_blocks, k, device
)
full_mask = create_full_mask_from_block_mask(
block_mask, q_variable_block_sizes, kv_variable_block_sizes, device
)
sum_diff = 0.0
sum_abs = 0.0
max_rel_diff = 0.0
for _ in range(num_iterations):
Q = ref.generate_tensor((1, h, S_q, d), torch.bfloat16, device)
K = ref.generate_tensor((1, h, S_kv, d), torch.bfloat16, device)
V = ref.generate_tensor((1, h, S_kv, d), torch.bfloat16, device)
pt_o = pytorch_forward(Q, K, V, full_mask)
bs_o = block_sparse_forward_test(
Q,
K,
V,
block_mask.unsqueeze(0),
kv_variable_block_sizes,
q_non_pad_index,
kv_non_pad_index,
num_q_blocks,
num_kv_blocks,
)
diff = (pt_o - bs_o).abs()
sum_diff += diff.sum().item()
sum_abs += pt_o.abs().sum().item()
rel_max = diff.max() / (pt_o.abs().mean() + 1e-6)
max_rel_diff = max(max_rel_diff, rel_max.item())
total_elems = h * S_q * d * num_iterations
avg_abs_err = sum_diff / total_elems
return avg_abs_err, max_rel_diff
def test_video_sparse_attention_forward():
if not torch.cuda.is_available():
return
h, d = 16, 128
print("Forward Block Sparse Attention Check (QK Equal)")
print("=" * 80)
avg_err_eq, max_rel_eq = run_forward_equal_qk(h, d, num_blocks=32, k=2)
print(f"QK equal: avg |ΔO| = {avg_err_eq:.6e}, max rel ΔO = {max_rel_eq:.6e}")
print("\nForward Block Sparse Attention Check (QK Different)")
print("=" * 80)
avg_err_diff, max_rel_diff = run_forward_qk_diff(
h, d, num_q_blocks=32, num_kv_blocks=48, k=2
)
print(
f"QK diff: avg |ΔO| = {avg_err_diff:.6e}, max rel ΔO = {max_rel_diff:.6e}"
)
if __name__ == "__main__":
test_video_sparse_attention_forward()
-60
View File
@@ -1,60 +0,0 @@
import torch
def generate_block_sparse_mask_for_function(h, num_q_blocks, num_kv_blocks, k, device="cuda"):
"""
Generate block sparse mask of shape [h, num_q_blocks, num_kv_blocks].
Args:
h: number of heads
num_q_blocks: number of query blocks
num_kv_blocks: number of key/value blocks
k: number of kv blocks each q block attends to
device: device to create tensors on
Returns:
block_sparse_mask: [h, num_q_blocks, num_kv_blocks] bool tensor
"""
k = min(k, num_kv_blocks)
scores = torch.rand(h, num_q_blocks, num_kv_blocks, device=device)
_, indices = torch.topk(scores, k, dim=-1)
block_sparse_mask = torch.zeros(h, num_q_blocks, num_kv_blocks, dtype=torch.bool, device=device)
block_sparse_mask = block_sparse_mask.scatter_(2, indices, 1).bool()
return block_sparse_mask
def create_full_mask_from_block_mask(block_sparse_mask, q_variable_block_sizes,
kv_variable_block_sizes, device="cuda"):
"""
Convert block-level sparse mask to full attention mask.
Args:
block_sparse_mask: [h, num_q_blocks, num_kv_blocks] bool tensor
q_variable_block_sizes: [num_q_blocks] tensor
kv_variable_block_sizes: [num_kv_blocks] tensor
device: device to create tensors on
Returns:
full_mask: [h, S_q, S_kv] bool tensor where S = total sequence length
"""
h, num_q_blocks, num_kv_blocks = block_sparse_mask.shape
total_q_seq_len = q_variable_block_sizes.sum().item()
total_kv_seq_len = kv_variable_block_sizes.sum().item()
q_cumsum = torch.cat([torch.tensor([0], device=device), q_variable_block_sizes.cumsum(dim=0)[:-1]])
kv_cumsum = torch.cat([torch.tensor([0], device=device), kv_variable_block_sizes.cumsum(dim=0)[:-1]])
full_mask = torch.zeros(h, total_q_seq_len, total_kv_seq_len, dtype=torch.bool, device=device)
for head in range(h):
for q_block in range(num_q_blocks):
q_start = q_cumsum[q_block]
q_end = q_start + q_variable_block_sizes[q_block]
for kv_block in range(num_kv_blocks):
if block_sparse_mask[head, q_block, kv_block]:
kv_start = kv_cumsum[kv_block]
kv_end = kv_start + kv_variable_block_sizes[kv_block]
full_mask[head, q_start:q_end, kv_start:kv_end] = True
return full_mask
+2 -1
View File
@@ -82,7 +82,8 @@ class SDPAImpl(AttentionImpl):
key = key.transpose(1, 2)
value = value.transpose(1, 2)
attn_mask = attn_metadata.attn_mask if attn_metadata is not None else None
if attn_metadata is not None:
attn_mask = getattr(attn_metadata, "attn_mask", None)
attn_kwargs = {
"attn_mask": attn_mask,
"dropout_p": self.dropout,
+12 -11
View File
@@ -53,6 +53,7 @@ def run_test(pytest_command: str):
git clone {git_repo} /FastVideo &&
cd /FastVideo &&
{checkout_command} &&
uv pip install -e fastvideo-kernel &&
uv pip install -e .[test] &&
{pytest_command}
"""
@@ -63,11 +64,11 @@ def run_test(pytest_command: str):
sys.exit(result.returncode)
@app.function(gpu="H100:1", image=image, timeout=1200, secrets=[modal.Secret.from_dict({"HF_API_KEY": os.environ.get("HF_API_KEY", "")})])
@app.function(gpu="H100:1", image=image, timeout=900, secrets=[modal.Secret.from_dict({"HF_API_KEY": os.environ.get("HF_API_KEY", "")})])
def run_encoder_tests():
run_test("hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/encoders -vs")
@app.function(gpu="L40S:1", image=image, timeout=1200, secrets=[modal.Secret.from_dict({"HF_API_KEY": os.environ.get("HF_API_KEY", "")})])
@app.function(gpu="L40S:1", image=image, timeout=900, secrets=[modal.Secret.from_dict({"HF_API_KEY": os.environ.get("HF_API_KEY", "")})])
def run_vae_tests():
run_test("hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/vaes -vs")
@@ -102,17 +103,17 @@ def run_inference_tests_STA():
run_test("pytest ./fastvideo/tests/inference/STA -srP")
@app.function(gpu="H100:1", image=image, timeout=900)
def run_kernel_tests():
run_test("pytest fastvideo-kernel/tests/ -vs")
def run_precision_tests_STA():
run_test("pytest fastvideo-kernel/tests/test_correctness.py")
# @app.function(gpu="H100:1", image=image, timeout=900)
# def run_precision_tests_VSA():
# # VSA correctness is covered by the same file now
# run_test("pytest fastvideo-kernel/tests/test_correctness.py")
@app.function(gpu="H100:1", image=image, timeout=900)
def run_precision_tests_VSA():
# VSA correctness is covered by the same file now
run_test("pytest fastvideo-kernel/tests/test_correctness.py")
# @app.function(gpu="L40S:1", image=image, timeout=900)
# def run_precision_tests_vmoba():
# run_test("pytest fastvideo-kernel/tests/test_vmoba_correctness.py")
@app.function(gpu="L40S:1", image=image, timeout=900)
def run_precision_tests_vmoba():
run_test("pytest fastvideo-kernel/tests/test_vmoba_correctness.py")
@app.function(gpu="L40S:1", image=image, timeout=900)
def run_inference_tests_vmoba():
@@ -1,10 +1,10 @@
{
"step_time": 0.6983645600266755,
"grad_norm": 1.118342604637146,
"grad_norm": 1.278342604637146,
"avg_step_time": 1.002151239803061,
"_timestamp": 1751181952.70901,
"vsa_sparsity": 0.05,
"learning_rate": 1e-05,
"train_loss": 0.2765433095693588,
"train_loss": 0.3085433095693588,
"_runtime": 107.325113071
}
+1 -1
View File
@@ -27,7 +27,7 @@ dependencies = [
"timm==1.0.11",
"peft>=0.15.0",
"diffusers>=0.33.1",
"torch>=2.9.1",
"torch>=2.9.0",
"torchvision",
# Acceleration & Optimization