1+ import random
2+ import torch
3+
4+ from vllm_flash_attn .utils .benchmark import benchmark_forward
5+
6+ from vllm_flash_attn .flash_attn_interface import (
7+ flash_attn_varlen_func ,
8+ tree_attention ,
9+ )
10+ from vllm_flash_attn .utils .tree import (
11+ create_tree_mask ,
12+ generate_q_and_block_kvcache ,
13+ treeify_output ,
14+ )
15+
16+
17+ def run_tree_attention_benchmark (
18+ seqlen_q : int = 1024 ,
19+ seqlen_k : int = 1024 ,
20+ spec_len : tuple [int ] = (8 ,8 ),
21+ random_seq_len : bool = False ,
22+ random_spec_len : bool = False ,
23+ batch_size : int = 8 ,
24+ nheads : int = 16 ,
25+ head_dim : int = 128 ,
26+ paged_kv_block_size : int = 256 ,
27+ dtype : torch .dtype = torch .float16 ,
28+ device : str = "cuda" ,
29+ ):
30+ """
31+ Benchmark tree_attention vs flash_attn_varlen_func performance.
32+
33+ Similar to test_paged_tree_attention but focused on performance measurement.
34+ """
35+ print ("Benchmarking with:" )
36+ print (f" seqlen_q: { seqlen_q } , seqlen_k: { seqlen_k } " )
37+ print (f" spec_len: { spec_len } , random_seq_len: { random_seq_len } , random_spec_len: { random_spec_len } " )
38+ print (f" batch_size: { batch_size } , nheads: { nheads } , head_dim: { head_dim } " )
39+ print (f" paged_kv_block_size: { paged_kv_block_size } , dtype: { dtype } " )
40+
41+ torch .set_default_device (device )
42+ torch .cuda .manual_seed_all (42 ) # Fixed seed for reproducibility
43+
44+ # Generate random sequence lengths and spec lengths similar to the test
45+ if random_seq_len :
46+ q_seqlens = [seqlen_q + random .randint (0 , 20 ) for _ in range (batch_size )]
47+ k_seqlens = [seqlen_k + random .randint (0 , 20 ) for _ in range (batch_size )]
48+ else :
49+ q_seqlens = [seqlen_q ]* batch_size
50+ k_seqlens = [seqlen_k ]* batch_size
51+
52+ if random_spec_len :
53+ speclens = [(spec_len [0 ]+ random .randint (0 , 7 ), spec_len [1 ]+ random .randint (1 , 2 )) for _ in range (batch_size )]
54+ else :
55+ speclens = [spec_len ]* batch_size
56+
57+ # Generate test data using the utility function
58+ (
59+ q_spec_tree ,
60+ q_seqlens_tree ,
61+ q_spec_batch ,
62+ q_seqlens_batch ,
63+ tree_block_table ,
64+ k_spec_tree ,
65+ v_spec_tree ,
66+ k_seqlens_tree ,
67+ batch_block_table ,
68+ k_spec_batch ,
69+ v_spec_batch ,
70+ k_seqlens_batch ,
71+ ) = generate_q_and_block_kvcache (
72+ q_seqlens , k_seqlens , speclens , paged_kv_block_size , nheads , head_dim , device , dtype
73+ )
74+
75+ # Create tree mask and cumulative sequence lengths
76+ tree_mask = create_tree_mask (speclens , device )
77+ tree_mask_lens = torch .tensor ([0 ] + [i * j for i ,j in speclens ], dtype = torch .int32 ).cumsum (dim = 0 , dtype = torch .int32 )
78+ cu_seqlens_q_tree = torch .tensor ([0 ] + q_seqlens_tree , dtype = torch .int32 ).cumsum (dim = 0 , dtype = torch .int32 )
79+ seqused_k_tree = torch .tensor (k_seqlens_tree , dtype = torch .int32 )
80+ cu_seqlens_q_batch = torch .tensor ([0 ] + q_seqlens_batch , dtype = torch .int32 ).cumsum (dim = 0 , dtype = torch .int32 )
81+ seqused_k_batch = torch .tensor (k_seqlens_batch , dtype = torch .int32 )
82+
83+
84+ print ("\n Running benchmarks..." )
85+
86+ # Benchmark tree_attention
87+ _ , tree_measurement = benchmark_forward (
88+ tree_attention ,
89+ q_spec_tree ,
90+ k_spec_tree ,
91+ v_spec_tree ,
92+ max (q_seqlens_tree ),
93+ cu_seqlens_q_tree ,
94+ max (k_seqlens_tree ),
95+ tree_mask ,
96+ tree_mask_lens ,
97+ seqused_k = seqused_k_tree ,
98+ block_table = tree_block_table ,
99+ desc = "tree_attention" ,
100+ verbose = False
101+ )
102+ tree_time = tree_measurement .mean
103+ print (f"tree_attention average time: { tree_time :.6f} seconds" )
104+
105+ # Benchmark flash_attn_varlen_func
106+ _ , varlen_measurement = benchmark_forward (
107+ flash_attn_varlen_func ,
108+ q_spec_batch ,
109+ k_spec_batch ,
110+ v_spec_batch ,
111+ max (q_seqlens_batch ),
112+ cu_seqlens_q_batch ,
113+ max (k_seqlens_batch ),
114+ seqused_k = seqused_k_batch ,
115+ causal = True ,
116+ block_table = batch_block_table ,
117+ desc = "flash_attn_varlen_func" ,
118+ verbose = False
119+ )
120+ varlen_time = varlen_measurement .mean
121+ print (f"flash_attn_varlen_func average time: { varlen_time :.6f} seconds" )
122+
123+ # Calculate speedup
124+ if varlen_time > 0 :
125+ speedup = varlen_time / tree_time
126+ print (f"Speedup (varlen/tree): { speedup :.2f} x" )
127+ if speedup > 1 :
128+ print (f"tree_attention is { speedup :.2f} x faster" )
129+ else :
130+ print (f"flash_attn_varlen_func is { 1 / speedup :.2f} x faster" )
131+
132+ # Verify correctness
133+ print ("\n Verifying correctness..." )
134+ tree_output = tree_attention (
135+ q_spec_tree ,
136+ k_spec_tree ,
137+ v_spec_tree ,
138+ max (q_seqlens_tree ),
139+ cu_seqlens_q_tree ,
140+ max (k_seqlens_tree ),
141+ tree_mask ,
142+ tree_mask_lens ,
143+ seqused_k = seqused_k_tree ,
144+ block_table = tree_block_table ,
145+ )
146+ varlen_output = flash_attn_varlen_func (
147+ q_spec_batch ,
148+ k_spec_batch ,
149+ v_spec_batch ,
150+ max (q_seqlens_batch ),
151+ cu_seqlens_q_batch ,
152+ max (k_seqlens_batch ),
153+ seqused_k = seqused_k_batch ,
154+ causal = True ,
155+ block_table = batch_block_table ,
156+ )
157+ varlen_output_treeified = treeify_output (varlen_output , q_seqlens , speclens )
158+ try :
159+ torch .testing .assert_close (tree_output , varlen_output_treeified , atol = 2e-2 , rtol = 1e-2 )
160+ except AssertionError as e :
161+ print ("✗ Outputs differ significantly!" )
162+ print (e )
163+ else :
164+ print ("✓ Outputs match within tolerance" )
165+ finally :
166+ max_diff = torch .max (torch .abs (tree_output - varlen_output_treeified )).item ()
167+ print (f"Maximum difference between outputs: { max_diff :.6f} " )
168+
169+ return {
170+ 'tree_time' : tree_time ,
171+ 'varlen_time' : varlen_time ,
172+ 'speedup' : varlen_time / tree_time if varlen_time > 0 else float ('inf' ),
173+ 'max_diff' : max_diff ,
174+ 'config' : {
175+ 'seqlen_q' : seqlen_q ,
176+ 'seqlen_k' : seqlen_k ,
177+ 'batch_size' : batch_size ,
178+ 'nheads' : nheads ,
179+ 'head_dim' : head_dim ,
180+ 'paged_kv_block_size' : paged_kv_block_size ,
181+ 'dtype' : str (dtype ),
182+ 'q_spec_tree.shape' : q_spec_tree .shape ,
183+ 'k_spec_tree.shape' : k_spec_tree .shape ,
184+ 'tree_mask.shape' : tree_mask .shape ,
185+ }
186+ }
187+
188+
189+ def run_decoding_benchmark ():
190+ """Run benchmarks for decoding scenario with seqlen_q=0."""
191+ configs = [
192+ # Small sequences with different spec_len and block sizes
193+ {'seqlen_q' : 0 , 'seqlen_k' : 128 , 'batch_size' : 4 , 'nheads' : 8 , 'head_dim' : 128 , 'spec_len' : (1 , 2 ), 'paged_kv_block_size' : 16 },
194+ {'seqlen_q' : 0 , 'seqlen_k' : 256 , 'batch_size' : 4 , 'nheads' : 8 , 'head_dim' : 128 , 'spec_len' : (2 , 3 ), 'paged_kv_block_size' : 16 },
195+
196+ # Medium sequences with varied spec_len and block sizes
197+ {'seqlen_q' : 0 , 'seqlen_k' : 512 , 'batch_size' : 8 , 'nheads' : 16 , 'head_dim' : 128 , 'spec_len' : (1 , 2 ), 'paged_kv_block_size' : 256 },
198+ {'seqlen_q' : 0 , 'seqlen_k' : 1024 , 'batch_size' : 8 , 'nheads' : 16 , 'head_dim' : 128 , 'spec_len' : (3 , 4 ), 'paged_kv_block_size' : 256 },
199+
200+ # Large sequences with larger block sizes
201+ {'seqlen_q' : 0 , 'seqlen_k' : 2048 , 'batch_size' : 4 , 'nheads' : 16 , 'head_dim' : 128 , 'spec_len' : (2 , 3 ), 'paged_kv_block_size' : 512 },
202+
203+ # Different head dimensions with varied block sizes
204+ {'seqlen_q' : 0 , 'seqlen_k' : 1024 , 'batch_size' : 8 , 'nheads' : 16 , 'head_dim' : 64 , 'spec_len' : (1 , 2 ), 'paged_kv_block_size' : 256 },
205+ {'seqlen_q' : 0 , 'seqlen_k' : 1024 , 'batch_size' : 8 , 'nheads' : 16 , 'head_dim' : 256 , 'spec_len' : (2 , 3 ), 'paged_kv_block_size' : 512 },
206+
207+ # Different batch sizes with randomization and block sizes
208+ {'seqlen_q' : 0 , 'seqlen_k' : 1024 , 'batch_size' : 2 , 'nheads' : 16 , 'head_dim' : 128 , 'spec_len' : (1 , 2 ), 'random_spec_len' : True , 'paged_kv_block_size' : 16 },
209+ {'seqlen_q' : 0 , 'seqlen_k' : 1024 , 'batch_size' : 16 , 'nheads' : 16 , 'head_dim' : 128 , 'spec_len' : (2 , 3 ), 'random_seq_len' : True , 'paged_kv_block_size' : 256 },
210+
211+ # High spec_len scenarios with different block sizes
212+ {'seqlen_q' : 0 , 'seqlen_k' : 1024 , 'batch_size' : 8 , 'nheads' : 16 , 'head_dim' : 128 , 'spec_len' : (4 , 5 ), 'paged_kv_block_size' : 256 },
213+ {'seqlen_q' : 0 , 'seqlen_k' : 1024 , 'batch_size' : 8 , 'nheads' : 16 , 'head_dim' : 128 , 'spec_len' : (6 , 8 ), 'paged_kv_block_size' : 512 },
214+
215+ # Block size comparison scenarios
216+ {'seqlen_q' : 0 , 'seqlen_k' : 1024 , 'batch_size' : 8 , 'nheads' : 16 , 'head_dim' : 128 , 'spec_len' : (2 , 3 ), 'paged_kv_block_size' : 16 },
217+ {'seqlen_q' : 0 , 'seqlen_k' : 1024 , 'batch_size' : 8 , 'nheads' : 16 , 'head_dim' : 128 , 'spec_len' : (2 , 3 ), 'paged_kv_block_size' : 256 },
218+ {'seqlen_q' : 0 , 'seqlen_k' : 1024 , 'batch_size' : 8 , 'nheads' : 16 , 'head_dim' : 128 , 'spec_len' : (2 , 3 ), 'paged_kv_block_size' : 512 },
219+ ]
220+
221+ print ("=" * 80 )
222+ print ("DECODING BENCHMARK (seqlen_q=0)" )
223+ print ("=" * 80 )
224+ print ("This benchmark represents the decoding scenario where tree attention" )
225+ print ("can be compared against batch expansion for generation tasks." )
226+ print ("=" * 80 )
227+
228+ results = []
229+ for i , config in enumerate (configs ):
230+ print (f"\n [{ i + 1 } /{ len (configs )} ] Decoding Configuration:" )
231+ result = run_tree_attention_benchmark (** config )
232+ results .append (result )
233+ print ("-" * 80 )
234+
235+ # Summary
236+ print ("\n " + "=" * 80 )
237+ print ("DECODING BENCHMARK SUMMARY" )
238+ print ("=" * 80 )
239+ print (f"{ 'Config' :<18} { 'Tree(ms)' :<10} { 'Varlen(ms)' :<12} { 'Speedup' :<10} { 'Max Diff' :<12} " )
240+ print ("-" * 80 )
241+
242+ for i , result in enumerate (results ):
243+ config = result ['config' ]
244+ config_str = f"{ config ['seqlen_q' ]} :{ config ['seqlen_k' ]} :{ config ['tree_mask.shape' ][0 ]} :{ config ['paged_kv_block_size' ]} "
245+ tree_ms = result ['tree_time' ] * 1000
246+ varlen_ms = result ['varlen_time' ] * 1000
247+ speedup = result ['speedup' ]
248+ max_diff = result ['max_diff' ]
249+
250+ print (f"{ config_str :<18} { tree_ms :<10.3f} { varlen_ms :<12.3f} { speedup :<10.2f} x { max_diff :<12.6f} " )
251+
252+ return results
253+
254+
255+ def run_comprehensive_benchmark ():
256+ """Run benchmarks across different configurations."""
257+ configs = [
258+ # Small sequences with different spec_len and block sizes
259+ {'seqlen_q' : 128 , 'seqlen_k' : 128 , 'batch_size' : 4 , 'nheads' : 8 , 'head_dim' : 128 , 'spec_len' : (1 , 2 ), 'paged_kv_block_size' : 16 },
260+ {'seqlen_q' : 256 , 'seqlen_k' : 256 , 'batch_size' : 4 , 'nheads' : 8 , 'head_dim' : 128 , 'spec_len' : (2 , 3 ), 'paged_kv_block_size' : 16 },
261+
262+ # Medium sequences with varied spec_len and block sizes
263+ {'seqlen_q' : 512 , 'seqlen_k' : 512 , 'batch_size' : 8 , 'nheads' : 16 , 'head_dim' : 128 , 'spec_len' : (1 , 2 ), 'paged_kv_block_size' : 256 },
264+ {'seqlen_q' : 1024 , 'seqlen_k' : 1024 , 'batch_size' : 8 , 'nheads' : 16 , 'head_dim' : 128 , 'spec_len' : (3 , 4 ), 'paged_kv_block_size' : 256 },
265+
266+ # Large sequences with larger block sizes
267+ {'seqlen_q' : 2048 , 'seqlen_k' : 2048 , 'batch_size' : 4 , 'nheads' : 16 , 'head_dim' : 128 , 'spec_len' : (2 , 3 ), 'paged_kv_block_size' : 512 },
268+
269+ # Different head dimensions with varied block sizes
270+ {'seqlen_q' : 1024 , 'seqlen_k' : 1024 , 'batch_size' : 8 , 'nheads' : 16 , 'head_dim' : 64 , 'spec_len' : (1 , 2 ), 'paged_kv_block_size' : 256 },
271+ {'seqlen_q' : 1024 , 'seqlen_k' : 1024 , 'batch_size' : 8 , 'nheads' : 16 , 'head_dim' : 256 , 'spec_len' : (2 , 3 ), 'paged_kv_block_size' : 512 },
272+
273+ # Different batch sizes with randomization and block sizes
274+ {'seqlen_q' : 1024 , 'seqlen_k' : 1024 , 'batch_size' : 2 , 'nheads' : 16 , 'head_dim' : 128 , 'spec_len' : (1 , 2 ), 'random_spec_len' : True , 'paged_kv_block_size' : 16 },
275+ {'seqlen_q' : 1024 , 'seqlen_k' : 1024 , 'batch_size' : 16 , 'nheads' : 16 , 'head_dim' : 128 , 'spec_len' : (2 , 3 ), 'random_seq_len' : True , 'paged_kv_block_size' : 256 },
276+
277+ # High spec_len scenarios with different block sizes
278+ {'seqlen_q' : 1024 , 'seqlen_k' : 1024 , 'batch_size' : 8 , 'nheads' : 16 , 'head_dim' : 128 , 'spec_len' : (4 , 5 ), 'paged_kv_block_size' : 256 },
279+ {'seqlen_q' : 1024 , 'seqlen_k' : 1024 , 'batch_size' : 8 , 'nheads' : 16 , 'head_dim' : 128 , 'spec_len' : (6 , 8 ), 'paged_kv_block_size' : 512 },
280+
281+ # Mixed randomization scenarios with block sizes
282+ {'seqlen_q' : 512 , 'seqlen_k' : 1024 , 'batch_size' : 8 , 'nheads' : 16 , 'head_dim' : 128 , 'spec_len' : (2 , 3 ), 'random_seq_len' : True , 'random_spec_len' : True , 'paged_kv_block_size' : 256 },
283+
284+ # Block size comparison scenarios
285+ {'seqlen_q' : 1024 , 'seqlen_k' : 1024 , 'batch_size' : 8 , 'nheads' : 16 , 'head_dim' : 128 , 'spec_len' : (2 , 3 ), 'paged_kv_block_size' : 16 },
286+ {'seqlen_q' : 1024 , 'seqlen_k' : 1024 , 'batch_size' : 8 , 'nheads' : 16 , 'head_dim' : 128 , 'spec_len' : (2 , 3 ), 'paged_kv_block_size' : 256 },
287+ {'seqlen_q' : 1024 , 'seqlen_k' : 1024 , 'batch_size' : 8 , 'nheads' : 16 , 'head_dim' : 128 , 'spec_len' : (2 , 3 ), 'paged_kv_block_size' : 512 },
288+ ]
289+
290+ print ("=" * 80 )
291+ print ("COMPREHENSIVE TREE ATTENTION BENCHMARK" )
292+ print ("=" * 80 )
293+
294+ results = []
295+ for i , config in enumerate (configs ):
296+ print (f"\n [{ i + 1 } /{ len (configs )} ] Configuration:" )
297+ result = run_tree_attention_benchmark (** config )
298+ results .append (result )
299+ print ("-" * 80 )
300+
301+ # Summary
302+ print ("\n " + "=" * 80 )
303+ print ("BENCHMARK SUMMARY" )
304+ print ("=" * 80 )
305+ print (f"{ 'Config' :<18} { 'Tree(ms)' :<10} { 'Varlen(ms)' :<12} { 'Speedup' :<10} { 'Max Diff' :<12} " )
306+ print ("-" * 80 )
307+
308+ for i , result in enumerate (results ):
309+ config = result ['config' ]
310+ config_str = f"{ config ['seqlen_q' ]} :{ config ['seqlen_k' ]} :{ config ['tree_mask.shape' ][0 ]} :{ config ['paged_kv_block_size' ]} "
311+ tree_ms = result ['tree_time' ] * 1000
312+ varlen_ms = result ['varlen_time' ] * 1000
313+ speedup = result ['speedup' ]
314+ max_diff = result ['max_diff' ]
315+
316+ print (f"{ config_str :<18} { tree_ms :<10.3f} { varlen_ms :<12.3f} { speedup :<10.2f} x { max_diff :<12.6f} " )
317+
318+ return results
319+
320+
321+ if __name__ == "__main__" :
322+ if not torch .cuda .is_available ():
323+ print ("CUDA is not available. This benchmark requires GPU." )
324+ exit (1 )
325+
326+ print ("Tree Attention vs Flash Attention Varlen Benchmark" )
327+ print (f"PyTorch version: { torch .__version__ } " )
328+ print (f"CUDA version: { torch .version .cuda } " )
329+ print (f"Device: { torch .cuda .get_device_name ()} " )
330+
331+ # Run single benchmark
332+ print ("\n " + "=" * 80 )
333+ print ("SINGLE BENCHMARK (1024x1024, batch=8)" )
334+ print ("=" * 80 )
335+ run_tree_attention_benchmark ()
336+
337+ # Run decoding benchmark
338+ run_decoding_benchmark ()
339+
340+ # Run comprehensive benchmark
341+ run_comprehensive_benchmark ()
0 commit comments