22"""
33Benchmark VSA *wrapper* performance (forward + backward) and report TFLOPs.
44
5- This script benchmarks the autograd-enabled wrapper:
6- - fastvideo_kernel.block_sparse_attn.block_sparse_attn
5+ This script benchmarks the autograd-enabled wrappers:
6+ - 64-token TK/Triton: fastvideo_kernel.block_sparse_attn.block_sparse_attn
7+ - 128/256-token Triton/CuTe: fastvideo_kernel.block_sparse_attn_256
78
89So measured time includes wrapper overhead (map->index conversion, dispatch) plus kernel time.
910"""
2324except Exception as e : # pragma: no cover
2425 raise ImportError ("This benchmark requires triton (for triton.testing.do_bench)." ) from e
2526
26- BLOCK_M = 64
27- BLOCK_N = 64
28-
2927
3028def set_seed (seed : int = 42 ) -> None :
3129 random .seed (seed )
@@ -41,7 +39,11 @@ def parse_arguments() -> argparse.Namespace:
4139 p .add_argument ("--num_heads" , type = int , default = 12 )
4240 p .add_argument ("--head_dim" , type = int , default = 128 , choices = [64 , 128 ])
4341 p .add_argument ("--topk" , type = int , default = None , help = "KV blocks per Q block (default: ~90%% sparsity)" )
44- p .add_argument ("--q_seq_lens" , type = int , nargs = "+" , default = [49152 ], help = "Q sequence lengths (must be /64)" )
42+ p .add_argument ("--q_seq_lens" ,
43+ type = int ,
44+ nargs = "+" ,
45+ default = [49152 ],
46+ help = "Q sequence lengths (must be divisible by --block_size)" )
4547 p .add_argument ("--kv_seq_lens" ,
4648 type = int ,
4749 nargs = "+" ,
@@ -51,9 +53,13 @@ def parse_arguments() -> argparse.Namespace:
5153 p .add_argument ("--rep" , type = int , default = 20 )
5254 p .add_argument ("--seed" , type = int , default = 42 )
5355 p .add_argument ("--dtype" , type = str , default = "bf16" , choices = ["bf16" , "fp16" ])
56+ p .add_argument ("--block_size" , type = int , default = 64 , choices = [64 , 128 , 256 ])
5457 p .add_argument ("--force_triton" ,
5558 action = "store_true" ,
5659 help = "Force wrapper to use Triton path (if supported by shapes)." )
60+ p .add_argument ("--use_cute" ,
61+ action = "store_true" ,
62+ help = "Use the optional FA4 CuTe forward/backward path (requires --block_size 128 or 256)." )
5763 return p .parse_args ()
5864
5965
@@ -84,18 +90,38 @@ def bench_ms(fn: Callable[[], object], warmup: int, rep: int) -> float:
8490 return do_bench (fn , warmup = warmup , rep = rep , quantiles = None )
8591
8692
93+ def _configure_backend (args : argparse .Namespace ) -> None :
94+ if args .use_cute and args .block_size not in (128 , 256 ):
95+ raise ValueError ("--use_cute requires --block_size 128 or 256" )
96+ if args .use_cute and args .force_triton :
97+ raise ValueError ("--use_cute and --force_triton are mutually exclusive" )
98+
99+ if args .force_triton :
100+ os .environ .pop ("FASTVIDEO_VSA_CUTEDSL" , None )
101+ os .environ ["FASTVIDEO_KERNEL_VSA_FORCE_TRITON" ] = "1"
102+ elif args .use_cute :
103+ os .environ .pop ("FASTVIDEO_VSA_TRITON" , None )
104+ os .environ .pop ("FASTVIDEO_KERNEL_VSA_FORCE_TRITON" , None )
105+ os .environ ["FASTVIDEO_VSA_CUTEDSL" ] = "1"
106+
107+
87108def main () -> None :
88109 args = parse_arguments ()
89110 set_seed (args .seed )
111+ _configure_backend (args )
90112
91113 dtype = torch .bfloat16 if args .dtype == "bf16" else torch .float16
92114
93- if args .force_triton :
94- os .environ ["FASTVIDEO_KERNEL_VSA_FORCE_TRITON" ] = "1"
95-
96115 from fastvideo_kernel .block_sparse_attn import block_sparse_attn
116+ from fastvideo_kernel .block_sparse_attn_256 import block_sparse_attn_128 , block_sparse_attn_256
97117
98118 bs , h , d = args .batch_size , args .num_heads , args .head_dim
119+ block_size = args .block_size
120+ attention = {
121+ 64 : block_sparse_attn ,
122+ 128 : block_sparse_attn_128 ,
123+ 256 : block_sparse_attn_256 ,
124+ }[block_size ]
99125 kv_seq_lens = args .kv_seq_lens
100126 if kv_seq_lens is None :
101127 kv_seq_lens = args .q_seq_lens
@@ -105,20 +131,22 @@ def main() -> None:
105131 print ("VSA Block-Sparse Attention Benchmark (WRAPPER)" )
106132 print (f"device: { torch .cuda .get_device_name (0 )} " )
107133 print (f"batch={ bs } , heads={ h } , head_dim={ d } , dtype={ args .dtype } " )
108- print (f"BLOCK_M= { BLOCK_M } , BLOCK_N= { BLOCK_N } " )
134+ print (f"block_size= { block_size } " )
109135 print ("NOTE: timings include wrapper overhead (map->index + dispatch)." )
110- if args .force_triton :
136+ if args .use_cute :
137+ print ("dispatch: FA4 CuTe" )
138+ elif args .force_triton :
111139 print ("dispatch: forced Triton (FASTVIDEO_KERNEL_VSA_FORCE_TRITON=1)" )
112140 else :
113141 print ("dispatch: SM90 if available, else Triton" )
114142
115143 for q_len , kv_len in zip (args .q_seq_lens , kv_seq_lens ):
116- if q_len % BLOCK_M != 0 or kv_len % BLOCK_N != 0 :
117- print (f"[skip] q_len={ q_len } , kv_len={ kv_len } must be divisible by 64 " )
144+ if q_len % block_size != 0 or kv_len % block_size != 0 :
145+ print (f"[skip] q_len={ q_len } , kv_len={ kv_len } must be divisible by { block_size } " )
118146 continue
119147
120- num_q_blocks = q_len // BLOCK_M
121- num_kv_blocks = kv_len // BLOCK_N
148+ num_q_blocks = q_len // block_size
149+ num_kv_blocks = kv_len // block_size
122150 topk = args .topk if args .topk is not None else max (1 , num_kv_blocks // 10 )
123151 topk = min (topk , num_kv_blocks )
124152
@@ -129,11 +157,11 @@ def main() -> None:
129157 q , k , v = create_qkv (bs , h , q_len , kv_len , d , dtype )
130158 block_map = make_block_map (bs , h , num_q_blocks , num_kv_blocks , topk )
131159
132- # Variable block sizes: default full blocks (64 tokens per KV block)
133- variable_block_sizes = torch .full ((num_kv_blocks , ), BLOCK_N , dtype = torch .int32 , device = "cuda" )
160+ # Variable block sizes: default full logical blocks.
161+ variable_block_sizes = torch .full ((num_kv_blocks , ), block_size , dtype = torch .int32 , device = "cuda" )
134162
135163 def _fwd ():
136- return block_sparse_attn (q , k , v , block_map , variable_block_sizes )
164+ return attention (q , k , v , block_map , variable_block_sizes )
137165
138166 fwd_ms = bench_ms (_fwd , warmup = args .warmup , rep = args .rep )
139167
@@ -142,7 +170,7 @@ def _fwd():
142170 q_ = q .detach ().requires_grad_ (True )
143171 k_ = k .detach ().requires_grad_ (True )
144172 v_ = v .detach ().requires_grad_ (True )
145- o_ , _aux_ = block_sparse_attn (q_ , k_ , v_ , block_map , variable_block_sizes )
173+ o_ , _aux_ = attention (q_ , k_ , v_ , block_map , variable_block_sizes )
146174 og = torch .randn_like (o_ )
147175 loss = (o_ * og ).sum ()
148176
@@ -156,7 +184,7 @@ def _fwd():
156184 rep = max (5 , args .rep // 2 ),
157185 )
158186
159- flops = flops_sparse_attention (bs , h , d , q_len , topk , BLOCK_N )
187+ flops = flops_sparse_attention (bs , h , d , q_len , topk , block_size )
160188 fwd_tflops = flops / fwd_ms * 1e-12 * 1e3
161189 # Rough backward multiplier (attention backward typically ~2-3x forward)
162190 bwd_tflops = (2.5 * flops ) / bwd_ms * 1e-12 * 1e3
0 commit comments