File size: 5,357 Bytes
eafbe80 | 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 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 |
import os
import torch
import triton
from torch.nn import functional as F
from fla.ops.gla import chunk_gla
from fla.ops.retention import chunk_retention
from fla.ops.rwkv6 import chunk_rwkv6
from fla.ops.rwkv7 import chunk_rwkv7
try:
from flash_attn import flash_attn_func
HAS_FLASH = True
except BaseException:
HAS_FLASH = False
@triton.testing.perf_report(
triton.testing.Benchmark(
# argument names to use as an x-axis for the plot
x_names=['T'],
# different possible values for `x_name`
x_vals=[128 * 2 ** i for i in range(0, 8)],
# argument name whose value corresponds to a different line in the plot
line_arg='provider',
# possible values for `line_arg``
# line styles
line_vals=['rwkv6', 'rwkv7', 'gla', 'flash', 'rwkv6_bwd', 'rwkv7_bwd', 'gla_bwd', 'retention_bwd', 'flash_bwd'],
# label name for the lines
line_names=['rwkv6', 'rwkv7', 'gla', 'flash', 'rwkv6_bwd', 'rwkv7_bwd', 'gla_bwd', 'retention_bwd', 'flash_bwd'],
# # line styles
styles=[
('green', '-'), # rwkv6
('blue', '--'), # rwkv7
('red', '-.'), # gla
('cyan', ':'), # rwkv6_bwd
('magenta', '-'), # rwkv7_bwd
('yellow', 'dotted'), # gla_bwd
('black', ':'), # retention_bwd
('gray', ':'), # flash
('gray', '--'), # flash_bwd
],
ylabel="Execution Time (ms)", # label name for the y-axis
# name for the plot. Used also as a file name for saving the plot.
plot_name="Performance",
args={},
),
)
def benchmark(T, provider):
from fla.utils import device
dtype = torch.bfloat16
requires_grad = True
# Read B, H, D from environment variables, default to 16, 8, 128 if not set
B = int(os.getenv('BENCH_B', '8')) # Batch size
H = int(os.getenv('BENCH_H', '64')) # Number of heads
D = int(os.getenv('BENCH_D', '64')) # Dimension per head
with torch.no_grad():
q = torch.randn(B, T, H, D, device=device, requires_grad=requires_grad, dtype=dtype)
k = torch.randn(B, T, H, D, device=device, requires_grad=requires_grad, dtype=dtype)
v = torch.randn(B, T, H, D, device=device, requires_grad=requires_grad, dtype=dtype)
if provider.startswith('flash'):
q = torch.randn(B, T, H, D, device=device, requires_grad=requires_grad, dtype=dtype)
k = torch.randn(B, T, H, D, device=device, requires_grad=requires_grad, dtype=dtype)
v = torch.randn(B, T, H, D, device=device, requires_grad=requires_grad, dtype=dtype)
if provider.startswith('gla'):
g = F.logsigmoid(torch.randn(B, T, H, D, device=device, dtype=dtype))
g = g.clamp_min(-5).requires_grad_(requires_grad)
if provider.startswith('rwkv6'):
w = F.logsigmoid(torch.randn(B, T, H, D, device=device, dtype=dtype)).requires_grad_(True)
u = torch.randn(H, D, device=device, dtype=dtype).requires_grad_(True)
if provider.startswith('rwkv7'):
q = torch.empty(B, T, H, D, device=device).uniform_(-1, 1).to(dtype=dtype).requires_grad_(True)
k = torch.empty(B, T, H, D, device=device).uniform_(-1, 1).to(dtype=dtype).requires_grad_(True)
v = torch.empty(B, T, H, D, device=device).uniform_(-1, 1).to(dtype=dtype).requires_grad_(True)
w = F.logsigmoid(torch.randn(B, T, H, D, device=device, dtype=dtype)).requires_grad_(True)
kk = torch.empty(B, T, H, D, device=device).uniform_(-1, 1)
kk = torch.nn.functional.normalize(kk, dim=-1).to(dtype=dtype)
a = -kk.clone().requires_grad_(True) # -kk
a_scale = torch.empty(B, T, H, D, device=device).uniform_(0, 0.1).to(dtype=dtype)
b = (kk * a_scale).requires_grad_(True) # kk*a
do = torch.ones_like(v, dtype=dtype)
quantiles = [0.5, 0.2, 0.8]
if provider == 'rwkv6':
results = triton.testing.do_bench(lambda: chunk_rwkv6(q, k, v, w, u), quantiles=quantiles)
elif provider == 'rwkv7':
results = triton.testing.do_bench(lambda: chunk_rwkv7(q, w, k, v, a, b), quantiles=quantiles)
elif provider == 'gla':
results = triton.testing.do_bench(lambda: chunk_gla(q, k, v, g), quantiles=quantiles)
elif provider == 'rwkv6_bwd':
results = triton.testing.do_bench(lambda: chunk_rwkv6(q, k, v, w, u)[0].backward(do), quantiles=quantiles)
elif provider == 'rwkv7_bwd':
results = triton.testing.do_bench(lambda: chunk_rwkv7(q, w, k, v, a, b)[0].backward(do), quantiles=quantiles)
elif provider == 'gla_bwd':
results = triton.testing.do_bench(lambda: chunk_gla(q, k, v, g)[0].backward(do), quantiles=quantiles)
elif provider == 'retention_bwd':
results = triton.testing.do_bench(lambda: chunk_retention(q, k, v)[0].backward(do), quantiles=quantiles)
elif provider == 'flash':
results = triton.testing.do_bench(lambda: flash_attn_func(q, k, v, causal=True), quantiles=quantiles)
elif provider == 'flash_bwd':
results = triton.testing.do_bench(lambda: flash_attn_func(q, k, v, causal=True).backward(do), quantiles=quantiles)
return results
if __name__ == '__main__':
benchmark.run(print_data=True)
|