diff --git a/scripts/test_block_sparse_attn.py b/scripts/test_block_sparse_attn.py deleted file mode 100644 index c5fe0d4..0000000 --- a/scripts/test_block_sparse_attn.py +++ /dev/null @@ -1,84 +0,0 @@ -import torch -import block_sparse_attn -from block_sparse_attn import block_sparse_attn_func, token_streaming_attn_func, block_streaming_attn_func - -print("block_sparse_attn version:", block_sparse_attn.__version__) -print("Available funcs:", [k for k in dir(block_sparse_attn) if not k.startswith("_")]) - -device = "cuda:0" -batch_size = 1 -nheads = 4 -head_dim = 64 -seqlen = 256 - -total_len = batch_size * seqlen -q = torch.randn(total_len, nheads, head_dim, device=device, dtype=torch.float16) -k = torch.randn(total_len, nheads, head_dim, device=device, dtype=torch.float16) -v = torch.randn(total_len, nheads, head_dim, device=device, dtype=torch.float16) - -cu_seqlens_q = torch.arange(0, total_len + 1, seqlen, device=device, dtype=torch.int32) -cu_seqlens_k = cu_seqlens_q.clone() - -# 1. dense attention (all heads mask_type=0) -head_mask_type_dense = torch.zeros(nheads, device=device, dtype=torch.int32) -streaming_info = torch.tensor([1, 3] * nheads, device=device, dtype=torch.int32) -out_dense = block_sparse_attn_func( - q, k, v, - cu_seqlens_q, cu_seqlens_k, - head_mask_type_dense, - streaming_info, - None, - seqlen, seqlen, - 0.0, - is_causal=False, -) -print("dense output shape:", out_dense.shape) -assert out_dense.shape == q.shape - -# 2. block-sparse attention (2 dense heads + 2 blocksparse heads) -num_blocksparse_heads = 2 -nrow = seqlen // 128 -ncol = seqlen // 128 -base_blockmask = torch.zeros(batch_size, num_blocksparse_heads, nrow, ncol, device=device, dtype=torch.bool) -base_blockmask[:, :, :, :ncol//2] = True # attend to first half - -head_mask_type_sparse = torch.tensor([0, 0, 1, 1], device=device, dtype=torch.int32) -out_sparse = block_sparse_attn_func( - q, k, v, - cu_seqlens_q, cu_seqlens_k, - head_mask_type_sparse, - streaming_info, - base_blockmask, - seqlen, seqlen, - 0.0, - is_causal=False, -) -print("block-sparse output shape:", out_sparse.shape) -assert out_sparse.shape == q.shape - -# 3. token-level streaming attention -head_mask_type_stream = torch.full((nheads,), -1, device=device, dtype=torch.int32) -out_stream = token_streaming_attn_func( - q, k, v, - cu_seqlens_q, cu_seqlens_k, - head_mask_type_stream, - streaming_info, - seqlen, seqlen, -) -print("token-streaming output shape:", out_stream.shape) -assert out_stream.shape == q.shape - -# 4. block-level streaming attention -out_block_stream = block_streaming_attn_func( - q, k, v, - cu_seqlens_q, cu_seqlens_k, - head_mask_type_stream, - streaming_info, - seqlen, seqlen, - 0.0, - is_causal=True, -) -print("block-streaming output shape:", out_block_stream.shape) -assert out_block_stream.shape == q.shape - -print("All smoke tests passed!")