forked from flashinfer-ai/flashinfer-bench-starter-kit
-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtest_kernel.py
More file actions
38 lines (32 loc) · 1.57 KB
/
Copy pathtest_kernel.py
File metadata and controls
38 lines (32 loc) · 1.57 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
#!/usr/bin/env python3
"""Test script to debug kernel compilation errors."""
import torch
import sys
sys.path.insert(0, 'solution/triton')
from kernel import gdn_decode
# Create dummy tensors for testing - GVA configuration
batch_size, seq_len = 1, 1
num_q_heads, num_k_heads, num_v_heads = 4, 4, 8 # GVA: more v heads than q/k heads
head_size = 128
device = 'cuda'
print(f"Creating test tensors: batch={batch_size}, seq={seq_len}")
print(f" q/k heads={num_q_heads}, v heads={num_v_heads}, head_size={head_size}")
q = torch.randn(batch_size, seq_len, num_q_heads, head_size, dtype=torch.bfloat16, device=device)
k = torch.randn(batch_size, seq_len, num_k_heads, head_size, dtype=torch.bfloat16, device=device)
v = torch.randn(batch_size, seq_len, num_v_heads, head_size, dtype=torch.bfloat16, device=device)
state = torch.randn(batch_size, num_v_heads, head_size, head_size, dtype=torch.float32, device=device)
A_log = torch.randn(num_v_heads, dtype=torch.float32, device=device)
a = torch.randn(batch_size, seq_len, num_v_heads, dtype=torch.bfloat16, device=device)
dt_bias = torch.randn(num_v_heads, dtype=torch.float32, device=device)
b = torch.randn(batch_size, seq_len, num_v_heads, dtype=torch.bfloat16, device=device)
scale = 0.08838834764831843
print('Running gdn_decode...')
try:
output, new_state = gdn_decode(q, k, v, state, A_log, a, dt_bias, b, scale)
print(f'Output shape: {output.shape}')
print(f'New state shape: {new_state.shape}')
print('Success!')
except Exception as e:
print(f'Error: {type(e).__name__}: {e}')
import traceback
traceback.print_exc()