Skip to content

Commit 753db41

Browse files
committed
feat: implement tree attention mask support for FlashAttention-2
Add comprehensive tree attention implementation including: - Core tree attention CUDA kernels for multiple head dimensions (32-256) and precisions (FP16/BF16) - Tree mask utilities for structured attention patterns in speculative decoding - Python interface and C++ bindings for tree_attention function - Benchmarking suite comparing tree attention vs varlen flash attention - Test utilities for paged KV cache with tree attention patterns This enables efficient speculative decoding by avoiding batch expansion, providing memory-efficient attention computation for tree-structured token generation patterns commonly used in speculative sampling. Signed-off-by: a <a.oneill@samsung.com> Signed-off-by: Andrew O'Neill <a.oneill@samsung.com>
1 parent 57b4e68 commit 753db41

44 files changed

Lines changed: 3719 additions & 1 deletion

Some content is hidden

Large Commits have some content hidden by default. Use the searchbox below for content that may be hidden.

CMakeLists.txt

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -150,6 +150,7 @@ if (FA2_ENABLED)
150150
SOURCES
151151
csrc/flash_attn/flash_api.cpp
152152
csrc/flash_attn/flash_api_sparse.cpp
153+
csrc/flash_attn/tree_attention.cpp
153154
csrc/flash_attn/flash_api_torch_lib.cpp
154155
${FA2_GEN_SRCS}
155156
COMPILE_FLAGS ${VLLM_FA_GPU_FLAGS}
Lines changed: 341 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,341 @@
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("\nRunning 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("\nVerifying 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()

csrc/flash_attn/flash_api_torch_lib.cpp

Lines changed: 32 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -101,6 +101,31 @@ mha_varlen_fwd_sparse(at::Tensor &q, // total_q x num_heads x head_size, total_
101101
const bool return_softmax,
102102
std::optional<at::Generator> gen_);
103103

104+
/////////////////////////// From flash_api_tree.cpp //////////////////////////
105+
106+
std::vector<at::Tensor>
107+
tree_attention(at::Tensor &q, // total_q x num_heads x head_size, total_q := \sum_{i=0}^{b} s_i
108+
const at::Tensor &k, // total_k x num_heads_k x head_size, total_k := \sum_{i=0}^{b} s_i or num_blocks x page_block_size x num_heads_k x head_size if there's a block_table.
109+
const at::Tensor &v, // total_k x num_heads_k x head_size, total_k := \sum_{i=0}^{b} s_i or num_blocks x page_block_size x num_heads_k x head_size if there's a block_table.
110+
std::optional<at::Tensor> &out_, // total_q x num_heads x head_size, total_k := \sum_{i=0}^{b} s_i
111+
const at::Tensor &cu_seqlens_q, // b+1
112+
const at::Tensor &cu_seqlens_k, // b+1
113+
std::optional<at::Tensor> &seqused_k, // b. If given, only this many elements of each batch element's keys are used.
114+
std::optional<const at::Tensor> &leftpad_k_, // batch_size
115+
std::optional<at::Tensor> &block_table_, // batch_size x max_num_blocks_per_seq
116+
std::optional<at::Tensor> &alibi_slopes_, // num_heads or b x num_heads
117+
int max_seqlen_q,
118+
const int max_seqlen_k,
119+
const float p_dropout,
120+
const float softmax_scale,
121+
const bool zero_tensors,
122+
const float softcap,
123+
const bool return_softmax,
124+
std::optional<at::Generator> gen_,
125+
const at::Tensor &tree_mask,
126+
const at::Tensor &tree_mask_lens);
127+
128+
104129
/**
105130
* Torch Library Registration
106131
*/
@@ -134,6 +159,13 @@ TORCH_LIBRARY_EXPAND(TORCH_EXTENSION_NAME, ops) {
134159
"bool is_causal, float softcap, bool return_softmax, "
135160
"Generator? gen) -> Tensor[]");
136161
ops.impl("varlen_fwd_sparse", torch::kCUDA, &mha_varlen_fwd_sparse);
162+
163+
ops.def("tree_attention(Tensor! q, Tensor k, Tensor v, Tensor!? out, Tensor cu_seqlens_q, "
164+
"Tensor cu_seqlens_k, Tensor? seqused_k, Tensor? leftpad_k, Tensor? block_table, Tensor? alibi_slopes, "
165+
"int max_seqlen_q, int max_seqlen_k, float p_dropout, float softmax_scale, bool zero_tensors, "
166+
"float softcap, bool return_softmax, "
167+
"Generator? gen, Tensor tree_mask, Tensor tree_mask_lens) -> Tensor[]");
168+
ops.impl("tree_attention", torch::kCUDA, make_pytorch_shim(&tree_attention));
137169
}
138170

139171
REGISTER_EXTENSION(TORCH_EXTENSION_NAME);

0 commit comments

Comments
 (0)