forked from metax-maca/op_optimization
add .gitignore & delete temp file
This commit is contained in:
parent
9828f90095
commit
97c91e6a48
|
|
@ -0,0 +1,79 @@
|
|||
# =========================
|
||||
# macOS
|
||||
# =========================
|
||||
.DS_Store
|
||||
.AppleDouble
|
||||
.LSOverride
|
||||
Icon?
|
||||
._*
|
||||
.Spotlight-V100
|
||||
.Trashes
|
||||
.fseventsd
|
||||
|
||||
# =========================
|
||||
# IDE / Editor
|
||||
# =========================
|
||||
.vscode/
|
||||
.idea/
|
||||
*.swp
|
||||
*.swo
|
||||
*~
|
||||
|
||||
# =========================
|
||||
# Python
|
||||
# =========================
|
||||
__pycache__/
|
||||
*.py[cod]
|
||||
*.pyo
|
||||
*.pyd
|
||||
.pytest_cache/
|
||||
.mypy_cache/
|
||||
.ruff_cache/
|
||||
.coverage
|
||||
htmlcov/
|
||||
.env
|
||||
.venv/
|
||||
venv/
|
||||
env/
|
||||
|
||||
# =========================
|
||||
# Jupyter
|
||||
# =========================
|
||||
.ipynb_checkpoints/
|
||||
|
||||
# =========================
|
||||
# Build / Packaging
|
||||
# =========================
|
||||
build/
|
||||
dist/
|
||||
*.egg-info/
|
||||
.eggs/
|
||||
pip-wheel-metadata/
|
||||
|
||||
# =========================
|
||||
# C / C++ / CMake
|
||||
# =========================
|
||||
CMakeFiles/
|
||||
CMakeCache.txt
|
||||
cmake-build-*/
|
||||
Makefile
|
||||
*.o
|
||||
*.so
|
||||
*.dylib
|
||||
*.dll
|
||||
*.a
|
||||
*.lib
|
||||
|
||||
# =========================
|
||||
# Logs / Temp
|
||||
# =========================
|
||||
*.log
|
||||
*.tmp
|
||||
*.temp
|
||||
logs/
|
||||
tmp/
|
||||
|
||||
# =========================
|
||||
# OS / Tool caches
|
||||
# =========================
|
||||
.cache/
|
||||
|
|
@ -1,49 +0,0 @@
|
|||
batch_size,seq_len_kv,heads,headdim,time_ms,bandwidth_GB_s
|
||||
1,512,8,128,0.0322,65.27
|
||||
2,512,8,128,0.0324,129.61
|
||||
4,512,8,128,0.0332,253.01
|
||||
8,512,8,128,0.0355,472.66
|
||||
16,512,8,128,0.0546,614.98
|
||||
32,512,8,128,0.0817,822.60
|
||||
64,512,8,128,0.1297,1035.62
|
||||
128,512,8,128,0.2453,1095.50
|
||||
1,1024,8,128,0.0578,72.55
|
||||
2,1024,8,128,0.0586,143.14
|
||||
4,1024,8,128,0.0597,281.24
|
||||
8,1024,8,128,0.0625,536.80
|
||||
16,1024,8,128,0.0982,683.59
|
||||
32,1024,8,128,0.1493,899.64
|
||||
64,1024,8,128,0.2403,1117.62
|
||||
128,1024,8,128,0.4594,1169.27
|
||||
1,2048,8,128,0.1101,76.24
|
||||
2,2048,8,128,0.1107,151.64
|
||||
4,2048,8,128,0.1119,299.88
|
||||
8,2048,8,128,0.1159,578.98
|
||||
16,2048,8,128,0.1849,726.23
|
||||
32,2048,8,128,0.2843,944.47
|
||||
64,2048,8,128,0.4607,1165.56
|
||||
128,2048,8,128,0.8868,1211.07
|
||||
1,4096,8,128,0.2139,78.46
|
||||
2,4096,8,128,0.2151,156.01
|
||||
4,4096,8,128,0.2163,310.36
|
||||
8,4096,8,128,0.2227,602.81
|
||||
16,4096,8,128,0.3574,751.13
|
||||
32,4096,8,128,0.5540,969.23
|
||||
64,4096,8,128,0.9016,1191.07
|
||||
128,4096,8,128,1.7414,1233.34
|
||||
1,8192,8,128,0.4215,79.61
|
||||
2,8192,8,128,0.4226,158.81
|
||||
4,8192,8,128,0.4242,316.39
|
||||
8,8192,8,128,0.4362,615.46
|
||||
16,8192,8,128,0.7035,763.14
|
||||
32,8192,8,128,1.0934,982.11
|
||||
64,8192,8,128,1.7814,1205.57
|
||||
128,8192,8,128,3.4505,1244.82
|
||||
1,16384,8,128,0.8356,80.32
|
||||
2,16384,8,128,0.8377,160.23
|
||||
4,16384,8,128,0.8407,319.30
|
||||
8,16384,8,128,0.8625,622.51
|
||||
16,16384,8,128,1.3934,770.60
|
||||
32,16384,8,128,2.1695,989.88
|
||||
64,16384,8,128,3.5397,1213.41
|
||||
128,16384,8,128,6.8668,1250.98
|
||||
|
|
|
@ -1,49 +0,0 @@
|
|||
batch_size,seq_len_kv,heads,headdim,time_ms,bandwidth_GB_s
|
||||
1,512,8,160,0.0580,45.25
|
||||
2,512,8,160,0.0611,85.84
|
||||
4,512,8,160,0.0656,159.96
|
||||
8,512,8,160,0.0699,300.40
|
||||
16,512,8,160,0.1321,317.92
|
||||
32,512,8,160,0.2002,419.52
|
||||
64,512,8,160,0.3383,496.43
|
||||
128,512,8,160,0.6669,503.61
|
||||
1,1024,8,160,0.1129,46.45
|
||||
2,1024,8,160,0.1190,88.18
|
||||
4,1024,8,160,0.1224,171.49
|
||||
8,1024,8,160,0.1287,326.07
|
||||
16,1024,8,160,0.2479,338.49
|
||||
32,1024,8,160,0.3767,445.54
|
||||
64,1024,8,160,0.6419,523.01
|
||||
128,1024,8,160,1.2804,524.37
|
||||
1,2048,8,160,0.2270,46.20
|
||||
2,2048,8,160,0.2299,91.22
|
||||
4,2048,8,160,0.2349,178.63
|
||||
8,2048,8,160,0.2447,342.96
|
||||
16,2048,8,160,0.4773,351.60
|
||||
32,2048,8,160,0.7279,461.07
|
||||
64,2048,8,160,1.2559,534.49
|
||||
128,2048,8,160,2.5613,524.15
|
||||
1,4096,8,160,0.4460,47.02
|
||||
2,4096,8,160,0.4513,92.94
|
||||
4,4096,8,160,0.4593,182.64
|
||||
8,4096,8,160,0.4813,348.64
|
||||
16,4096,8,160,0.9363,358.43
|
||||
32,4096,8,160,1.4552,461.21
|
||||
64,4096,8,160,2.5615,524.05
|
||||
128,4096,8,160,5.1420,522.11
|
||||
1,8192,8,160,0.8847,47.41
|
||||
2,8192,8,160,0.8944,93.80
|
||||
4,8192,8,160,0.9094,184.51
|
||||
8,8192,8,160,0.9625,348.64
|
||||
16,8192,8,160,1.8550,361.80
|
||||
32,8192,8,160,2.9567,453.97
|
||||
64,8192,8,160,5.1398,522.30
|
||||
128,8192,8,160,10.2972,521.41
|
||||
1,16384,8,160,1.7608,47.64
|
||||
2,16384,8,160,1.7786,94.33
|
||||
4,16384,8,160,1.8143,184.95
|
||||
8,16384,8,160,1.9317,347.42
|
||||
16,16384,8,160,3.7301,359.83
|
||||
32,16384,8,160,5.9216,453.33
|
||||
64,16384,8,160,10.2668,522.94
|
||||
128,16384,8,160,20.6062,521.09
|
||||
|
|
|
@ -1,49 +0,0 @@
|
|||
batch_size,seq_len_kv,heads,headdim,time_ms,bandwidth_GB_s
|
||||
1,512,8,192,0.0458,68.82
|
||||
2,512,8,192,0.0515,122.32
|
||||
4,512,8,192,0.0574,219.28
|
||||
8,512,8,192,0.0607,414.80
|
||||
16,512,8,192,0.1147,439.27
|
||||
32,512,8,192,0.1763,571.40
|
||||
64,512,8,192,0.2978,676.79
|
||||
128,512,8,192,0.5874,686.11
|
||||
1,1024,8,192,0.0946,66.55
|
||||
2,1024,8,192,0.1033,121.85
|
||||
4,1024,8,192,0.1073,234.66
|
||||
8,1024,8,192,0.1131,445.23
|
||||
16,1024,8,192,0.2165,465.24
|
||||
32,1024,8,192,0.3347,601.80
|
||||
64,1024,8,192,0.5701,706.63
|
||||
128,1024,8,192,1.1302,712.88
|
||||
1,2048,8,192,0.1943,64.79
|
||||
2,2048,8,192,0.1992,126.38
|
||||
4,2048,8,192,0.2059,244.52
|
||||
8,2048,8,192,0.2174,463.13
|
||||
16,2048,8,192,0.4202,479.24
|
||||
32,2048,8,192,0.6503,619.36
|
||||
64,2048,8,192,1.1158,721.93
|
||||
128,2048,8,192,2.2250,724.05
|
||||
1,4096,8,192,0.3834,65.65
|
||||
2,4096,8,192,0.3904,128.95
|
||||
4,4096,8,192,0.4043,249.04
|
||||
8,4096,8,192,0.4267,471.92
|
||||
16,4096,8,192,0.8271,486.90
|
||||
32,4096,8,192,1.2840,627.28
|
||||
64,4096,8,192,2.2148,727.29
|
||||
128,4096,8,192,4.3819,735.21
|
||||
1,8192,8,192,0.7566,66.52
|
||||
2,8192,8,192,0.7712,130.54
|
||||
4,8192,8,192,0.7974,252.49
|
||||
8,8192,8,192,0.8433,477.47
|
||||
16,8192,8,192,1.6433,490.09
|
||||
32,8192,8,192,2.5573,629.84
|
||||
64,8192,8,192,4.3785,735.73
|
||||
128,8192,8,192,8.7303,737.99
|
||||
1,16384,8,192,1.5068,66.81
|
||||
2,16384,8,192,1.5350,131.16
|
||||
4,16384,8,192,1.5868,253.76
|
||||
8,16384,8,192,1.6778,479.99
|
||||
16,16384,8,192,3.2750,491.81
|
||||
32,16384,8,192,5.0659,635.88
|
||||
64,16384,8,192,8.7435,736.85
|
||||
128,16384,8,192,17.5040,736.13
|
||||
|
|
|
@ -1,49 +0,0 @@
|
|||
batch_size,seq_len_kv,heads,headdim,time_ms,bandwidth_GB_s
|
||||
1,512,8,224,0.1254,29.29
|
||||
2,512,8,224,0.1412,52.05
|
||||
4,512,8,224,0.1497,98.13
|
||||
8,512,8,224,0.1533,191.70
|
||||
16,512,8,224,0.1913,307.17
|
||||
32,512,8,224,0.3292,357.08
|
||||
64,512,8,224,0.5187,453.28
|
||||
128,512,8,224,0.9522,493.84
|
||||
1,1024,8,224,0.2727,26.93
|
||||
2,1024,8,224,0.2836,51.78
|
||||
4,1024,8,224,0.2890,101.63
|
||||
8,1024,8,224,0.2959,198.55
|
||||
16,1024,8,224,0.3696,317.93
|
||||
32,1024,8,224,0.6408,366.75
|
||||
64,1024,8,224,1.0081,466.21
|
||||
128,1024,8,224,1.8548,506.78
|
||||
1,2048,8,224,0.5515,26.63
|
||||
2,2048,8,224,0.5575,52.67
|
||||
4,2048,8,224,0.5666,103.65
|
||||
8,2048,8,224,0.5803,202.42
|
||||
16,2048,8,224,0.7250,324.05
|
||||
32,2048,8,224,1.2593,373.14
|
||||
64,2048,8,224,1.9890,472.48
|
||||
128,2048,8,224,3.6905,509.28
|
||||
1,4096,8,224,1.0939,26.84
|
||||
2,4096,8,224,1.1044,53.18
|
||||
4,4096,8,224,1.1219,104.69
|
||||
8,4096,8,224,1.1500,204.26
|
||||
16,4096,8,224,1.4390,326.48
|
||||
32,4096,8,224,2.4992,375.97
|
||||
64,4096,8,224,4.0082,468.86
|
||||
128,4096,8,224,7.3372,512.26
|
||||
1,8192,8,224,2.1775,26.97
|
||||
2,8192,8,224,2.1989,53.41
|
||||
4,8192,8,224,2.2338,105.15
|
||||
8,8192,8,224,2.3268,201.90
|
||||
16,8192,8,224,2.8806,326.18
|
||||
32,8192,8,224,5.0187,374.43
|
||||
64,8192,8,224,8.0323,467.90
|
||||
128,8192,8,224,14.6300,513.78
|
||||
1,16384,8,224,4.3360,27.09
|
||||
2,16384,8,224,4.3820,53.60
|
||||
4,16384,8,224,4.5006,104.38
|
||||
8,16384,8,224,4.6987,199.96
|
||||
16,16384,8,224,5.7361,327.59
|
||||
32,16384,8,224,10.1291,371.03
|
||||
64,16384,8,224,16.0745,467.60
|
||||
128,16384,8,224,OOM,OOM
|
||||
|
|
|
@ -1,49 +0,0 @@
|
|||
batch_size,seq_len_kv,heads,headdim,time_ms,bandwidth_GB_s
|
||||
1,512,8,256,0.0877,47.89
|
||||
2,512,8,256,0.0921,91.17
|
||||
4,512,8,256,0.0940,178.74
|
||||
8,512,8,256,0.0964,348.52
|
||||
16,512,8,256,0.1450,463.27
|
||||
32,512,8,256,0.2250,597.21
|
||||
64,512,8,256,0.3609,744.43
|
||||
128,512,8,256,0.6932,775.25
|
||||
1,1024,8,256,0.1747,48.04
|
||||
2,1024,8,256,0.1762,95.27
|
||||
4,1024,8,256,0.1784,188.22
|
||||
8,1024,8,256,0.1817,369.53
|
||||
16,1024,8,256,0.2796,480.25
|
||||
32,1024,8,256,0.4339,619.00
|
||||
64,1024,8,256,0.6960,771.73
|
||||
128,1024,8,256,1.3439,799.36
|
||||
1,2048,8,256,0.3410,49.21
|
||||
2,2048,8,256,0.3439,97.60
|
||||
4,2048,8,256,0.3469,193.52
|
||||
8,2048,8,256,0.3533,379.94
|
||||
16,2048,8,256,0.5461,491.67
|
||||
32,2048,8,256,0.8493,632.28
|
||||
64,2048,8,256,1.3667,785.82
|
||||
128,2048,8,256,2.6465,811.64
|
||||
1,4096,8,256,0.6742,49.77
|
||||
2,4096,8,256,0.6777,99.03
|
||||
4,4096,8,256,0.6836,196.36
|
||||
8,4096,8,256,0.6950,386.31
|
||||
16,4096,8,256,1.0803,497.02
|
||||
32,4096,8,256,1.6794,639.44
|
||||
64,4096,8,256,2.7101,792.50
|
||||
128,4096,8,256,5.2543,817.52
|
||||
1,8192,8,256,1.3375,50.18
|
||||
2,8192,8,256,1.3448,99.81
|
||||
4,8192,8,256,1.3564,197.91
|
||||
8,8192,8,256,1.3799,389.08
|
||||
16,8192,8,256,2.1465,500.25
|
||||
32,8192,8,256,3.3342,644.12
|
||||
64,8192,8,256,5.3983,795.67
|
||||
128,8192,8,256,10.4691,820.55
|
||||
1,16384,8,256,2.6697,50.28
|
||||
2,16384,8,256,2.6817,100.10
|
||||
4,16384,8,256,2.7049,198.49
|
||||
8,16384,8,256,2.7533,390.00
|
||||
16,16384,8,256,4.2789,501.89
|
||||
32,16384,8,256,6.6476,646.11
|
||||
64,16384,8,256,10.7723,797.43
|
||||
128,16384,8,256,OOM,OOM
|
||||
|
|
|
@ -1,49 +0,0 @@
|
|||
batch_size,seq_len_kv,heads,headdim,time_ms,bandwidth_GB_s
|
||||
1,512,8,32,0.0257,20.45
|
||||
2,512,8,32,0.0256,41.02
|
||||
4,512,8,32,0.0258,81.28
|
||||
8,512,8,32,0.0265,158.45
|
||||
16,512,8,32,0.0396,212.30
|
||||
32,512,8,32,0.0516,325.43
|
||||
64,512,8,32,0.0721,465.83
|
||||
128,512,8,32,0.1270,529.03
|
||||
1,1024,8,32,0.0461,22.75
|
||||
2,1024,8,32,0.0465,45.15
|
||||
4,1024,8,32,0.0477,88.04
|
||||
8,1024,8,32,0.0548,153.23
|
||||
16,1024,8,32,0.0734,228.71
|
||||
32,1024,8,32,0.0958,350.42
|
||||
64,1024,8,32,0.1334,503.15
|
||||
128,1024,8,32,0.2381,564.04
|
||||
1,2048,8,32,0.0872,24.06
|
||||
2,2048,8,32,0.0904,46.42
|
||||
4,2048,8,32,0.1028,81.59
|
||||
8,2048,8,32,0.1067,157.25
|
||||
16,2048,8,32,0.1428,235.10
|
||||
32,2048,8,32,0.1818,369.13
|
||||
64,2048,8,32,0.2554,525.57
|
||||
128,2048,8,32,0.4622,580.86
|
||||
1,4096,8,32,0.1730,24.25
|
||||
2,4096,8,32,0.1955,42.91
|
||||
4,4096,8,32,0.2020,83.05
|
||||
8,4096,8,32,0.2140,156.83
|
||||
16,4096,8,32,0.2777,241.65
|
||||
32,4096,8,32,0.3542,378.99
|
||||
64,4096,8,32,0.4990,538.05
|
||||
128,4096,8,32,0.9099,590.13
|
||||
1,8192,8,32,0.3820,21.96
|
||||
2,8192,8,32,0.3913,42.88
|
||||
4,8192,8,32,0.4127,81.31
|
||||
8,8192,8,32,0.4224,158.88
|
||||
16,8192,8,32,0.5490,244.51
|
||||
32,8192,8,32,0.6960,385.70
|
||||
64,8192,8,32,0.9870,543.98
|
||||
128,8192,8,32,1.8100,593.25
|
||||
1,16384,8,32,0.7655,21.92
|
||||
2,16384,8,32,0.8067,41.59
|
||||
4,16384,8,32,0.8228,81.56
|
||||
8,16384,8,32,0.8397,159.85
|
||||
16,16384,8,32,1.0910,246.04
|
||||
32,16384,8,32,1.3824,388.37
|
||||
64,16384,8,32,1.9663,546.08
|
||||
128,16384,8,32,3.6107,594.78
|
||||
|
|
|
@ -1,49 +0,0 @@
|
|||
batch_size,seq_len_kv,heads,headdim,time_ms,bandwidth_GB_s
|
||||
1,512,8,512,0.3588,23.40
|
||||
2,512,8,512,0.3651,46.00
|
||||
4,512,8,512,0.3736,89.89
|
||||
8,512,8,512,0.3856,174.22
|
||||
16,512,8,512,0.7472,179.80
|
||||
32,512,8,512,1.1447,234.72
|
||||
64,512,8,512,1.9549,274.89
|
||||
128,512,8,512,3.8962,275.85
|
||||
1,1024,8,512,0.7261,23.12
|
||||
2,1024,8,512,0.7354,45.65
|
||||
4,1024,8,512,0.7496,89.57
|
||||
8,1024,8,512,0.7746,173.35
|
||||
16,1024,8,512,1.5049,178.46
|
||||
32,1024,8,512,2.3111,232.42
|
||||
64,1024,8,512,3.9538,271.70
|
||||
128,1024,8,512,7.8811,272.62
|
||||
1,2048,8,512,1.4636,22.93
|
||||
2,2048,8,512,1.4826,45.27
|
||||
4,2048,8,512,1.5109,88.86
|
||||
8,2048,8,512,1.5549,172.68
|
||||
16,2048,8,512,3.0237,177.60
|
||||
32,2048,8,512,4.6439,231.27
|
||||
64,2048,8,512,7.9560,269.99
|
||||
128,2048,8,512,15.8741,270.63
|
||||
1,4096,8,512,2.9312,22.90
|
||||
2,4096,8,512,2.9675,45.24
|
||||
4,4096,8,512,3.0243,88.77
|
||||
8,4096,8,512,3.1127,172.50
|
||||
16,4096,8,512,6.0753,176.76
|
||||
32,4096,8,512,9.3182,230.49
|
||||
64,4096,8,512,15.9642,269.07
|
||||
128,4096,8,512,31.8313,269.89
|
||||
1,8192,8,512,5.8843,22.81
|
||||
2,8192,8,512,5.9344,45.24
|
||||
4,8192,8,512,6.0465,88.80
|
||||
8,8192,8,512,6.2334,172.27
|
||||
16,8192,8,512,12.1594,176.62
|
||||
32,8192,8,512,18.6826,229.90
|
||||
64,8192,8,512,32.0055,268.41
|
||||
128,8192,8,512,OOM,OOM
|
||||
1,16384,8,512,11.8153,22.72
|
||||
2,16384,8,512,11.9237,45.03
|
||||
4,16384,8,512,12.1671,88.25
|
||||
8,16384,8,512,12.4948,171.88
|
||||
16,16384,8,512,24.3414,176.45
|
||||
32,16384,8,512,37.3907,229.74
|
||||
64,16384,8,512,OOM,OOM
|
||||
128,16384,8,512,OOM,OOM
|
||||
|
|
|
@ -1,49 +0,0 @@
|
|||
batch_size,seq_len_kv,heads,headdim,time_ms,bandwidth_GB_s
|
||||
1,512,8,64,0.0404,25.99
|
||||
2,512,8,64,0.0399,52.60
|
||||
4,512,8,64,0.0413,101.69
|
||||
8,512,8,64,0.0482,174.25
|
||||
16,512,8,64,0.0540,310.86
|
||||
32,512,8,64,0.0629,533.75
|
||||
64,512,8,64,0.0833,806.14
|
||||
128,512,8,64,0.1104,1216.59
|
||||
1,1024,8,64,0.0747,28.08
|
||||
2,1024,8,64,0.0766,54.77
|
||||
4,1024,8,64,0.0891,94.17
|
||||
8,1024,8,64,0.0918,182.94
|
||||
16,1024,8,64,0.1044,321.41
|
||||
32,1024,8,64,0.1179,569.43
|
||||
64,1024,8,64,0.1566,857.28
|
||||
128,1024,8,64,0.2078,1292.17
|
||||
1,2048,8,64,0.1455,28.84
|
||||
2,2048,8,64,0.1684,49.82
|
||||
4,2048,8,64,0.1730,97.01
|
||||
8,2048,8,64,0.1850,181.39
|
||||
16,2048,8,64,0.2009,334.18
|
||||
32,2048,8,64,0.2268,592.01
|
||||
64,2048,8,64,0.3002,894.44
|
||||
128,2048,8,64,0.4027,1333.64
|
||||
1,4096,8,64,0.3265,25.69
|
||||
2,4096,8,64,0.3322,50.51
|
||||
4,4096,8,64,0.3522,95.27
|
||||
8,4096,8,64,0.3632,184.79
|
||||
16,4096,8,64,0.3942,340.56
|
||||
32,4096,8,64,0.4456,602.47
|
||||
64,4096,8,64,0.5927,905.94
|
||||
128,4096,8,64,0.7938,1352.87
|
||||
1,8192,8,64,0.6508,25.78
|
||||
2,8192,8,64,0.6879,48.78
|
||||
4,8192,8,64,0.7008,95.77
|
||||
8,8192,8,64,0.7199,186.44
|
||||
16,8192,8,64,0.7786,344.79
|
||||
32,8192,8,64,0.8798,610.25
|
||||
64,8192,8,64,1.1745,914.30
|
||||
128,8192,8,64,1.5728,1365.50
|
||||
1,16384,8,64,1.3524,24.81
|
||||
2,16384,8,64,1.3698,48.99
|
||||
4,16384,8,64,1.3923,96.40
|
||||
8,16384,8,64,1.4267,188.16
|
||||
16,16384,8,64,1.5451,347.47
|
||||
32,16384,8,64,1.7622,609.32
|
||||
64,16384,8,64,2.3392,918.09
|
||||
128,16384,8,64,3.1332,1370.84
|
||||
|
|
|
@ -1,49 +0,0 @@
|
|||
batch_size,seq_len_kv,heads,headdim,time_ms,bandwidth_GB_s
|
||||
1,512,8,96,0.0407,38.67
|
||||
2,512,8,96,0.0398,79.02
|
||||
4,512,8,96,0.0431,146.08
|
||||
8,512,8,96,0.0495,254.61
|
||||
16,512,8,96,0.0698,360.64
|
||||
32,512,8,96,0.1117,450.87
|
||||
64,512,8,96,0.1780,566.16
|
||||
128,512,8,96,0.3329,605.28
|
||||
1,1024,8,96,0.0732,43.01
|
||||
2,1024,8,96,0.0794,79.29
|
||||
4,1024,8,96,0.0871,144.54
|
||||
8,1024,8,96,0.0934,269.52
|
||||
16,1024,8,96,0.1297,388.14
|
||||
32,1024,8,96,0.2114,476.36
|
||||
64,1024,8,96,0.3379,596.08
|
||||
128,1024,8,96,0.6327,636.68
|
||||
1,2048,8,96,0.1505,41.80
|
||||
2,2048,8,96,0.1619,77.76
|
||||
4,2048,8,96,0.1713,146.94
|
||||
8,2048,8,96,0.1780,282.84
|
||||
16,2048,8,96,0.2492,404.09
|
||||
32,2048,8,96,0.4088,492.55
|
||||
64,2048,8,96,0.6575,612.55
|
||||
128,2048,8,96,1.2457,646.61
|
||||
1,4096,8,96,0.3099,40.61
|
||||
2,4096,8,96,0.3259,77.23
|
||||
4,4096,8,96,0.3346,150.42
|
||||
8,4096,8,96,0.3467,290.41
|
||||
16,4096,8,96,0.4888,411.94
|
||||
32,4096,8,96,0.8055,499.94
|
||||
64,4096,8,96,1.3209,609.72
|
||||
128,4096,8,96,2.4810,649.25
|
||||
1,8192,8,96,0.6343,39.68
|
||||
2,8192,8,96,0.6437,78.20
|
||||
4,8192,8,96,0.6601,152.50
|
||||
8,8192,8,96,0.6826,294.97
|
||||
16,8192,8,96,0.9688,415.64
|
||||
32,8192,8,96,1.6057,501.55
|
||||
64,8192,8,96,2.6527,607.19
|
||||
128,8192,8,96,4.9464,651.27
|
||||
1,16384,8,96,1.2581,40.01
|
||||
2,16384,8,96,1.2812,78.57
|
||||
4,16384,8,96,1.3112,153.55
|
||||
8,16384,8,96,1.3653,294.92
|
||||
16,16384,8,96,1.9351,416.16
|
||||
32,16384,8,96,3.2277,499.01
|
||||
64,16384,8,96,5.3192,605.60
|
||||
128,16384,8,96,9.8747,652.44
|
||||
|
|
|
@ -1,16 +0,0 @@
|
|||
{
|
||||
"id": 201,
|
||||
"displayId": 10008,
|
||||
"type": "Traditional",
|
||||
"isPublic": false,
|
||||
"locales": [
|
||||
"zh_CN"
|
||||
],
|
||||
"samples": [
|
||||
{
|
||||
"inputData": "1\n",
|
||||
"outputData": ""
|
||||
}
|
||||
],
|
||||
"problemTagIds": []
|
||||
}
|
||||
|
|
@ -1,207 +0,0 @@
|
|||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
"""
|
||||
SPJ: problem_10008 — Paged KV Cache (flash-attn)
|
||||
"""
|
||||
|
||||
import base64
|
||||
import hashlib
|
||||
import json
|
||||
import math
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from typing import Callable, List, Optional, Tuple
|
||||
|
||||
|
||||
CHAL_PREFIX = "OJCHAL v1 "
|
||||
RESULT_PREFIX = "OJRESULT v1 "
|
||||
|
||||
|
||||
def eprint(msg: str) -> None:
|
||||
print(msg, file=sys.stderr, flush=True)
|
||||
|
||||
|
||||
def read_text(path: Path) -> str:
|
||||
try:
|
||||
return path.read_text(encoding="utf-8", errors="replace")
|
||||
except Exception as ex:
|
||||
raise RuntimeError(f"Failed to read {path}: {ex}") from ex
|
||||
|
||||
|
||||
def parse_nonce_from_text(text: str) -> Optional[bytes]:
|
||||
for line in text.splitlines():
|
||||
line = line.strip()
|
||||
if line.startswith(CHAL_PREFIX):
|
||||
parts = line.split()
|
||||
if len(parts) != 3:
|
||||
return None
|
||||
try:
|
||||
return base64.b64decode(parts[2], validate=True)
|
||||
except Exception:
|
||||
return None
|
||||
return None
|
||||
|
||||
|
||||
def parse_results_from_text(text: str) -> List[Tuple[str, str]]:
|
||||
out = []
|
||||
for line in text.splitlines():
|
||||
line = line.strip()
|
||||
if line.startswith(RESULT_PREFIX):
|
||||
parts = line.split()
|
||||
if len(parts) < 4:
|
||||
continue
|
||||
sig_hex = parts[2]
|
||||
payload_b64 = parts[3]
|
||||
out.append((sig_hex, payload_b64))
|
||||
return out
|
||||
|
||||
|
||||
def verify_and_decode_payload(nonce: bytes, sig_hex: str, payload_b64: str) -> Optional[dict]:
|
||||
try:
|
||||
payload = base64.b64decode(payload_b64, validate=True)
|
||||
except Exception:
|
||||
return None
|
||||
sig_calc = hashlib.sha256(nonce + payload).hexdigest()
|
||||
if sig_calc.lower() != sig_hex.lower():
|
||||
return None
|
||||
try:
|
||||
obj = json.loads(payload.decode("utf-8"))
|
||||
except Exception:
|
||||
return None
|
||||
if not isinstance(obj, dict):
|
||||
return None
|
||||
return obj
|
||||
|
||||
|
||||
def _read_testcase_id() -> Optional[int]:
|
||||
try:
|
||||
text = Path("input").read_text(encoding="utf-8").strip()
|
||||
return int(text.split()[0])
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
def _format_diagnostic(payload: dict, describe_fn: Optional[Callable[[int], str]] = None) -> List[str]:
|
||||
tk = payload.get("tk_time_ms")
|
||||
tb = payload.get("tb_time_ms")
|
||||
th = payload.get("th_time_ms")
|
||||
score = payload.get("score_ratio")
|
||||
passed = payload.get("pass", False)
|
||||
lines: List[str] = []
|
||||
lines.append("-" * 64)
|
||||
tc_id = _read_testcase_id()
|
||||
if tc_id is not None:
|
||||
lines.append(f" Testcase #{tc_id}")
|
||||
if describe_fn:
|
||||
desc = describe_fn(tc_id)
|
||||
if desc:
|
||||
lines.append(f" Config: {desc}")
|
||||
else:
|
||||
lines.append(" Testcase <?> (no 'input' file)")
|
||||
lines.append("")
|
||||
if tk is not None and tb is not None:
|
||||
speedup = tb / tk if tk > 0 else float("inf")
|
||||
lines.append(f" Baseline: {tb:<12.6f} ms")
|
||||
lines.append(f" User kernel: {tk:<12.6f} ms")
|
||||
if th is not None:
|
||||
lines.append(f" Hardware bound: {th:<12.6f} ms")
|
||||
lines.append(f" Speedup vs base: {speedup:<12.3f}x")
|
||||
else:
|
||||
lines.append(" (timing data unavailable)")
|
||||
lines.append("")
|
||||
if score is not None:
|
||||
pct = float(score) * 100.0
|
||||
raw_display = int(pct)
|
||||
final_display = raw_display
|
||||
if raw_display > 150:
|
||||
final_display = int(150 + 10 * math.log10(raw_display / 150))
|
||||
lines.append(f" Score ratio: {score:<12.6f} ({pct:.2f}%)")
|
||||
lines.append(f" Display score: {final_display:<12d} / 100")
|
||||
else:
|
||||
lines.append(" Score: (unavailable)")
|
||||
lines.append(f" Pass: {'OK' if passed else 'FAIL'}")
|
||||
lines.append("-" * 64)
|
||||
return lines
|
||||
|
||||
|
||||
def run_spj(testcases: list, describe_fn: Optional[Callable[[int], str]] = None, problem_name: str = "") -> int:
|
||||
cwd = Path(".")
|
||||
user_out_path = cwd / "user_out"
|
||||
try:
|
||||
user_text = read_text(user_out_path)
|
||||
except Exception as ex:
|
||||
eprint(str(ex))
|
||||
print("0", flush=True)
|
||||
return 0
|
||||
nonce_ans = parse_nonce_from_text(user_text)
|
||||
if nonce_ans is None:
|
||||
eprint("SPJ FAIL: cannot find/parse nonce from user_out (expected 'OJCHAL v1 <b64>').")
|
||||
print("0", flush=True)
|
||||
return 0
|
||||
results = parse_results_from_text(user_text)
|
||||
if not results:
|
||||
eprint("SPJ FAIL: no OJRESULT line found in user_out.")
|
||||
print("0", flush=True)
|
||||
return 0
|
||||
payload_obj = None
|
||||
for sig_hex, payload_b64 in reversed(results):
|
||||
obj = verify_and_decode_payload(nonce_ans, sig_hex, payload_b64)
|
||||
if obj is not None:
|
||||
payload_obj = obj
|
||||
break
|
||||
if payload_obj is None:
|
||||
eprint("SPJ FAIL: no OJRESULT line passes signature verification.")
|
||||
print("0", flush=True)
|
||||
return 0
|
||||
ok = bool(payload_obj.get("pass", False))
|
||||
tk_time_ms = payload_obj.get("tk_time_ms", None)
|
||||
score_ratio = payload_obj.get("score_ratio", None)
|
||||
if not isinstance(tk_time_ms, (int, float)):
|
||||
eprint("SPJ FAIL: payload missing/invalid 'tk_time_ms'.")
|
||||
print("0", flush=True)
|
||||
return 0
|
||||
if not isinstance(score_ratio, (int, float)) or not (0.0 <= score_ratio):
|
||||
eprint("SPJ FAIL: payload missing/invalid 'score_ratio'.")
|
||||
print("0", flush=True)
|
||||
return 0
|
||||
time_us = float(tk_time_ms) * 1000.0
|
||||
if not ok:
|
||||
eprint("SPJ FAIL: payload pass=false.")
|
||||
print("0", flush=True)
|
||||
return 0
|
||||
displayScore = int(score_ratio * 100)
|
||||
if displayScore > 150:
|
||||
displayScore = int(150 + 10 * math.log10(displayScore / 150))
|
||||
extraInfo = {"rewriteTimeUs": time_us, "displayScore": displayScore}
|
||||
print(f"{100} {json.dumps(extraInfo)}", flush=True)
|
||||
header = f"=== SPJ Report{' - ' + problem_name if problem_name else ''} ==="
|
||||
eprint(header)
|
||||
for line in _format_diagnostic(payload_obj, describe_fn=describe_fn):
|
||||
eprint(line)
|
||||
return 0
|
||||
|
||||
HEAD_DIMS = [128]
|
||||
BATCH_SIZES = [1, 4, 16]
|
||||
SEQ_LENS_KV = [1024, 4096, 8192, 16384]
|
||||
SEQ_LEN_Q = 1
|
||||
NUM_HEADS = 8
|
||||
NUM_HEADS_K = 8
|
||||
PAGE_BLOCK_SIZE = 16
|
||||
CAUSAL = 0
|
||||
|
||||
TESTCASES = []
|
||||
for headdim in HEAD_DIMS:
|
||||
for seqlen_k in SEQ_LENS_KV:
|
||||
for batch_size in BATCH_SIZES:
|
||||
TESTCASES.append((batch_size, seqlen_k, SEQ_LEN_Q, NUM_HEADS, NUM_HEADS_K, headdim, PAGE_BLOCK_SIZE, CAUSAL))
|
||||
|
||||
|
||||
def describe(tc_id: int) -> str:
|
||||
if tc_id < 1 or tc_id > len(TESTCASES):
|
||||
return f"testcase #{tc_id}"
|
||||
b, slk, slq, nh, nhk, hd, pbs, c = TESTCASES[tc_id - 1]
|
||||
return f"batch={b}, seqlen_k={slk}, seqlen_q={slq}, heads={nh}, kv_heads={nhk}, headdim={hd}, page_size={pbs}, causal={c}"
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(run_spj(TESTCASES, describe_fn=describe, problem_name="Paged KV Cache (flash-attn)"))
|
||||
|
|
@ -1,350 +0,0 @@
|
|||
from __future__ import annotations
|
||||
|
||||
|
||||
HEAD_DIMS = [128]
|
||||
BATCH_SIZES = [1, 4, 16]
|
||||
SEQ_LENS_KV = [1024, 4096, 8192, 16384]
|
||||
SEQ_LEN_Q = 1
|
||||
NUM_HEADS = 8
|
||||
NUM_HEADS_K = 8
|
||||
PAGE_BLOCK_SIZE = 16
|
||||
CAUSAL = 0
|
||||
|
||||
|
||||
def _build_cases():
|
||||
cases = []
|
||||
for headdim in HEAD_DIMS:
|
||||
for seqlen_k in SEQ_LENS_KV:
|
||||
for batch_size in BATCH_SIZES:
|
||||
cases.append(
|
||||
(
|
||||
batch_size,
|
||||
seqlen_k,
|
||||
SEQ_LEN_Q,
|
||||
NUM_HEADS,
|
||||
NUM_HEADS_K,
|
||||
headdim,
|
||||
PAGE_BLOCK_SIZE,
|
||||
CAUSAL,
|
||||
)
|
||||
)
|
||||
return cases
|
||||
|
||||
|
||||
TESTCASES = _build_cases()
|
||||
|
||||
|
||||
def getNumOfTestcases() -> int:
|
||||
return len(TESTCASES)
|
||||
|
||||
|
||||
try:
|
||||
from pathlib import Path
|
||||
from typing import List, Tuple, Union
|
||||
import math
|
||||
import sys
|
||||
|
||||
import torch
|
||||
|
||||
KernelArg = Union[torch.Tensor, int, float]
|
||||
CURRENT_CASE = None
|
||||
|
||||
def _ensure_flashattn_importable():
|
||||
try:
|
||||
from flash_attn.flash_attn_interface import flash_attn_with_kvcache # noqa: F401
|
||||
|
||||
return
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
here = Path(__file__).resolve()
|
||||
for parent in here.parents:
|
||||
candidate = parent / "flashattn"
|
||||
if (candidate / "flash_attn").is_dir():
|
||||
sys.path.insert(0, str(candidate))
|
||||
return
|
||||
|
||||
def _get_testcase_index() -> int:
|
||||
try:
|
||||
raw = input().strip()
|
||||
except EOFError:
|
||||
return 0
|
||||
if raw == "":
|
||||
return 0
|
||||
try:
|
||||
testcase_id = int(raw.split()[0])
|
||||
except ValueError:
|
||||
return 0
|
||||
if 1 <= testcase_id <= len(TESTCASES):
|
||||
return testcase_id - 1
|
||||
if 0 <= testcase_id < len(TESTCASES):
|
||||
return testcase_id
|
||||
return 0
|
||||
|
||||
def _compute_reps(batch_size: int, seq_len: int, head_dim: int, base_reps: int = 100) -> int:
|
||||
workload = batch_size * seq_len * head_dim
|
||||
if workload < 1e5:
|
||||
return base_reps
|
||||
if workload < 1e6:
|
||||
return base_reps // 2
|
||||
if workload < 1e7:
|
||||
return base_reps // 4
|
||||
if workload < 1e8:
|
||||
return base_reps // 8
|
||||
if workload < 1e9:
|
||||
return base_reps // 16
|
||||
return base_reps // 32
|
||||
|
||||
def _get_num_blocks(batch_size: int, seqlen_k: int, page_block_size: int) -> int:
|
||||
num_blocks = math.ceil(seqlen_k / page_block_size) * batch_size * 3
|
||||
return max(1024, num_blocks)
|
||||
|
||||
def getTestCaseSize() -> Tuple[List[Tuple[int, ...]], Tuple[int, int]]:
|
||||
testcase_id = _get_testcase_index()
|
||||
global CURRENT_CASE
|
||||
(
|
||||
batch_size,
|
||||
seqlen_k,
|
||||
seqlen_q,
|
||||
num_heads,
|
||||
num_heads_k,
|
||||
headdim,
|
||||
page_block_size,
|
||||
causal,
|
||||
) = TESTCASES[testcase_id]
|
||||
num_blocks = _get_num_blocks(batch_size, seqlen_k, page_block_size)
|
||||
blocks_per_batch = num_blocks // batch_size
|
||||
CURRENT_CASE = (
|
||||
batch_size,
|
||||
seqlen_k,
|
||||
seqlen_q,
|
||||
num_heads,
|
||||
num_heads_k,
|
||||
headdim,
|
||||
page_block_size,
|
||||
num_blocks,
|
||||
causal,
|
||||
20260720 + testcase_id,
|
||||
)
|
||||
warmup = 3
|
||||
iters = max(1, _compute_reps(batch_size, seqlen_k, headdim))
|
||||
return [
|
||||
(batch_size, seqlen_q, num_heads, headdim),
|
||||
(num_blocks, page_block_size, num_heads_k, headdim),
|
||||
(num_blocks, page_block_size, num_heads_k, headdim),
|
||||
(batch_size, seqlen_q, num_heads, headdim),
|
||||
(batch_size,),
|
||||
(batch_size, blocks_per_batch),
|
||||
(), (), (), (), (), (), (), (), (),
|
||||
], (warmup, iters)
|
||||
|
||||
def genTestCase(testcase_sizes, device: str = "cuda") -> List[KernelArg]:
|
||||
del testcase_sizes
|
||||
(
|
||||
batch_size,
|
||||
seqlen_k,
|
||||
seqlen_q,
|
||||
num_heads,
|
||||
num_heads_k,
|
||||
headdim,
|
||||
page_block_size,
|
||||
num_blocks,
|
||||
causal,
|
||||
seed,
|
||||
) = CURRENT_CASE
|
||||
gen = torch.Generator(device=device)
|
||||
gen.manual_seed(seed)
|
||||
dtype = torch.bfloat16
|
||||
blocks_per_batch = num_blocks // batch_size
|
||||
q = torch.randn(
|
||||
batch_size,
|
||||
seqlen_q,
|
||||
num_heads,
|
||||
headdim,
|
||||
dtype=dtype,
|
||||
device=device,
|
||||
generator=gen,
|
||||
).contiguous()
|
||||
k_cache_paged = torch.randn(
|
||||
num_blocks,
|
||||
page_block_size,
|
||||
num_heads_k,
|
||||
headdim,
|
||||
dtype=dtype,
|
||||
device=device,
|
||||
generator=gen,
|
||||
).contiguous()
|
||||
v_cache_paged = torch.randn(
|
||||
num_blocks,
|
||||
page_block_size,
|
||||
num_heads_k,
|
||||
headdim,
|
||||
dtype=dtype,
|
||||
device=device,
|
||||
generator=gen,
|
||||
).contiguous()
|
||||
output = torch.empty(
|
||||
batch_size,
|
||||
seqlen_q,
|
||||
num_heads,
|
||||
headdim,
|
||||
dtype=dtype,
|
||||
device=device,
|
||||
)
|
||||
cache_seqlens = torch.full((batch_size,), seqlen_k, dtype=torch.int32, device=device)
|
||||
block_table = torch.randperm(num_blocks, dtype=torch.int32, device=device, generator=gen).reshape(
|
||||
batch_size,
|
||||
blocks_per_batch,
|
||||
)
|
||||
return [
|
||||
q,
|
||||
k_cache_paged,
|
||||
v_cache_paged,
|
||||
output,
|
||||
cache_seqlens,
|
||||
block_table,
|
||||
batch_size,
|
||||
seqlen_k,
|
||||
seqlen_q,
|
||||
num_heads,
|
||||
num_heads_k,
|
||||
headdim,
|
||||
page_block_size,
|
||||
num_blocks,
|
||||
causal,
|
||||
]
|
||||
|
||||
def baseline(
|
||||
q,
|
||||
k_cache_paged,
|
||||
v_cache_paged,
|
||||
output,
|
||||
cache_seqlens,
|
||||
block_table,
|
||||
batch_size,
|
||||
seqlen_k,
|
||||
seqlen_q,
|
||||
num_heads,
|
||||
num_heads_k,
|
||||
headdim,
|
||||
page_block_size,
|
||||
num_blocks,
|
||||
causal,
|
||||
):
|
||||
_ensure_flashattn_importable()
|
||||
from flash_attn.flash_attn_interface import flash_attn_with_kvcache
|
||||
|
||||
out = flash_attn_with_kvcache(
|
||||
q,
|
||||
k_cache_paged,
|
||||
v_cache_paged,
|
||||
None,
|
||||
None,
|
||||
cache_seqlens=cache_seqlens,
|
||||
cache_batch_idx=None,
|
||||
block_table=block_table,
|
||||
causal=bool(causal),
|
||||
window_size=(-1, -1),
|
||||
rotary_interleaved=False,
|
||||
alibi_slopes=None,
|
||||
num_splits=1,
|
||||
)
|
||||
output.copy_(out)
|
||||
return [
|
||||
q,
|
||||
k_cache_paged,
|
||||
v_cache_paged,
|
||||
output,
|
||||
cache_seqlens,
|
||||
block_table,
|
||||
batch_size,
|
||||
seqlen_k,
|
||||
seqlen_q,
|
||||
num_heads,
|
||||
num_heads_k,
|
||||
headdim,
|
||||
page_block_size,
|
||||
num_blocks,
|
||||
causal,
|
||||
]
|
||||
|
||||
def check(
|
||||
testcase_sizes,
|
||||
original_input_tensors,
|
||||
target_kernel_input_tensors,
|
||||
baseline_input_tensors,
|
||||
rtol=1e-2,
|
||||
atol=1e-2,
|
||||
) -> bool:
|
||||
del testcase_sizes, original_input_tensors
|
||||
output_t = target_kernel_input_tensors[3]
|
||||
output_ref = baseline_input_tensors[3]
|
||||
if output_t.shape != output_ref.shape:
|
||||
print(f"[FAIL] shape mismatch: target {output_t.shape}, ref {output_ref.shape}", file=sys.stderr)
|
||||
return False
|
||||
if output_t.dtype != output_ref.dtype:
|
||||
print(f"[FAIL] dtype mismatch: target {output_t.dtype}, ref {output_ref.dtype}", file=sys.stderr)
|
||||
return False
|
||||
if not torch.allclose(output_t.float(), output_ref.float(), rtol=rtol, atol=atol):
|
||||
diff = (output_t.float() - output_ref.float()).abs()
|
||||
print(
|
||||
f"[FAIL] allclose failed: max_abs_diff={float(diff.max().item()):.6f}, "
|
||||
f"mean_abs_diff={float(diff.mean().item()):.6f} (rtol={rtol}, atol={atol})",
|
||||
file=sys.stderr,
|
||||
)
|
||||
return False
|
||||
return True
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
INPUT_CLASS = [
|
||||
"INPUT", # q
|
||||
"INPUT", # k_cache_paged
|
||||
"INPUT", # v_cache_paged
|
||||
"OUTPUT", # output
|
||||
"INPUT", # cache_seqlens
|
||||
"INPUT", # block_table
|
||||
"INPUT", # batch_size
|
||||
"INPUT", # seqlen_k
|
||||
"INPUT", # seqlen_q
|
||||
"INPUT", # num_heads
|
||||
"INPUT", # num_heads_k
|
||||
"INPUT", # headdim
|
||||
"INPUT", # page_block_size
|
||||
"INPUT", # num_blocks
|
||||
"INPUT", # causal
|
||||
]
|
||||
|
||||
|
||||
def getWorkload(testcase_sizes) -> dict:
|
||||
raw_sizes = testcase_sizes[0] if isinstance(testcase_sizes, tuple) and len(testcase_sizes) == 2 else testcase_sizes
|
||||
q_shape, k_shape, v_shape, output_shape, cache_seqlens_shape, block_table_shape = raw_sizes[:6]
|
||||
batch_size, seqlen_q, num_heads, headdim = q_shape
|
||||
num_blocks, page_block_size, num_heads_k, headdim_k = k_shape
|
||||
assert v_shape == k_shape
|
||||
assert output_shape == q_shape
|
||||
blocks_per_batch = block_table_shape[1]
|
||||
# KV 长度按 padded (blocks_per_batch * page_block_size) 估计
|
||||
seqlen_k = blocks_per_batch * page_block_size
|
||||
# QK: 2 * batch * seqlen_q * num_heads * seqlen_k * headdim
|
||||
# PV: 2 * batch * seqlen_q * num_heads * seqlen_k * headdim
|
||||
flops = 4 * batch_size * seqlen_q * num_heads * seqlen_k * headdim
|
||||
# memory_bytes 只算输入/输出变量的 IO,不算 softmax 等中间结果
|
||||
# k_cache / v_cache 按 tensor 实际占用 (padded) 计算 IO
|
||||
memory_bytes = (
|
||||
batch_size * seqlen_q * num_heads * headdim * 2 # q (bf16)
|
||||
+ num_blocks * page_block_size * num_heads_k * headdim * 2 # k_cache (bf16)
|
||||
+ num_blocks * page_block_size * num_heads_k * headdim * 2 # v_cache (bf16)
|
||||
+ batch_size * seqlen_q * num_heads * headdim * 2 # output (bf16)
|
||||
+ batch_size * 4 # cache_seqlens (int32)
|
||||
+ batch_size * blocks_per_batch * 4 # block_table (int32)
|
||||
)
|
||||
return {
|
||||
"flops": flops,
|
||||
"memory_bytes": memory_bytes,
|
||||
"dtype": "bf16",
|
||||
}
|
||||
|
||||
|
||||
DESIGNED_VRAM_SIZE = 128
|
||||
|
|
@ -1,28 +0,0 @@
|
|||
---
|
||||
sectionTitle: "题目描述"
|
||||
type: "Text"
|
||||
---
|
||||
你需要实现 FlashAttention paged KV cache decode 的 CUDA C++ 前向算子。
|
||||
|
||||
本题输入采用 `flash_attn_with_kvcache` 在 `flashattn/benchmarks/benchmark_kvcache.py` 中使用的 paged KV cache 配置。每个 batch 只有 1 个 query token,KV cache 长度为 `seqlen_k`,K/V cache 按 page 存储。
|
||||
|
||||
评测程序会调用你提交代码中的 `run_kernel` 函数。你需要根据 `cache_seqlens` 和 `block_table` 读取 paged KV cache,并将结果写入 `output`。
|
||||
|
||||
baseline 使用 benchmark 中的 FlashAttention Python API:
|
||||
|
||||
```python
|
||||
out = flash_attn_with_kvcache(
|
||||
q, k_cache_paged, v_cache_paged, None, None,
|
||||
cache_seqlens=cache_seqlens,
|
||||
cache_batch_idx=None,
|
||||
block_table=block_table,
|
||||
causal=False,
|
||||
window_size=(-1, -1),
|
||||
rotary_interleaved=False,
|
||||
alibi_slopes=None,
|
||||
num_splits=1,
|
||||
)
|
||||
output.copy_(out)
|
||||
```
|
||||
|
||||
如何提交代码详见[评测指南](/d/2)。
|
||||
|
|
@ -1,43 +0,0 @@
|
|||
---
|
||||
sectionTitle: "接口约定"
|
||||
type: "codeSample"
|
||||
lang: "cuda"
|
||||
---
|
||||
你必须在提交的 CUDA 源码中提供如下 **C 符号**,函数名、参数类型、顺序必须完全一致,并使用 `extern "C"` 防止 name mangling:
|
||||
|
||||
```cpp
|
||||
#include <stdint.h>
|
||||
#include <cuda_bf16.h>
|
||||
|
||||
extern "C" void run_kernel(
|
||||
const __nv_bfloat16* q,
|
||||
const __nv_bfloat16* k_cache_paged,
|
||||
const __nv_bfloat16* v_cache_paged,
|
||||
__nv_bfloat16* output,
|
||||
const int32_t* cache_seqlens,
|
||||
const int32_t* block_table,
|
||||
int64_t batch_size,
|
||||
int64_t seqlen_k,
|
||||
int64_t seqlen_q,
|
||||
int64_t num_heads,
|
||||
int64_t num_heads_k,
|
||||
int64_t headdim,
|
||||
int64_t page_block_size,
|
||||
int64_t num_blocks,
|
||||
int64_t causal
|
||||
);
|
||||
```
|
||||
|
||||
### 参数说明
|
||||
|
||||
* `q`:decode query tensor,shape `(batch_size, seqlen_q, num_heads, headdim)`,连续 `bf16`
|
||||
* `k_cache_paged`:paged key cache,shape `(num_blocks, page_block_size, num_heads_k, headdim)`,连续 `bf16`
|
||||
* `v_cache_paged`:paged value cache,shape `(num_blocks, page_block_size, num_heads_k, headdim)`,连续 `bf16`
|
||||
* `output`:输出缓冲区,shape `(batch_size, seqlen_q, num_heads, headdim)`,连续 `bf16`
|
||||
* `cache_seqlens`:每个 batch 的 KV 长度,shape `(batch_size)`,连续 `int32`
|
||||
* `block_table`:每个 batch 的 page 映射表,shape `(batch_size, num_blocks / batch_size)`,连续 `int32`
|
||||
* `seqlen_q`:query 长度,评测中固定为 `1`
|
||||
* `page_block_size`:page size,评测中固定为 `16`
|
||||
* `causal`:是否启用 causal mask,评测中固定为 `0`
|
||||
|
||||
`run_kernel` 内部需要自行计算合适的 launch 配置并启动 CUDA kernel。为保证计时准确,不建议在 `run_kernel` 内部做 `cudaDeviceSynchronize()` 或显式同步。
|
||||
|
|
@ -1,57 +0,0 @@
|
|||
---
|
||||
sectionTitle: "接口约定"
|
||||
type: "codeSample"
|
||||
lang: "tilelang"
|
||||
---
|
||||
你必须在提交的 Python 代码中提供 `run_kernel` 函数,函数名、参数顺序、类型必须完全一致:
|
||||
|
||||
```python
|
||||
import tilelang
|
||||
import tilelang.language as T
|
||||
from tilelang import jit
|
||||
|
||||
real_kernel = None
|
||||
|
||||
@jit
|
||||
def build_kernel(*args):
|
||||
@T.prim_func
|
||||
def kernel(*args):
|
||||
...
|
||||
return kernel
|
||||
|
||||
def run_kernel(
|
||||
q, # Tensor[bf16], shape (batch_size, seqlen_q, num_heads, headdim)
|
||||
k_cache_paged, # Tensor[bf16], shape (num_blocks, page_block_size, num_heads_k, headdim)
|
||||
v_cache_paged, # Tensor[bf16], shape (num_blocks, page_block_size, num_heads_k, headdim)
|
||||
output, # Tensor[bf16], shape (batch_size, seqlen_q, num_heads, headdim)
|
||||
cache_seqlens, # Tensor[int32], shape (batch_size)
|
||||
block_table, # Tensor[int32], shape (batch_size, num_blocks / batch_size)
|
||||
batch_size, # int64
|
||||
seqlen_k, # int64
|
||||
seqlen_q, # int64
|
||||
num_heads, # int64
|
||||
num_heads_k, # int64
|
||||
headdim, # int64
|
||||
page_block_size, # int64
|
||||
num_blocks, # int64
|
||||
causal, # int64
|
||||
):
|
||||
global real_kernel
|
||||
if real_kernel is None:
|
||||
real_kernel = build_kernel(...)
|
||||
real_kernel(q, k_cache_paged, v_cache_paged, output,
|
||||
cache_seqlens, block_table,
|
||||
batch_size, seqlen_k, seqlen_q, num_heads,
|
||||
num_heads_k, headdim, page_block_size, num_blocks, causal)
|
||||
```
|
||||
|
||||
### 参数说明
|
||||
|
||||
* `q`:decode query tensor,连续 `bfloat16`
|
||||
* `k_cache_paged/v_cache_paged`:paged KV cache,连续 `bfloat16`
|
||||
* `output`:输出缓冲区,连续 `bfloat16`,需要写入结果
|
||||
* `cache_seqlens/block_table`:paged KV metadata,连续 `int32`
|
||||
* `page_block_size`:评测中固定为 `16`
|
||||
* `causal`:评测中固定为 `0`
|
||||
|
||||
`run_kernel` 内部需要自行计算合适的 grid/block,并 launch 你实现的 TileLang kernel。
|
||||
|
|
@ -1,45 +0,0 @@
|
|||
---
|
||||
sectionTitle: "接口约定"
|
||||
type: "codeSample"
|
||||
lang: "triton"
|
||||
---
|
||||
你必须在提交的 Python 代码中提供 `run_kernel` 函数,函数名、参数顺序、类型必须完全一致:
|
||||
|
||||
```python
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
@triton.jit
|
||||
def your_kernel(...):
|
||||
...
|
||||
|
||||
def run_kernel(
|
||||
q, # Tensor[bf16], shape (batch_size, seqlen_q, num_heads, headdim)
|
||||
k_cache_paged, # Tensor[bf16], shape (num_blocks, page_block_size, num_heads_k, headdim)
|
||||
v_cache_paged, # Tensor[bf16], shape (num_blocks, page_block_size, num_heads_k, headdim)
|
||||
output, # Tensor[bf16], shape (batch_size, seqlen_q, num_heads, headdim)
|
||||
cache_seqlens, # Tensor[int32], shape (batch_size)
|
||||
block_table, # Tensor[int32], shape (batch_size, num_blocks / batch_size)
|
||||
batch_size, # int64
|
||||
seqlen_k, # int64
|
||||
seqlen_q, # int64
|
||||
num_heads, # int64
|
||||
num_heads_k, # int64
|
||||
headdim, # int64
|
||||
page_block_size, # int64
|
||||
num_blocks, # int64
|
||||
causal, # int64
|
||||
):
|
||||
...
|
||||
```
|
||||
|
||||
### 参数说明
|
||||
|
||||
* `q`:decode query tensor,连续 `bfloat16`
|
||||
* `k_cache_paged/v_cache_paged`:paged KV cache,连续 `bfloat16`
|
||||
* `output`:输出缓冲区,连续 `bfloat16`,需要写入结果
|
||||
* `cache_seqlens/block_table`:paged KV metadata,连续 `int32`
|
||||
* `page_block_size`:评测中固定为 `16`
|
||||
* `causal`:评测中固定为 `0`
|
||||
|
||||
`run_kernel` 内部需要自行计算合适的 grid/block,并 launch 你实现的 Triton kernel。
|
||||
|
|
@ -1,9 +0,0 @@
|
|||
---
|
||||
sectionTitle: "输入格式"
|
||||
type: "Text"
|
||||
---
|
||||
本题输入由评测程序在 GPU 上构造,并按接口约定中的顺序传入 `run_kernel`。
|
||||
|
||||
`q/k_cache_paged/v_cache_paged/output` 均为连续 `torch.bfloat16` CUDA tensor,`cache_seqlens/block_table` 均为连续 `torch.int32` CUDA tensor。
|
||||
|
||||
KV cache layout 固定为 `flash_attn_with_kvcache` 的 paged cache 布局:`(num_blocks, page_block_size, num_heads_k, headdim)`。
|
||||
|
|
@ -1,5 +0,0 @@
|
|||
---
|
||||
sectionTitle: "输出格式"
|
||||
type: "Text"
|
||||
---
|
||||
输出写入 `output`,shape 为 `(batch_size, 1, num_heads, headdim)`,类型为 `bfloat16`。
|
||||
|
|
@ -1,12 +0,0 @@
|
|||
---
|
||||
sectionTitle: "样例"
|
||||
type: "Text"
|
||||
---
|
||||
若 `batch_size = 1`、`seqlen_k = 512`、`page_block_size = 16`,则每个序列需要访问 `32` 个有效 page:
|
||||
|
||||
```text
|
||||
cache_seqlens = [512]
|
||||
block_table.shape = (1, num_blocks)
|
||||
```
|
||||
|
||||
第 `t` 个 KV token 位于 `block_table[0, t / 16]` 指向的物理 page 中,page 内偏移为 `t % 16`。
|
||||
|
|
@ -1 +0,0 @@
|
|||
Agent 推理算子库优化 - FlashAttention KV Cache Decode
|
||||
|
|
@ -1,145 +0,0 @@
|
|||
api,batch_size,seq_len_q,seq_len_kv,num_qo_heads,num_kv_heads,head_dim,time_ms,bandwidth_GB_s,tflops
|
||||
BatchDecodeWithPagedKVCacheWrapper,1,1,512,32,8,64,0.02042879999999998,51.528822055137894,0.8212531328320811
|
||||
BatchDecodeWithPagedKVCacheWrapper,1,1,512,32,4,128,0.02333952000000001,45.27805199078642,0.718832949435121
|
||||
BatchDecodeWithPagedKVCacheWrapper,1,1,512,32,4,256,0.0319488,66.15384615384615,1.0502564102564103
|
||||
BatchDecodeWithPagedKVCacheWrapper,1,1,1024,32,8,64,0.023262719999999973,90.32684054143292,1.4424122372620245
|
||||
BatchDecodeWithPagedKVCacheWrapper,1,1,1024,32,4,128,0.025041919999999992,84.07278675117566,1.3399304845634845
|
||||
BatchDecodeWithPagedKVCacheWrapper,1,1,1024,32,4,256,0.033387520000000004,126.11562643766291,2.0099984664928687
|
||||
BatchDecodeWithPagedKVCacheWrapper,1,1,2048,32,8,64,0.028298240000000037,148.36258368011562,2.371485435136599
|
||||
BatchDecodeWithPagedKVCacheWrapper,1,1,2048,32,4,128,0.027745280000000008,151.4670603432367,2.418748846650673
|
||||
BatchDecodeWithPagedKVCacheWrapper,1,1,2048,32,4,256,0.03723775999999999,225.7115358174069,3.604344837068611
|
||||
BatchDecodeWithPagedKVCacheWrapper,1,1,4096,32,8,64,0.03886591999999997,215.93992886312756,3.4533526544592306
|
||||
BatchDecodeWithPagedKVCacheWrapper,1,1,4096,32,4,128,0.03426815999999998,245.03212311370103,3.916689078141344
|
||||
BatchDecodeWithPagedKVCacheWrapper,1,1,4096,32,4,256,0.066048,254.26356589147287,4.064248062015504
|
||||
BatchDecodeWithPagedKVCacheWrapper,1,1,8192,32,8,64,0.052495359999999984,319.6722910367698,5.1135082414902975
|
||||
BatchDecodeWithPagedKVCacheWrapper,1,1,8192,32,4,128,0.04628480000000001,362.6548672566371,5.799646017699114
|
||||
BatchDecodeWithPagedKVCacheWrapper,1,1,8192,32,4,256,0.08975359999999999,374.0330861380491,5.981608670849972
|
||||
BatchDecodeWithPagedKVCacheWrapper,1,1,16384,32,8,64,0.08625152000000001,389.0775258221536,6.224480588863825
|
||||
BatchDecodeWithPagedKVCacheWrapper,1,1,16384,32,4,128,0.0638464,525.6776263031276,8.408789093825181
|
||||
BatchDecodeWithPagedKVCacheWrapper,1,1,16384,32,4,256,0.13059071999999994,514.0123892417473,8.222190857053247
|
||||
BatchDecodeWithPagedKVCacheWrapper,2,1,512,32,8,64,0.02342912,89.86013986013987,1.4321678321678322
|
||||
BatchDecodeWithPagedKVCacheWrapper,2,1,512,32,4,128,0.02486784,84.99073502161829,1.3493102738315832
|
||||
BatchDecodeWithPagedKVCacheWrapper,2,1,512,32,4,256,0.03340287999999998,126.54812998160644,2.009074187614961
|
||||
BatchDecodeWithPagedKVCacheWrapper,2,1,1024,32,8,64,0.02839040000000001,148.02524797114512,2.3637871956717755
|
||||
BatchDecodeWithPagedKVCacheWrapper,2,1,1024,32,4,128,0.028165120000000012,149.5000908925649,2.382694055626249
|
||||
BatchDecodeWithPagedKVCacheWrapper,2,1,1024,32,4,256,0.03740160000000001,225.16084873374396,3.5885557837097872
|
||||
BatchDecodeWithPagedKVCacheWrapper,2,1,2048,32,8,64,0.03881984000000001,216.30176734370872,3.457451859667633
|
||||
BatchDecodeWithPagedKVCacheWrapper,2,1,2048,32,4,128,0.03601408000000001,233.38072220642587,3.7268126243957904
|
||||
BatchDecodeWithPagedKVCacheWrapper,2,1,2048,32,4,256,0.06728704000000002,249.82498858621207,3.9894080048698815
|
||||
BatchDecodeWithPagedKVCacheWrapper,2,1,4096,32,8,64,0.052490240000000014,319.7815060476004,5.114007023019897
|
||||
BatchDecodeWithPagedKVCacheWrapper,2,1,4096,32,4,128,0.04626431999999999,362.9924745462595,5.802213368747235
|
||||
BatchDecodeWithPagedKVCacheWrapper,2,1,4096,32,4,256,0.08993791999999999,373.44870773084375,5.969349880450872
|
||||
BatchDecodeWithPagedKVCacheWrapper,2,1,8192,32,8,64,0.08536063999999999,393.18618042226495,6.289443378119003
|
||||
BatchDecodeWithPagedKVCacheWrapper,2,1,8192,32,4,128,0.0630784,532.2077922077922,8.51116883116883
|
||||
BatchDecodeWithPagedKVCacheWrapper,2,1,8192,32,4,256,0.12952576,518.3650881492608,8.289793659577834
|
||||
BatchDecodeWithPagedKVCacheWrapper,2,1,16384,32,8,64,0.15207424000000003,441.34401723789637,7.0606423809844445
|
||||
BatchDecodeWithPagedKVCacheWrapper,2,1,16384,32,4,128,0.10330112,649.8017446471055,10.394290245836638
|
||||
BatchDecodeWithPagedKVCacheWrapper,2,1,16384,32,4,256,0.2281984,588.3060354498541,9.410599057662106
|
||||
BatchDecodeWithPagedKVCacheWrapper,4,1,512,32,8,64,0.0283904,148.3137962128043,2.3637871956717764
|
||||
BatchDecodeWithPagedKVCacheWrapper,4,1,512,32,4,128,0.028078080000000036,150.5470459518598,2.3900802334062696
|
||||
BatchDecodeWithPagedKVCacheWrapper,4,1,512,32,4,256,0.03707903999999999,228.00331400165706,3.619773543220106
|
||||
BatchDecodeWithPagedKVCacheWrapper,4,1,1024,32,8,64,0.03844096000000004,218.64677677144357,3.4915290356952546
|
||||
BatchDecodeWithPagedKVCacheWrapper,4,1,1024,32,4,128,0.03641856000000004,231.23857725291697,3.6854210600309254
|
||||
BatchDecodeWithPagedKVCacheWrapper,4,1,1024,32,4,256,0.06640640000000002,253.63145720894363,4.04231303006939
|
||||
BatchDecodeWithPagedKVCacheWrapper,4,1,2048,32,8,64,0.059007999999999984,284.5986984815619,4.5491366594360105
|
||||
BatchDecodeWithPagedKVCacheWrapper,4,1,2048,32,4,128,0.04641792000000003,362.1442753143611,5.783013456871825
|
||||
BatchDecodeWithPagedKVCacheWrapper,4,1,2048,32,4,256,0.08961023999999998,375.1799794309223,5.991178151068451
|
||||
BatchDecodeWithPagedKVCacheWrapper,4,1,4096,32,8,64,0.09185279999999997,365.484949832776,5.84490523968785
|
||||
BatchDecodeWithPagedKVCacheWrapper,4,1,4096,32,4,128,0.06349823999999998,528.9469440412838,8.454894371875506
|
||||
BatchDecodeWithPagedKVCacheWrapper,4,1,4096,32,4,256,0.1303347200000001,515.3991200502825,8.238340666247638
|
||||
BatchDecodeWithPagedKVCacheWrapper,4,1,8192,32,8,64,0.16568319999999992,405.1421508034613,6.480692212608161
|
||||
BatchDecodeWithPagedKVCacheWrapper,4,1,8192,32,4,128,0.10290176,652.4828341128471,10.43463031147378
|
||||
BatchDecodeWithPagedKVCacheWrapper,4,1,8192,32,4,256,0.22947840000000008,585.1673360107093,9.358107987505575
|
||||
BatchDecodeWithPagedKVCacheWrapper,4,1,16384,32,8,64,0.30601215999999987,438.65613706331163,7.017641547316293
|
||||
BatchDecodeWithPagedKVCacheWrapper,4,1,16384,32,4,128,0.18384895999999992,730.2216776205863,11.680695109724857
|
||||
BatchDecodeWithPagedKVCacheWrapper,4,1,16384,32,4,256,0.4362026666666668,615.5418398787107,9.846265564630505
|
||||
BatchDecodeWithPagedKVCacheWrapper,8,1,512,32,8,64,0.038655999999999975,217.85430463576174,3.472105960264903
|
||||
BatchDecodeWithPagedKVCacheWrapper,8,1,512,32,4,128,0.03645951999999999,231.87754528858312,3.6812807190001418
|
||||
BatchDecodeWithPagedKVCacheWrapper,8,1,512,32,4,256,0.06676480000000001,253.25153374233125,4.020613496932515
|
||||
BatchDecodeWithPagedKVCacheWrapper,8,1,1024,32,8,64,0.05858303999999996,286.9428421604617,4.582135990211505
|
||||
BatchDecodeWithPagedKVCacheWrapper,8,1,1024,32,4,128,0.04676608000000001,360.14889424129615,5.73996058681848
|
||||
BatchDecodeWithPagedKVCacheWrapper,8,1,1024,32,4,256,0.08992768000000002,374.58437713504884,5.970029606012297
|
||||
BatchDecodeWithPagedKVCacheWrapper,8,1,2048,32,8,64,0.092416,363.43490304709144,5.8092853185595565
|
||||
BatchDecodeWithPagedKVCacheWrapper,8,1,2048,32,4,128,0.07130112000000002,471.5208961654457,7.5296280338934345
|
||||
BatchDecodeWithPagedKVCacheWrapper,8,1,2048,32,4,256,0.14862335999999993,452.41835469202175,7.224583161085851
|
||||
BatchDecodeWithPagedKVCacheWrapper,8,1,4096,32,8,64,0.16396288000000003,409.4928803397451,6.548688483637271
|
||||
BatchDecodeWithPagedKVCacheWrapper,8,1,4096,32,4,128,0.11201536000000002,599.6891854831337,9.585665965810401
|
||||
BatchDecodeWithPagedKVCacheWrapper,8,1,4096,32,4,256,0.24935424000000006,538.7869081351894,8.61218019793848
|
||||
BatchDecodeWithPagedKVCacheWrapper,8,1,8192,32,8,64,0.3056947200000001,439.16524302415155,7.024928817874248
|
||||
BatchDecodeWithPagedKVCacheWrapper,8,1,8192,32,4,128,0.20128768000000002,667.1211273337741,10.668728697156228
|
||||
BatchDecodeWithPagedKVCacheWrapper,8,1,8192,32,4,256,0.46690133333333317,575.2104541716168,9.198875628255509
|
||||
BatchDecodeWithPagedKVCacheWrapper,8,1,16384,32,8,64,0.5866495999999998,457.6296037702917,7.321179961598886
|
||||
BatchDecodeWithPagedKVCacheWrapper,8,1,16384,32,4,128,0.37337600000000015,719.1169009256082,11.503062050051419
|
||||
BatchDecodeWithPagedKVCacheWrapper,8,1,16384,32,4,256,0.8934826666666666,601.0211546726517,9.613991308915525
|
||||
BatchDecodeWithPagedKVCacheWrapper,16,1,512,32,8,64,0.0698112,241.26145947928126,3.845163182984965
|
||||
BatchDecodeWithPagedKVCacheWrapper,16,1,512,32,4,128,0.04724735999999999,357.8673602080625,5.681491114000869
|
||||
BatchDecodeWithPagedKVCacheWrapper,16,1,512,32,4,256,0.08954879999999998,377.6329331046313,5.995288736420813
|
||||
BatchDecodeWithPagedKVCacheWrapper,16,1,1024,32,8,64,0.12070911999999998,278.52052935188334,4.447641669494402
|
||||
BatchDecodeWithPagedKVCacheWrapper,16,1,1024,32,4,128,0.07076864000000004,475.9947909130369,7.586282737664589
|
||||
BatchDecodeWithPagedKVCacheWrapper,16,1,1024,32,4,256,0.14710784000000002,457.9702074342197,7.2990115550605585
|
||||
BatchDecodeWithPagedKVCacheWrapper,16,1,2048,32,8,64,0.22239232000000014,302.05359609540454,4.8281425545630325
|
||||
BatchDecodeWithPagedKVCacheWrapper,16,1,2048,32,4,128,0.11209728000000002,599.8355713894217,9.578660820316067
|
||||
BatchDecodeWithPagedKVCacheWrapper,16,1,2048,32,4,256,0.2504192,537.0190145164587,8.575555101206296
|
||||
BatchDecodeWithPagedKVCacheWrapper,16,1,4096,32,8,64,0.42098688000000006,318.97256275539985,5.101070247129791
|
||||
BatchDecodeWithPagedKVCacheWrapper,16,1,4096,32,4,128,0.2027008,662.7936347562515,10.594352109118466
|
||||
BatchDecodeWithPagedKVCacheWrapper,16,1,4096,32,4,256,0.46432,578.6905582356995,9.250015713301172
|
||||
BatchDecodeWithPagedKVCacheWrapper,16,1,8192,32,8,64,0.8234496000000004,326.06851955480926,5.215822918609709
|
||||
BatchDecodeWithPagedKVCacheWrapper,16,1,8192,32,4,128,0.3726506666666667,720.6924662239522,11.525451797572705
|
||||
BatchDecodeWithPagedKVCacheWrapper,16,1,8192,32,4,256,0.8939733333333334,600.8378952392316,9.608714568667223
|
||||
BatchDecodeWithPagedKVCacheWrapper,16,1,16384,32,8,64,1.6324906666666663,328.9062896122735,5.261858317107276
|
||||
BatchDecodeWithPagedKVCacheWrapper,16,1,16384,32,4,128,0.7114879999999999,754.7590177206082,12.073196725735361
|
||||
BatchDecodeWithPagedKVCacheWrapper,16,1,16384,32,4,256,1.742272,616.4387466480549,9.860612570253094
|
||||
BatchDecodeWithPagedKVCacheWrapper,32,1,512,32,8,64,0.08406016,400.730905104154,6.386746254111341
|
||||
BatchDecodeWithPagedKVCacheWrapper,32,1,512,32,4,128,0.08498175999999999,397.92746113989637,6.317484034220991
|
||||
BatchDecodeWithPagedKVCacheWrapper,32,1,512,32,4,256,0.1808896,373.8918765921313,5.935895839230116
|
||||
BatchDecodeWithPagedKVCacheWrapper,32,1,1024,32,8,64,0.14712832,457.0155902004454,7.297995545657015
|
||||
BatchDecodeWithPagedKVCacheWrapper,32,1,1024,32,4,128,0.14887935999999996,452.5208061077104,7.212160396175805
|
||||
BatchDecodeWithPagedKVCacheWrapper,32,1,1024,32,4,256,0.3279462400000001,410.8661712358707,6.548279522887651
|
||||
BatchDecodeWithPagedKVCacheWrapper,32,1,2048,32,8,64,0.27223039999999993,493.51137859695325,7.888478465299983
|
||||
BatchDecodeWithPagedKVCacheWrapper,32,1,2048,32,4,128,0.27833343999999993,483.1610316029581,7.715507155733786
|
||||
BatchDecodeWithPagedKVCacheWrapper,32,1,2048,32,4,256,0.6300373333333331,426.89493109403054,6.817004435715981
|
||||
BatchDecodeWithPagedKVCacheWrapper,32,1,4096,32,8,64,0.52494336,511.6104868913858,8.181772784019977
|
||||
BatchDecodeWithPagedKVCacheWrapper,32,1,4096,32,4,128,0.5080533333333332,528.8767583455806,8.453772496325847
|
||||
BatchDecodeWithPagedKVCacheWrapper,32,1,4096,32,4,256,1.2449493333333332,431.66029782202656,6.899826653186422
|
||||
BatchDecodeWithPagedKVCacheWrapper,32,1,8192,32,8,64,1.0273706666666667,522.6954607749491,8.361086091615102
|
||||
BatchDecodeWithPagedKVCacheWrapper,32,1,8192,32,4,128,1.0078719999999999,532.9377698755399,8.522842773685548
|
||||
BatchDecodeWithPagedKVCacheWrapper,32,1,8192,32,4,256,2.446784,439.05228741073995,7.021408176610604
|
||||
BatchDecodeWithPagedKVCacheWrapper,32,1,16384,32,8,64,2.0322986666666663,528.4030903594223,8.453417534430637
|
||||
BatchDecodeWithPagedKVCacheWrapper,32,1,16384,32,4,128,2.018026666666667,532.2050425498176,8.513202262276018
|
||||
BatchDecodeWithPagedKVCacheWrapper,32,1,16384,32,4,256,4.847957333333333,443.0748433429558,7.087467154826447
|
||||
BatchDecodeWithPagedKVCacheWrapper,64,1,512,32,8,64,0.13077504,515.1671756322919,8.210602145485865
|
||||
BatchDecodeWithPagedKVCacheWrapper,64,1,512,32,4,128,0.14377984000000002,470.39384659212305,7.4679581226408365
|
||||
BatchDecodeWithPagedKVCacheWrapper,64,1,512,32,4,256,0.27039743999999993,500.2499431947286,7.9419525865333656
|
||||
BatchDecodeWithPagedKVCacheWrapper,64,1,1024,32,8,64,0.231424,581.0973451327434,9.279433628318584
|
||||
BatchDecodeWithPagedKVCacheWrapper,64,1,1024,32,4,128,0.25729023999999995,523.6965692907746,8.34654143118682
|
||||
BatchDecodeWithPagedKVCacheWrapper,64,1,1024,32,4,256,0.502016,536.8036715961244,8.555439061703213
|
||||
BatchDecodeWithPagedKVCacheWrapper,64,1,2048,32,8,64,0.4335923199999999,619.7010131544766,9.905542828802874
|
||||
BatchDecodeWithPagedKVCacheWrapper,64,1,2048,32,4,128,0.47517866666666664,566.018137739068,9.038636616683128
|
||||
BatchDecodeWithPagedKVCacheWrapper,64,1,2048,32,4,256,0.9693866666666666,554.9070422535212,8.861205633802816
|
||||
BatchDecodeWithPagedKVCacheWrapper,64,1,4096,32,8,64,0.8388479999999999,640.3222705424583,10.240156252384224
|
||||
BatchDecodeWithPagedKVCacheWrapper,64,1,4096,32,4,128,0.9261013333333336,580.2768883462716,9.275372232844207
|
||||
BatchDecodeWithPagedKVCacheWrapper,64,1,4096,32,4,256,1.906474666666667,563.7580287805205,9.011328335161021
|
||||
BatchDecodeWithPagedKVCacheWrapper,64,1,8192,32,8,64,1.6543999999999999,649.1803481624759,10.384350328820116
|
||||
BatchDecodeWithPagedKVCacheWrapper,64,1,8192,32,4,128,1.8147413333333327,591.9665201137942,9.466841840453299
|
||||
BatchDecodeWithPagedKVCacheWrapper,64,1,8192,32,4,256,3.774634666666667,569.2026947596871,9.10279839037844
|
||||
BatchDecodeWithPagedKVCacheWrapper,64,1,16384,32,8,64,3.2680746666666667,657.1899393567508,10.513755612274872
|
||||
BatchDecodeWithPagedKVCacheWrapper,64,1,16384,32,4,128,3.5912106666666666,598.129192458031,9.567731207451676
|
||||
BatchDecodeWithPagedKVCacheWrapper,64,1,16384,32,4,256,7.526272,570.802632697835,9.130612969608327
|
||||
BatchDecodeWithPagedKVCacheWrapper,128,1,512,32,8,64,0.2176000000000001,619.2188235294115,9.86895058823529
|
||||
BatchDecodeWithPagedKVCacheWrapper,128,1,512,32,4,128,0.21536768,628.0715100798782,9.971243818942565
|
||||
BatchDecodeWithPagedKVCacheWrapper,128,1,512,32,4,256,0.45757866666666663,591.2264441232692,9.386292694298103
|
||||
BatchDecodeWithPagedKVCacheWrapper,128,1,1024,32,8,64,0.39856127999999985,674.826576229382,10.776177997019683
|
||||
BatchDecodeWithPagedKVCacheWrapper,128,1,1024,32,4,128,0.39381333333333335,684.2938244853738,10.906099241603465
|
||||
BatchDecodeWithPagedKVCacheWrapper,128,1,1024,32,4,256,0.8577493333333336,628.3514810853829,10.014504539010616
|
||||
BatchDecodeWithPagedKVCacheWrapper,128,1,2048,32,8,64,0.7606186666666664,706.5238121949853,11.293352330734283
|
||||
BatchDecodeWithPagedKVCacheWrapper,128,1,2048,32,4,128,0.7354026666666665,731.4625203063357,11.680586679043863
|
||||
BatchDecodeWithPagedKVCacheWrapper,128,1,2048,32,4,256,1.6673066666666665,645.2556074467406,10.30396478792144
|
||||
BatchDecodeWithPagedKVCacheWrapper,128,1,4096,32,8,64,1.4816639999999999,725.0403006349618,11.594983197270098
|
||||
BatchDecodeWithPagedKVCacheWrapper,128,1,4096,32,4,128,1.4354773333333333,748.7338009749138,11.968053263583403
|
||||
BatchDecodeWithPagedKVCacheWrapper,128,1,4096,32,4,256,3.2697173333333325,657.4209880731792,10.508473627893627
|
||||
BatchDecodeWithPagedKVCacheWrapper,128,1,8192,32,8,64,2.9226666666666676,734.9479708029195,11.75629734306569
|
||||
BatchDecodeWithPagedKVCacheWrapper,128,1,8192,32,4,128,2.825301333333333,760.4612643087983,12.161442024827087
|
||||
BatchDecodeWithPagedKVCacheWrapper,128,1,8192,32,4,256,6.484309333333334,662.6865294520187,10.5978097594357
|
||||
BatchDecodeWithPagedKVCacheWrapper,128,1,16384,32,8,64,5.794901333333332,741.2536188134122,11.858610316747416
|
||||
BatchDecodeWithPagedKVCacheWrapper,128,1,16384,32,4,128,5.61536,765.0472760428539,12.237768680191476
|
||||
BatchDecodeWithPagedKVCacheWrapper,128,1,16384,32,4,256,12.908458666666668,665.6125232199165,10.647200957222468
|
||||
|
|
|
@ -1,33 +0,0 @@
|
|||
api,batch_size,seq_len,num_heads,head_dim_ckv,head_dim_kpe,time_ms,bandwidth_GB_s,tflops
|
||||
BatchMLAPagedAttentionWrapper,1,1024,64,512,64,0.035975679999999996,34.83953604212624,3.963964989681919
|
||||
BatchMLAPagedAttentionWrapper,1,4096,64,512,64,0.05349631999999998,89.58223668469162,10.662889409963158
|
||||
BatchMLAPagedAttentionWrapper,1,8192,64,512,64,0.06174719999999999,154.0298507462687,18.47615257048093
|
||||
BatchMLAPagedAttentionWrapper,1,16384,64,512,64,0.08995584000000004,210.63775292410136,25.3646831156265
|
||||
BatchMLAPagedAttentionWrapper,4,1024,64,512,64,0.05086207999999998,98.5705657338434,11.215139923495071
|
||||
BatchMLAPagedAttentionWrapper,4,4096,64,512,64,0.08034559999999999,238.58531145451653,28.39858531145452
|
||||
BatchMLAPagedAttentionWrapper,4,8192,64,512,64,0.10866687999999997,350.0942329438373,41.99442140972485
|
||||
BatchMLAPagedAttentionWrapper,4,16384,64,512,64,0.16821760000000002,450.56155836250184,54.2559488662304
|
||||
BatchMLAPagedAttentionWrapper,16,1024,64,512,64,0.06735359999999997,297.7423033067276,33.87645762067656
|
||||
BatchMLAPagedAttentionWrapper,16,4096,64,512,64,0.14288383999999996,536.6395528003728,63.87570143691549
|
||||
BatchMLAPagedAttentionWrapper,16,8192,64,512,64,0.21618431999999987,703.9113289992544,84.43540682321462
|
||||
BatchMLAPagedAttentionWrapper,16,16384,64,512,64,0.39363328000000025,770.1826837405613,92.74424666532254
|
||||
BatchMLAPagedAttentionWrapper,64,1024,64,512,64,0.15278592,525.0226198853926,59.73590697362689
|
||||
BatchMLAPagedAttentionWrapper,64,4096,64,512,64,0.4850483199999999,632.3256206721838,75.26512413443676
|
||||
BatchMLAPagedAttentionWrapper,64,8192,64,512,64,0.9133465600000001,666.4484158127227,79.94166423750474
|
||||
BatchMLAPagedAttentionWrapper,64,16384,64,512,64,1.7720038399999998,684.3541287134007,82.40890045926764
|
||||
BatchMLAPagedAttentionWrapper,1,1024,128,512,64,0.04499968000000001,29.491409716691315,6.338104448742746
|
||||
BatchMLAPagedAttentionWrapper,1,4096,128,512,64,0.05375743999999999,90.51859612362495,21.222191532930147
|
||||
BatchMLAPagedAttentionWrapper,1,8192,128,512,64,0.08302080000000002,115.44865864939868,27.48349059512796
|
||||
BatchMLAPagedAttentionWrapper,1,16384,128,512,64,0.11321343999999998,168.01736613603475,40.30795947901592
|
||||
BatchMLAPagedAttentionWrapper,4,1024,128,512,64,0.05178880000000003,102.50123578843295,22.028907563025196
|
||||
BatchMLAPagedAttentionWrapper,4,4096,128,512,64,0.11032576,176.4247261926861,41.36298496380175
|
||||
BatchMLAPagedAttentionWrapper,4,8192,128,512,64,0.1688268800000001,227.08800873415404,54.06014435615937
|
||||
BatchMLAPagedAttentionWrapper,4,16384,128,512,64,0.30781695999999986,247.18357299091002,59.30021207408457
|
||||
BatchMLAPagedAttentionWrapper,16,1024,128,512,64,0.10527487999999995,201.69734698344004,43.34749896651511
|
||||
BatchMLAPagedAttentionWrapper,16,4096,128,512,64,0.2629478400000002,296.0920614521874,69.41913273750409
|
||||
BatchMLAPagedAttentionWrapper,16,8192,128,512,64,0.3962367999999998,387.02674764181444,92.13485980100793
|
||||
BatchMLAPagedAttentionWrapper,16,16384,128,512,64,0.7528985599999998,404.23663979381246,96.97779742333418
|
||||
BatchMLAPagedAttentionWrapper,64,1024,128,512,64,0.3242547199999998,261.9380714026308,56.29404872811108
|
||||
BatchMLAPagedAttentionWrapper,64,4096,128,512,64,1.1793126399999994,264.07507342582215,61.91271216426548
|
||||
BatchMLAPagedAttentionWrapper,64,8192,128,512,64,2.3186406399999986,264.55887532446616,62.98038839860932
|
||||
BatchMLAPagedAttentionWrapper,64,16384,128,512,64,4.6020608,264.53295358462015,63.462389746784744
|
||||
|
|
|
@ -1,33 +0,0 @@
|
|||
api,batch_size,seq_len,num_qo_heads,num_kv_heads,head_dim,time_ms,bandwidth_GB_s,tflops
|
||||
BatchPrefillWithPagedKVCacheWrapper,1,1024,32,4,128,0.3529011200000001,29.71302556364796,24.34091054174041
|
||||
BatchPrefillWithPagedKVCacheWrapper,1,4096,32,4,128,4.62532608,9.068126068205768,29.714435500296666
|
||||
BatchPrefillWithPagedKVCacheWrapper,1,8192,32,4,128,18.113853439999996,4.631045529757804,30.350019983820744
|
||||
BatchPrefillWithPagedKVCacheWrapper,1,16384,32,4,128,71.05519616000001,2.36115258372119,30.948099145350383
|
||||
BatchPrefillWithPagedKVCacheWrapper,4,1024,32,4,128,1.2374374399999997,33.89507917264893,27.766848858234006
|
||||
BatchPrefillWithPagedKVCacheWrapper,4,4096,32,4,128,17.896878079999997,9.374381344614939,30.71797279003423
|
||||
BatchPrefillWithPagedKVCacheWrapper,4,8192,32,4,128,71.25501952,4.709062214288198,30.861310127559136
|
||||
BatchPrefillWithPagedKVCacheWrapper,4,16384,32,4,128,283.27072767999994,2.3690716139159393,31.051895457919
|
||||
BatchPrefillWithPagedKVCacheWrapper,16,1024,32,4,128,4.752537600000002,35.301595509733566,28.919067041573737
|
||||
BatchPrefillWithPagedKVCacheWrapper,16,4096,32,4,128,70.51405312000001,9.517090711803915,31.185602844439067
|
||||
BatchPrefillWithPagedKVCacheWrapper,16,8192,32,4,128,284.16772266666663,4.723186952426669,30.953878011423416
|
||||
BatchPrefillWithPagedKVCacheWrapper,16,16384,32,4,128,1129.139136,2.377346134250013,31.160351250841774
|
||||
BatchPrefillWithPagedKVCacheWrapper,64,1024,32,4,128,18.757478399999997,35.77712449878125,29.3086203894016
|
||||
BatchPrefillWithPagedKVCacheWrapper,64,4096,32,4,128,281.4907093333333,9.536210151864244,31.248253425628754
|
||||
BatchPrefillWithPagedKVCacheWrapper,64,8192,32,4,128,1134.7048106666668,4.731370722616177,31.007511167737377
|
||||
BatchPrefillWithPagedKVCacheWrapper,64,16384,32,4,128,4514.139178666666,2.378619226173592,31.177037921302507
|
||||
BatchPrefillWithPagedKVCacheWrapper,1,1024,32,4,256,0.7928422399999997,26.4510629504301,21.668710768992337
|
||||
BatchPrefillWithPagedKVCacheWrapper,1,4096,32,4,256,12.533002240000002,6.69321511267838,21.932327281224513
|
||||
BatchPrefillWithPagedKVCacheWrapper,1,8192,32,4,256,49.81321727999999,3.368024977325858,22.072688491402744
|
||||
BatchPrefillWithPagedKVCacheWrapper,1,16384,32,4,256,190.01136128,1.765917141688929,23.14622915954513
|
||||
BatchPrefillWithPagedKVCacheWrapper,4,1024,32,4,256,3.111116800000001,26.963333552761494,22.088362846422218
|
||||
BatchPrefillWithPagedKVCacheWrapper,4,4096,32,4,256,47.738091520000026,7.02885912101079,23.032165567728153
|
||||
BatchPrefillWithPagedKVCacheWrapper,4,8192,32,4,256,190.14286336,3.529391680241077,23.130221315627924
|
||||
BatchPrefillWithPagedKVCacheWrapper,4,16384,32,4,256,759.6848640000004,1.76675532658763,23.157215416649382
|
||||
BatchPrefillWithPagedKVCacheWrapper,16,1024,32,4,256,12.28442624,27.31461066593534,22.376129057534232
|
||||
BatchPrefillWithPagedKVCacheWrapper,16,4096,32,4,256,191.34602666666663,7.014398487291994,22.984780963158407
|
||||
BatchPrefillWithPagedKVCacheWrapper,16,8192,32,4,256,759.7649706666668,3.5331380935403933,23.15477380982632
|
||||
BatchPrefillWithPagedKVCacheWrapper,16,16384,32,4,256,3028.668266666667,1.77263029400997,23.234219789647476
|
||||
BatchPrefillWithPagedKVCacheWrapper,64,1024,32,4,256,49.26948266666667,27.241554149868346,22.316281159572153
|
||||
BatchPrefillWithPagedKVCacheWrapper,64,4096,32,4,256,763.6229333333335,7.030576067909256,23.037791659325052
|
||||
BatchPrefillWithPagedKVCacheWrapper,64,8192,32,4,256,3037.7449386666663,3.534667477616765,23.16479678130923
|
||||
BatchPrefillWithPagedKVCacheWrapper,64,16384,32,4,256,12110.653866666667,1.7732185822854112,23.241930601731337
|
||||
|
|
|
@ -1,49 +0,0 @@
|
|||
api,batch_size,seq_len,num_qo_heads,num_kv_heads,head_dim_qk,head_dim_vo,time_ms,bandwidth_GB_s,tflops
|
||||
BatchPrefillWithRaggedKVCacheWrapper,1,1024,32,4,128,128,0.031580159999999996,66.66666666666667,272.00415045395596
|
||||
BatchPrefillWithRaggedKVCacheWrapper,1,4096,32,4,128,128,0.0424448,197.82870928829917,3238.0634016887816
|
||||
BatchPrefillWithRaggedKVCacheWrapper,1,8192,32,4,128,128,0.057313279999999994,292.871180989816,9592.119206717885
|
||||
BatchPrefillWithRaggedKVCacheWrapper,1,16384,32,4,128,128,0.06972416000000001,481.36290204141574,31538.89922161844
|
||||
BatchPrefillWithRaggedKVCacheWrapper,4,1024,32,4,128,128,0.04327423999999998,194.60482725982024,793.9998106956938
|
||||
BatchPrefillWithRaggedKVCacheWrapper,4,4096,32,4,128,128,0.06579199999999998,510.5058365758757,8355.967501945528
|
||||
BatchPrefillWithRaggedKVCacheWrapper,4,8192,32,4,128,128,0.09618432000000002,698.0517406579366,22862.596060896405
|
||||
BatchPrefillWithRaggedKVCacheWrapper,4,16384,32,4,128,128,0.15411199999999997,871.12292358804,57075.97735548174
|
||||
BatchPrefillWithRaggedKVCacheWrapper,16,1024,32,4,128,128,0.07452671999999999,451.99230557845567,1844.1567463588901
|
||||
BatchPrefillWithRaggedKVCacheWrapper,16,4096,32,4,128,128,0.1668906666666667,805.0108653969065,13176.43041083983
|
||||
BatchPrefillWithRaggedKVCacheWrapper,16,8192,32,4,128,128,0.2874026666666667,934.46080760095,30605.46766745843
|
||||
BatchPrefillWithRaggedKVCacheWrapper,16,16384,32,4,128,128,0.5342506666666667,1005.1498622369525,65857.42289917343
|
||||
BatchPrefillWithRaggedKVCacheWrapper,64,1024,32,4,128,128,0.15733333333333333,856.4111186440679,3494.2106814915255
|
||||
BatchPrefillWithRaggedKVCacheWrapper,64,4096,32,4,128,128,0.5614719999999999,957.1184315513509,15666.129428017784
|
||||
BatchPrefillWithRaggedKVCacheWrapper,64,8192,32,4,128,128,1.1031466666666667,973.8198414233224,31894.55505055115
|
||||
BatchPrefillWithRaggedKVCacheWrapper,64,16384,32,4,128,128,2.1813759999999998,984.7032038493136,64517.75776176505
|
||||
BatchPrefillWithRaggedKVCacheWrapper,1,1024,32,4,192,128,0.03564544000000001,73.88681413386956,301.22838264866414
|
||||
BatchPrefillWithRaggedKVCacheWrapper,1,4096,32,4,192,128,0.04922368,213.27231121281466,3490.1635115456625
|
||||
BatchPrefillWithRaggedKVCacheWrapper,1,8192,32,4,192,128,0.061327359999999984,342.16062781766584,11205.353815328106
|
||||
BatchPrefillWithRaggedKVCacheWrapper,1,16384,32,4,192,128,0.08377343999999999,500.8189707859675,32812.059161471705
|
||||
BatchPrefillWithRaggedKVCacheWrapper,4,1024,32,4,192,128,0.049623040000000056,212.29880313660726,865.5187783739157
|
||||
BatchPrefillWithRaggedKVCacheWrapper,4,4096,32,4,192,128,0.08634367999999998,486.3377609108161,7958.831119544594
|
||||
BatchPrefillWithRaggedKVCacheWrapper,4,8192,32,4,192,128,0.13644799999999999,615.1444652908068,20145.249981238278
|
||||
BatchPrefillWithRaggedKVCacheWrapper,4,16384,32,4,192,128,0.2321706666666666,722.8359827253516,47357.904577781876
|
||||
BatchPrefillWithRaggedKVCacheWrapper,16,1024,32,4,192,128,0.09042944,465.99479107688825,1899.8093081191257
|
||||
BatchPrefillWithRaggedKVCacheWrapper,16,4096,32,4,192,128,0.3087573333333334,544.0154770952807,8902.716705589717
|
||||
BatchPrefillWithRaggedKVCacheWrapper,16,8192,32,4,192,128,0.5995946666666665,559.9464882943145,18337.58185156195
|
||||
BatchPrefillWithRaggedKVCacheWrapper,16,16384,32,4,192,128,1.1809706666666668,568.4182232017052,37240.94624227753
|
||||
BatchPrefillWithRaggedKVCacheWrapper,64,1024,32,4,192,128,0.2555306666666667,659.6413424611787,2689.2849156787443
|
||||
BatchPrefillWithRaggedKVCacheWrapper,64,4096,32,4,192,128,0.9085866666666667,739.472740079831,12101.340115520075
|
||||
BatchPrefillWithRaggedKVCacheWrapper,64,8192,32,4,192,128,1.7810773333333334,754.017631276351,24693.1810808739
|
||||
BatchPrefillWithRaggedKVCacheWrapper,64,16384,32,4,192,128,3.5260586666666662,761.5134193267346,49891.92667360423
|
||||
BatchPrefillWithRaggedKVCacheWrapper,1,1024,32,4,256,256,0.044037119999999964,95.61678874549479,390.12245087780525
|
||||
BatchPrefillWithRaggedKVCacheWrapper,1,4096,32,4,256,256,0.08118271999999997,206.86175580222005,3385.916448032292
|
||||
BatchPrefillWithRaggedKVCacheWrapper,1,8192,32,4,256,256,0.11204607999999996,299.6161579235972,9813.030743922503
|
||||
BatchPrefillWithRaggedKVCacheWrapper,1,16384,32,4,256,256,0.14619648000000002,459.1440778875113,30083.12177628353
|
||||
BatchPrefillWithRaggedKVCacheWrapper,4,1024,32,4,256,256,0.07792639999999999,216.1366622864652,881.8510381077531
|
||||
BatchPrefillWithRaggedKVCacheWrapper,4,4096,32,4,256,256,0.13784064000000001,487.3337790654483,7976.686902904687
|
||||
BatchPrefillWithRaggedKVCacheWrapper,4,8192,32,4,256,256,0.22408533333333336,599.2505712109672,19626.65938766184
|
||||
BatchPrefillWithRaggedKVCacheWrapper,4,16384,32,4,256,256,0.3959893333333334,678.0510720827496,44425.908890852275
|
||||
BatchPrefillWithRaggedKVCacheWrapper,16,1024,32,4,256,256,0.15150079999999996,444.6907739101049,1814.366042581954
|
||||
BatchPrefillWithRaggedKVCacheWrapper,16,4096,32,4,256,256,0.4274346666666664,628.6284687562392,10289.400589339195
|
||||
BatchPrefillWithRaggedKVCacheWrapper,16,8192,32,4,256,256,0.7913173333333334,678.7833823093305,22231.518637802277
|
||||
BatchPrefillWithRaggedKVCacheWrapper,16,16384,32,4,256,256,1.5360853333333337,699.1824898616742,45810.43946625186
|
||||
BatchPrefillWithRaggedKVCacheWrapper,64,1024,32,4,256,256,0.43906133333333336,613.773091686507,2504.2324256352945
|
||||
BatchPrefillWithRaggedKVCacheWrapper,64,4096,32,4,256,256,1.6363946666666664,656.8039006075145,10750.576497692491
|
||||
BatchPrefillWithRaggedKVCacheWrapper,64,8192,32,4,256,256,3.234005333333333,664.3564257160574,21759.006842803803
|
||||
BatchPrefillWithRaggedKVCacheWrapper,64,16384,32,4,256,256,6.420821333333334,669.0757535484556,43837.84598543405
|
||||
|
Binary file not shown.
|
Binary file not shown.
|
Binary file not shown.
|
Binary file not shown.
|
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
|
|
@ -1,16 +0,0 @@
|
|||
{
|
||||
"id": 193,
|
||||
"displayId": 20001,
|
||||
"type": "Traditional",
|
||||
"isPublic": false,
|
||||
"locales": [
|
||||
"zh_CN"
|
||||
],
|
||||
"samples": [
|
||||
{
|
||||
"inputData": "1\n",
|
||||
"outputData": ""
|
||||
}
|
||||
],
|
||||
"problemTagIds": []
|
||||
}
|
||||
|
|
@ -1,307 +0,0 @@
|
|||
from __future__ import annotations
|
||||
|
||||
|
||||
HEAD_DIM_CONFIGS = [(128, 128)]
|
||||
BATCH_SIZES = [1, 4, 16]
|
||||
SEQ_LENS = [1024, 4096, 8192, 16384]
|
||||
NUM_QO_HEADS = 32
|
||||
NUM_KV_HEADS = 4
|
||||
CAUSAL = 1
|
||||
|
||||
|
||||
def _build_cases():
|
||||
cases = []
|
||||
for head_dim_qk, head_dim_vo in HEAD_DIM_CONFIGS:
|
||||
for batch_size in BATCH_SIZES:
|
||||
for seq_len in SEQ_LENS:
|
||||
cases.append(
|
||||
(
|
||||
batch_size,
|
||||
seq_len,
|
||||
NUM_QO_HEADS,
|
||||
NUM_KV_HEADS,
|
||||
head_dim_qk,
|
||||
head_dim_vo,
|
||||
CAUSAL,
|
||||
)
|
||||
)
|
||||
return cases
|
||||
|
||||
|
||||
TESTCASES = _build_cases()
|
||||
|
||||
|
||||
def getNumOfTestcases() -> int:
|
||||
return len(TESTCASES)
|
||||
|
||||
|
||||
try:
|
||||
from pathlib import Path
|
||||
from typing import List, Tuple, Union
|
||||
import sys
|
||||
|
||||
import torch
|
||||
|
||||
KernelArg = Union[torch.Tensor, int, float]
|
||||
CURRENT_CASE = None
|
||||
|
||||
def _ensure_flashinfer_importable():
|
||||
try:
|
||||
import flashinfer # noqa: F401
|
||||
|
||||
return
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
here = Path(__file__).resolve()
|
||||
for parent in here.parents:
|
||||
candidate = parent / "McFlashInfer"
|
||||
if (candidate / "flashinfer").is_dir():
|
||||
sys.path.insert(0, str(candidate))
|
||||
return
|
||||
|
||||
def _get_testcase_index() -> int:
|
||||
try:
|
||||
raw = input().strip()
|
||||
except EOFError:
|
||||
return 0
|
||||
if raw == "":
|
||||
return 0
|
||||
try:
|
||||
testcase_id = int(raw.split()[0])
|
||||
except ValueError:
|
||||
return 0
|
||||
if 1 <= testcase_id <= len(TESTCASES):
|
||||
return testcase_id - 1
|
||||
if 0 <= testcase_id < len(TESTCASES):
|
||||
return testcase_id
|
||||
return 0
|
||||
|
||||
def _compute_reps(batch_size: int, seq_len: int, head_dim: int, base_reps: int = 100) -> int:
|
||||
workload = batch_size * seq_len * head_dim
|
||||
if workload < 1e5:
|
||||
return base_reps
|
||||
if workload < 1e6:
|
||||
return base_reps // 2
|
||||
if workload < 1e7:
|
||||
return base_reps // 4
|
||||
if workload < 1e8:
|
||||
return base_reps // 8
|
||||
if workload < 1e9:
|
||||
return base_reps // 16
|
||||
return base_reps // 32
|
||||
|
||||
def getTestCaseSize() -> Tuple[List[Tuple[int, ...]], Tuple[int, int]]:
|
||||
testcase_id = _get_testcase_index()
|
||||
global CURRENT_CASE
|
||||
batch_size, seq_len, num_qo_heads, num_kv_heads, head_dim_qk, head_dim_vo, causal = TESTCASES[testcase_id]
|
||||
CURRENT_CASE = (
|
||||
batch_size,
|
||||
seq_len,
|
||||
num_qo_heads,
|
||||
num_kv_heads,
|
||||
head_dim_qk,
|
||||
head_dim_vo,
|
||||
causal,
|
||||
20260610 + testcase_id,
|
||||
)
|
||||
qo_len = batch_size * seq_len
|
||||
kv_len = batch_size * seq_len
|
||||
warmup = 3
|
||||
iters = max(1, _compute_reps(batch_size, seq_len, head_dim_qk + head_dim_vo))
|
||||
return [
|
||||
(qo_len, num_qo_heads, head_dim_qk),
|
||||
(kv_len, num_kv_heads, head_dim_qk),
|
||||
(kv_len, num_kv_heads, head_dim_vo),
|
||||
(qo_len, num_qo_heads, head_dim_vo),
|
||||
(batch_size + 1,),
|
||||
(batch_size + 1,),
|
||||
(), (), (), (), (), (), (),
|
||||
], (warmup, iters)
|
||||
|
||||
def genTestCase(testcase_sizes, device: str = "cuda") -> List[KernelArg]:
|
||||
del testcase_sizes
|
||||
batch_size, seq_len, num_qo_heads, num_kv_heads, head_dim_qk, head_dim_vo, causal, seed = CURRENT_CASE
|
||||
gen = torch.Generator(device=device)
|
||||
gen.manual_seed(seed)
|
||||
dtype = torch.bfloat16
|
||||
qo_len = batch_size * seq_len
|
||||
q = torch.rand(
|
||||
qo_len,
|
||||
num_qo_heads,
|
||||
head_dim_qk,
|
||||
dtype=dtype,
|
||||
device=device,
|
||||
generator=gen,
|
||||
).contiguous()
|
||||
kv_len = batch_size * seq_len
|
||||
k = torch.rand(
|
||||
kv_len,
|
||||
num_kv_heads,
|
||||
head_dim_qk,
|
||||
dtype=dtype,
|
||||
device=device,
|
||||
generator=gen,
|
||||
).contiguous()
|
||||
v = torch.rand(
|
||||
kv_len,
|
||||
num_kv_heads,
|
||||
head_dim_vo,
|
||||
dtype=dtype,
|
||||
device=device,
|
||||
generator=gen,
|
||||
).contiguous()
|
||||
output = torch.empty(
|
||||
qo_len,
|
||||
num_qo_heads,
|
||||
head_dim_vo,
|
||||
dtype=dtype,
|
||||
device=device,
|
||||
)
|
||||
qo_indptr = torch.arange(0, batch_size + 1, dtype=torch.int32, device=device) * seq_len
|
||||
kv_indptr = qo_indptr.clone()
|
||||
return [
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
output,
|
||||
qo_indptr,
|
||||
kv_indptr,
|
||||
batch_size,
|
||||
seq_len,
|
||||
num_qo_heads,
|
||||
num_kv_heads,
|
||||
head_dim_qk,
|
||||
head_dim_vo,
|
||||
causal,
|
||||
]
|
||||
|
||||
def baseline(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
output,
|
||||
qo_indptr,
|
||||
kv_indptr,
|
||||
batch_size,
|
||||
seq_len,
|
||||
num_qo_heads,
|
||||
num_kv_heads,
|
||||
head_dim_qk,
|
||||
head_dim_vo,
|
||||
causal,
|
||||
):
|
||||
_ensure_flashinfer_importable()
|
||||
import flashinfer
|
||||
|
||||
workspace_buffer = torch.empty(128 * 1024 * 1024, dtype=torch.uint8, device=q.device)
|
||||
wrapper = flashinfer.BatchPrefillWithRaggedKVCacheWrapper(
|
||||
workspace_buffer,
|
||||
kv_layout="NHD",
|
||||
backend="auto",
|
||||
)
|
||||
wrapper.plan(
|
||||
qo_indptr,
|
||||
kv_indptr,
|
||||
int(num_qo_heads),
|
||||
int(num_kv_heads),
|
||||
int(head_dim_qk),
|
||||
int(head_dim_vo),
|
||||
causal=bool(causal),
|
||||
q_data_type=torch.bfloat16,
|
||||
kv_data_type=torch.bfloat16,
|
||||
)
|
||||
wrapper.run(q, k, v, out=output)
|
||||
return [
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
output,
|
||||
qo_indptr,
|
||||
kv_indptr,
|
||||
batch_size,
|
||||
seq_len,
|
||||
num_qo_heads,
|
||||
num_kv_heads,
|
||||
head_dim_qk,
|
||||
head_dim_vo,
|
||||
causal,
|
||||
]
|
||||
|
||||
def check(
|
||||
testcase_sizes,
|
||||
original_input_tensors,
|
||||
target_kernel_input_tensors,
|
||||
baseline_input_tensors,
|
||||
rtol=1e-2,
|
||||
atol=1e-2,
|
||||
) -> bool:
|
||||
del testcase_sizes, original_input_tensors
|
||||
output_t = target_kernel_input_tensors[3]
|
||||
output_ref = baseline_input_tensors[3]
|
||||
if output_t.shape != output_ref.shape:
|
||||
print(f"[FAIL] shape mismatch: target {output_t.shape}, ref {output_ref.shape}", file=sys.stderr)
|
||||
return False
|
||||
if output_t.dtype != output_ref.dtype:
|
||||
print(f"[FAIL] dtype mismatch: target {output_t.dtype}, ref {output_ref.dtype}", file=sys.stderr)
|
||||
return False
|
||||
if not torch.allclose(output_t.float(), output_ref.float(), rtol=rtol, atol=atol):
|
||||
diff = (output_t.float() - output_ref.float()).abs()
|
||||
print(
|
||||
f"[FAIL] allclose failed: max_abs_diff={float(diff.max().item()):.6f}, "
|
||||
f"mean_abs_diff={float(diff.mean().item()):.6f} (rtol={rtol}, atol={atol})",
|
||||
file=sys.stderr,
|
||||
)
|
||||
return False
|
||||
return True
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
INPUT_CLASS = [
|
||||
"INPUT",
|
||||
"INPUT",
|
||||
"INPUT",
|
||||
"OUTPUT",
|
||||
"INPUT",
|
||||
"INPUT",
|
||||
"INPUT",
|
||||
"INPUT",
|
||||
"INPUT",
|
||||
"INPUT",
|
||||
"INPUT",
|
||||
"INPUT",
|
||||
"INPUT",
|
||||
]
|
||||
|
||||
|
||||
def getWorkload(testcase_sizes) -> dict:
|
||||
raw_sizes = testcase_sizes[0] if isinstance(testcase_sizes, tuple) and len(testcase_sizes) == 2 else testcase_sizes
|
||||
q_shape, k_shape, v_shape, output_shape, qo_indptr_shape, kv_indptr_shape = raw_sizes[:6]
|
||||
qo_len, num_qo_heads, head_dim_qk = q_shape
|
||||
kv_len, num_kv_heads, k_dim = k_shape
|
||||
v_len, v_heads, head_dim_vo = v_shape
|
||||
assert k_dim == head_dim_qk
|
||||
assert v_len == kv_len
|
||||
assert v_heads == num_kv_heads
|
||||
assert qo_len == kv_len
|
||||
assert output_shape == (qo_len, num_qo_heads, head_dim_vo)
|
||||
assert qo_indptr_shape == kv_indptr_shape
|
||||
batch_size = qo_indptr_shape[0] - 1
|
||||
seq_len = kv_len // batch_size
|
||||
flops = batch_size * seq_len * seq_len * num_qo_heads * (head_dim_qk + head_dim_vo)
|
||||
memory_bytes = (
|
||||
qo_len * num_qo_heads * head_dim_qk * 2
|
||||
+ kv_len * num_kv_heads * head_dim_qk * 2
|
||||
+ kv_len * num_kv_heads * head_dim_vo * 2
|
||||
+ qo_len * num_qo_heads * head_dim_vo * 2
|
||||
+ (batch_size + 1) * 4 * 2
|
||||
)
|
||||
return {
|
||||
"flops": flops,
|
||||
"memory_bytes": memory_bytes,
|
||||
"dtype": "bf16",
|
||||
}
|
||||
|
||||
|
||||
DESIGNED_VRAM_SIZE = 48
|
||||
|
|
@ -1,23 +0,0 @@
|
|||
---
|
||||
sectionTitle: "题目描述"
|
||||
type: "Text"
|
||||
---
|
||||
你需要实现 FlashInfer ragged KV cache prefill 的CUDA C++前向算子。
|
||||
|
||||
本题输入采用 FlashInfer `BatchPrefillWithRaggedKVCacheWrapper` 的 ragged `NHD` 布局。每个 batch 中有 `seq_len` 个 query token,KV cache 中也有 `seq_len` 个 token:
|
||||
|
||||
其中 query heads 采用 GQA 布局:`num_qo_heads` 个 query/output heads 共享 `num_kv_heads` 个 KV heads,`G = num_qo_heads / num_kv_heads`。
|
||||
|
||||
评测程序会调用你提交代码中的 `run_kernel` 函数。你需要根据 `qo_indptr` 和 `kv_indptr` 读取 ragged Q/K/V,并将结果写入 `output`。
|
||||
|
||||
baseline 使用 FlashInfer ragged prefill 的 Python API:
|
||||
|
||||
```python
|
||||
wrapper = flashinfer.BatchPrefillWithRaggedKVCacheWrapper(workspace, kv_layout="NHD", backend="auto")
|
||||
wrapper.plan(qo_indptr, kv_indptr, num_qo_heads, num_kv_heads,
|
||||
head_dim_qk, head_dim_vo, causal=True,
|
||||
q_data_type=torch.bfloat16, kv_data_type=torch.bfloat16)
|
||||
wrapper.run(q, k, v, out=output)
|
||||
```
|
||||
|
||||
如何提交代码详见[评测指南](/d/2)。
|
||||
|
|
@ -1,41 +0,0 @@
|
|||
---
|
||||
sectionTitle: "接口约定"
|
||||
type: "codeSample"
|
||||
lang: "cuda"
|
||||
---
|
||||
你必须在提交的 CUDA 源码中提供如下 **C 符号**,函数名、参数类型、顺序必须完全一致,并使用 `extern "C"` 防止 name mangling:
|
||||
|
||||
```cpp
|
||||
#include <stdint.h>
|
||||
#include <cuda_bf16.h>
|
||||
|
||||
extern "C" void run_kernel(
|
||||
const __nv_bfloat16* q,
|
||||
const __nv_bfloat16* k,
|
||||
const __nv_bfloat16* v,
|
||||
__nv_bfloat16* output,
|
||||
const int32_t* qo_indptr,
|
||||
const int32_t* kv_indptr,
|
||||
int64_t batch_size,
|
||||
int64_t seq_len,
|
||||
int64_t num_qo_heads,
|
||||
int64_t num_kv_heads,
|
||||
int64_t head_dim_qk,
|
||||
int64_t head_dim_vo,
|
||||
int64_t causal
|
||||
);
|
||||
```
|
||||
|
||||
### 参数说明
|
||||
|
||||
* `q`:query tensor,shape `(batch_size * seq_len, num_qo_heads, head_dim_qk)`,连续 `bf16`
|
||||
* `k`:key tensor,shape `(batch_size * seq_len, num_kv_heads, head_dim_qk)`,连续 `bf16`
|
||||
* `v`:value tensor,shape `(batch_size * seq_len, num_kv_heads, head_dim_vo)`,连续 `bf16`
|
||||
* `output`:输出缓冲区,shape `(batch_size * seq_len, num_qo_heads, head_dim_vo)`,连续 `bf16`
|
||||
* `qo_indptr`:query/output ragged indptr,shape `(batch_size + 1)`,连续 `int32`
|
||||
* `kv_indptr`:KV ragged indptr,shape `(batch_size + 1)`,连续 `int32`
|
||||
* `causal`:是否启用 causal mask,评测中固定为 `1`
|
||||
|
||||
本题测试中 `qo_indptr[b + 1] - qo_indptr[b] == seq_len`,`kv_indptr[b + 1] - kv_indptr[b] == seq_len`。
|
||||
|
||||
`run_kernel` 内部需要自行计算合适的 launch 配置并启动 CUDA kernel。为保证计时准确,不建议在 `run_kernel` 内部做 `cudaDeviceSynchronize()` 或显式同步。
|
||||
|
|
@ -1,52 +0,0 @@
|
|||
---
|
||||
sectionTitle: "接口约定"
|
||||
type: "codeSample"
|
||||
lang: "tilelang"
|
||||
---
|
||||
你必须在提交的 Python 代码中提供 `run_kernel` 函数,函数名、参数顺序、类型必须完全一致:
|
||||
|
||||
```python
|
||||
import tilelang
|
||||
import tilelang.language as T
|
||||
from tilelang import jit
|
||||
|
||||
real_kernel = None
|
||||
|
||||
@jit
|
||||
def build_kernel(*args):
|
||||
@T.prim_func
|
||||
def kernel(*args):
|
||||
...
|
||||
return kernel
|
||||
|
||||
def run_kernel(
|
||||
q, # Tensor[bf16], shape (batch_size * seq_len, num_qo_heads, head_dim_qk)
|
||||
k, # Tensor[bf16], shape (batch_size * seq_len, num_kv_heads, head_dim_qk)
|
||||
v, # Tensor[bf16], shape (batch_size * seq_len, num_kv_heads, head_dim_vo)
|
||||
output, # Tensor[bf16], shape (batch_size * seq_len, num_qo_heads, head_dim_vo)
|
||||
qo_indptr, # Tensor[int32], shape (batch_size + 1)
|
||||
kv_indptr, # Tensor[int32], shape (batch_size + 1)
|
||||
batch_size, # int64
|
||||
seq_len, # int64
|
||||
num_qo_heads, # int64
|
||||
num_kv_heads, # int64
|
||||
head_dim_qk, # int64
|
||||
head_dim_vo, # int64
|
||||
causal, # int64
|
||||
):
|
||||
global real_kernel
|
||||
if real_kernel is None:
|
||||
real_kernel = build_kernel(...)
|
||||
real_kernel(q, k, v, output, qo_indptr, kv_indptr,
|
||||
batch_size, seq_len, num_qo_heads, num_kv_heads,
|
||||
head_dim_qk, head_dim_vo, causal)
|
||||
```
|
||||
|
||||
### 参数说明
|
||||
|
||||
* `q/k/v`:FlashInfer ragged prefill 输入 tensor,连续 `bfloat16`
|
||||
* `output`:输出缓冲区,连续 `bfloat16`,需要写入结果
|
||||
* `qo_indptr/kv_indptr`:ragged indptr,连续 `int32`
|
||||
* `causal`:是否启用 causal mask,评测中固定为 `1`
|
||||
|
||||
`run_kernel` 内部需要自行计算合适的 grid/block,并 launch 你实现的 TileLang kernel。
|
||||
|
|
@ -1,41 +0,0 @@
|
|||
---
|
||||
sectionTitle: "接口约定"
|
||||
type: "codeSample"
|
||||
lang: "triton"
|
||||
---
|
||||
你必须在提交的 Python 代码中提供 `run_kernel` 函数,函数名、参数顺序、类型必须完全一致:
|
||||
|
||||
```python
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
@triton.jit
|
||||
def your_kernel(...):
|
||||
...
|
||||
|
||||
def run_kernel(
|
||||
q, # Tensor[bf16], shape (batch_size * seq_len, num_qo_heads, head_dim_qk)
|
||||
k, # Tensor[bf16], shape (batch_size * seq_len, num_kv_heads, head_dim_qk)
|
||||
v, # Tensor[bf16], shape (batch_size * seq_len, num_kv_heads, head_dim_vo)
|
||||
output, # Tensor[bf16], shape (batch_size * seq_len, num_qo_heads, head_dim_vo)
|
||||
qo_indptr, # Tensor[int32], shape (batch_size + 1)
|
||||
kv_indptr, # Tensor[int32], shape (batch_size + 1)
|
||||
batch_size, # int64
|
||||
seq_len, # int64
|
||||
num_qo_heads, # int64
|
||||
num_kv_heads, # int64
|
||||
head_dim_qk, # int64
|
||||
head_dim_vo, # int64
|
||||
causal, # int64
|
||||
):
|
||||
...
|
||||
```
|
||||
|
||||
### 参数说明
|
||||
|
||||
* `q/k/v`:FlashInfer ragged prefill 输入 tensor,连续 `bfloat16`
|
||||
* `output`:输出缓冲区,连续 `bfloat16`,需要写入结果
|
||||
* `qo_indptr/kv_indptr`:ragged indptr,连续 `int32`
|
||||
* `causal`:是否启用 causal mask,评测中固定为 `1`
|
||||
|
||||
`run_kernel` 内部需要自行计算合适的 grid/block,并 launch 你实现的 Triton kernel。
|
||||
|
|
@ -1,26 +0,0 @@
|
|||
---
|
||||
sectionTitle: "输入格式"
|
||||
type: "Text"
|
||||
---
|
||||
## 数据范围
|
||||
|
||||
- 数据类型:`q/kv_data/output` 均为 `bfloat16`
|
||||
- KV layout:`NHD`
|
||||
- `num_qo_heads = 32`
|
||||
- `num_kv_heads = 4`
|
||||
- `page_block_size = 16`
|
||||
- `causal = 0`
|
||||
- `head_dim` 固定为 `128`
|
||||
- `batch_size` 取值为 `1, 4, 16`
|
||||
- `seq_len` 取值为 `1024, 4096, 8192, 16384`
|
||||
|
||||
测试点顺序与 `McFlashInfer/benchmarks/bench_batch_prefill_paged.py` 中的 cases 一致,即:
|
||||
|
||||
```python
|
||||
for head_dim in [128]:
|
||||
for batch_size in [1, 4, 16]:
|
||||
for seq_len in [1024, 4096, 8192, 16384]:
|
||||
...
|
||||
```
|
||||
|
||||
输出参与 `torch.allclose` 校验,容差为 `rtol=1e-2, atol=1e-2`。
|
||||
|
|
@ -1,28 +0,0 @@
|
|||
pytorch参考实现:
|
||||
|
||||
```python
|
||||
def baseline(q, kv_data, output, qo_indptr, kv_indptr, kv_indices, last_page_len,
|
||||
batch_size, seq_len, num_qo_heads, num_kv_heads,
|
||||
head_dim, page_block_size, causal):
|
||||
workspace_buffer = torch.empty(128 * 1024 * 1024, dtype=torch.uint8, device=q.device)
|
||||
wrapper = flashinfer.BatchPrefillWithPagedKVCacheWrapper(
|
||||
workspace_buffer,
|
||||
kv_layout="NHD",
|
||||
backend="auto",
|
||||
)
|
||||
wrapper.plan(
|
||||
qo_indptr,
|
||||
kv_indptr,
|
||||
kv_indices,
|
||||
last_page_len,
|
||||
num_qo_heads,
|
||||
num_kv_heads,
|
||||
head_dim,
|
||||
page_block_size,
|
||||
causal=bool(causal),
|
||||
q_data_type=torch.bfloat16,
|
||||
kv_data_type=torch.bfloat16,
|
||||
)
|
||||
wrapper.run(q, kv_data, out=output)
|
||||
```
|
||||
|
||||
|
|
@ -1 +0,0 @@
|
|||
FlashInfer Ragged Prefill
|
||||
|
|
@ -1,16 +0,0 @@
|
|||
{
|
||||
"id": 194,
|
||||
"displayId": 20002,
|
||||
"type": "Traditional",
|
||||
"isPublic": false,
|
||||
"locales": [
|
||||
"zh_CN"
|
||||
],
|
||||
"samples": [
|
||||
{
|
||||
"inputData": "1\n",
|
||||
"outputData": ""
|
||||
}
|
||||
],
|
||||
"problemTagIds": []
|
||||
}
|
||||
|
|
@ -1,273 +0,0 @@
|
|||
from __future__ import annotations
|
||||
|
||||
|
||||
HEAD_DIMS = [128]
|
||||
BATCH_SIZES = [1, 4, 16]
|
||||
SEQ_LENS = [1024, 4096, 8192, 16384]
|
||||
NUM_QO_HEADS = 32
|
||||
PAGE_BLOCK_SIZE = 16
|
||||
CAUSAL = 0
|
||||
|
||||
|
||||
def _build_cases():
|
||||
cases = []
|
||||
for head_dim in HEAD_DIMS:
|
||||
for batch_size in BATCH_SIZES:
|
||||
for seq_len in SEQ_LENS:
|
||||
num_kv_heads = 8 if head_dim == 64 else 4
|
||||
cases.append(
|
||||
(
|
||||
batch_size,
|
||||
seq_len,
|
||||
NUM_QO_HEADS,
|
||||
num_kv_heads,
|
||||
head_dim,
|
||||
PAGE_BLOCK_SIZE,
|
||||
CAUSAL,
|
||||
)
|
||||
)
|
||||
return cases
|
||||
|
||||
|
||||
TESTCASES = _build_cases()
|
||||
|
||||
|
||||
def getNumOfTestcases() -> int:
|
||||
return len(TESTCASES)
|
||||
|
||||
|
||||
try:
|
||||
from pathlib import Path
|
||||
from typing import List, Tuple, Union
|
||||
import sys
|
||||
|
||||
import torch
|
||||
|
||||
KernelArg = Union[torch.Tensor, int, float]
|
||||
CURRENT_CASE = None
|
||||
|
||||
def _ensure_flashinfer_importable():
|
||||
try:
|
||||
import flashinfer # noqa: F401
|
||||
|
||||
return
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
here = Path(__file__).resolve()
|
||||
for parent in here.parents:
|
||||
candidate = parent / "McFlashInfer"
|
||||
if (candidate / "flashinfer").is_dir():
|
||||
sys.path.insert(0, str(candidate))
|
||||
return
|
||||
|
||||
def _get_testcase_index() -> int:
|
||||
try:
|
||||
raw = input().strip()
|
||||
except EOFError:
|
||||
return 0
|
||||
if raw == "":
|
||||
return 0
|
||||
try:
|
||||
testcase_id = int(raw.split()[0])
|
||||
except ValueError:
|
||||
return 0
|
||||
if 1 <= testcase_id <= len(TESTCASES):
|
||||
return testcase_id - 1
|
||||
if 0 <= testcase_id < len(TESTCASES):
|
||||
return testcase_id
|
||||
return 0
|
||||
|
||||
def _compute_reps(batch_size: int, seq_len: int, head_dim: int, base_reps: int = 100) -> int:
|
||||
workload = batch_size * seq_len * head_dim
|
||||
if workload < 1e5:
|
||||
return base_reps
|
||||
if workload < 1e6:
|
||||
return base_reps // 2
|
||||
if workload < 1e7:
|
||||
return base_reps // 4
|
||||
if workload < 1e8:
|
||||
return base_reps // 8
|
||||
if workload < 1e9:
|
||||
return base_reps // 16
|
||||
return base_reps // 32
|
||||
|
||||
def _setup_paged_kv_indptr(batch_size: int, seq_len: int, page_block_size: int, device: str):
|
||||
seq_lens = torch.full((batch_size,), seq_len, dtype=torch.int32, device=device)
|
||||
seq_lens_blocks = torch.div(seq_lens + page_block_size - 1, page_block_size, rounding_mode="floor")
|
||||
kv_indptr = torch.empty((batch_size + 1,), dtype=torch.int32, device=device)
|
||||
kv_indptr[0] = 0
|
||||
kv_indptr[1:] = torch.cumsum(seq_lens_blocks, dim=0)
|
||||
num_blocks = int(kv_indptr[-1].item())
|
||||
last_page_len = (seq_lens - 1) % page_block_size + 1
|
||||
return kv_indptr, last_page_len, num_blocks
|
||||
|
||||
def getTestCaseSize() -> Tuple[List[Tuple[int, ...]], Tuple[int, int]]:
|
||||
testcase_id = _get_testcase_index()
|
||||
global CURRENT_CASE
|
||||
batch_size, seq_len, num_qo_heads, num_kv_heads, head_dim, page_block_size, causal = TESTCASES[testcase_id]
|
||||
num_blocks = batch_size * ((seq_len + page_block_size - 1) // page_block_size)
|
||||
qo_len = batch_size * seq_len
|
||||
CURRENT_CASE = (
|
||||
batch_size,
|
||||
seq_len,
|
||||
num_qo_heads,
|
||||
num_kv_heads,
|
||||
head_dim,
|
||||
page_block_size,
|
||||
causal,
|
||||
num_blocks,
|
||||
20260620 + testcase_id,
|
||||
)
|
||||
warmup = 3
|
||||
iters = max(1, _compute_reps(batch_size, seq_len, head_dim))
|
||||
return [
|
||||
(qo_len, num_qo_heads, head_dim),
|
||||
(num_blocks, 2, page_block_size, num_kv_heads, head_dim),
|
||||
(qo_len, num_qo_heads, head_dim),
|
||||
(batch_size + 1,),
|
||||
(batch_size + 1,),
|
||||
(num_blocks,),
|
||||
(batch_size,),
|
||||
(), (), (), (), (), (), (),
|
||||
], (warmup, iters)
|
||||
|
||||
def genTestCase(testcase_sizes, device: str = "cuda") -> List[KernelArg]:
|
||||
del testcase_sizes
|
||||
(
|
||||
batch_size,
|
||||
seq_len,
|
||||
num_qo_heads,
|
||||
num_kv_heads,
|
||||
head_dim,
|
||||
page_block_size,
|
||||
causal,
|
||||
num_blocks,
|
||||
seed,
|
||||
) = CURRENT_CASE
|
||||
gen = torch.Generator(device=device)
|
||||
gen.manual_seed(seed)
|
||||
dtype = torch.bfloat16
|
||||
qo_len = batch_size * seq_len
|
||||
q = torch.rand(qo_len, num_qo_heads, head_dim, dtype=dtype, device=device, generator=gen).contiguous()
|
||||
kv_data = torch.randn(
|
||||
num_blocks,
|
||||
2,
|
||||
page_block_size,
|
||||
num_kv_heads,
|
||||
head_dim,
|
||||
dtype=dtype,
|
||||
device=device,
|
||||
generator=gen,
|
||||
).contiguous()
|
||||
output = torch.empty(qo_len, num_qo_heads, head_dim, dtype=dtype, device=device)
|
||||
qo_indptr = torch.arange(0, batch_size + 1, dtype=torch.int32, device=device) * seq_len
|
||||
kv_indptr, last_page_len, num_blocks_check = _setup_paged_kv_indptr(
|
||||
batch_size,
|
||||
seq_len,
|
||||
page_block_size,
|
||||
device,
|
||||
)
|
||||
assert num_blocks_check == num_blocks
|
||||
kv_indices = torch.arange(num_blocks, dtype=torch.int32, device=device)
|
||||
return [
|
||||
q,
|
||||
kv_data,
|
||||
output,
|
||||
qo_indptr,
|
||||
kv_indptr,
|
||||
kv_indices,
|
||||
last_page_len,
|
||||
batch_size,
|
||||
seq_len,
|
||||
num_qo_heads,
|
||||
num_kv_heads,
|
||||
head_dim,
|
||||
page_block_size,
|
||||
causal,
|
||||
]
|
||||
|
||||
def baseline(
|
||||
q,
|
||||
kv_data,
|
||||
output,
|
||||
qo_indptr,
|
||||
kv_indptr,
|
||||
kv_indices,
|
||||
last_page_len,
|
||||
batch_size,
|
||||
seq_len,
|
||||
num_qo_heads,
|
||||
num_kv_heads,
|
||||
head_dim,
|
||||
page_block_size,
|
||||
causal,
|
||||
):
|
||||
_ensure_flashinfer_importable()
|
||||
import flashinfer
|
||||
|
||||
workspace_buffer = torch.empty(128 * 1024 * 1024, dtype=torch.uint8, device=q.device)
|
||||
wrapper = flashinfer.BatchPrefillWithPagedKVCacheWrapper(
|
||||
workspace_buffer,
|
||||
kv_layout="NHD",
|
||||
backend="auto",
|
||||
)
|
||||
wrapper.plan(
|
||||
qo_indptr,
|
||||
kv_indptr,
|
||||
kv_indices,
|
||||
last_page_len,
|
||||
int(num_qo_heads),
|
||||
int(num_kv_heads),
|
||||
int(head_dim),
|
||||
int(page_block_size),
|
||||
causal=bool(causal),
|
||||
q_data_type=torch.bfloat16,
|
||||
kv_data_type=torch.bfloat16,
|
||||
)
|
||||
wrapper.run(q, kv_data, out=output)
|
||||
return [
|
||||
q,
|
||||
kv_data,
|
||||
output,
|
||||
qo_indptr,
|
||||
kv_indptr,
|
||||
kv_indices,
|
||||
last_page_len,
|
||||
batch_size,
|
||||
seq_len,
|
||||
num_qo_heads,
|
||||
num_kv_heads,
|
||||
head_dim,
|
||||
page_block_size,
|
||||
causal,
|
||||
]
|
||||
|
||||
def check(
|
||||
testcase_sizes,
|
||||
original_input_tensors,
|
||||
target_kernel_input_tensors,
|
||||
baseline_input_tensors,
|
||||
rtol=1e-2,
|
||||
atol=1e-2,
|
||||
) -> bool:
|
||||
del testcase_sizes, original_input_tensors
|
||||
output_t = target_kernel_input_tensors[2]
|
||||
output_ref = baseline_input_tensors[2]
|
||||
if output_t.shape != output_ref.shape:
|
||||
print(f"[FAIL] shape mismatch: target {output_t.shape}, ref {output_ref.shape}", file=sys.stderr)
|
||||
return False
|
||||
if output_t.dtype != output_ref.dtype:
|
||||
print(f"[FAIL] dtype mismatch: target {output_t.dtype}, ref {output_ref.dtype}", file=sys.stderr)
|
||||
return False
|
||||
if not torch.allclose(output_t.float(), output_ref.float(), rtol=rtol, atol=atol):
|
||||
diff = (output_t.float() - output_ref.float()).abs()
|
||||
print(
|
||||
f"[FAIL] allclose failed: max_abs_diff={float(diff.max().item()):.6f}, "
|
||||
f"mean_abs_diff={float(diff.mean().item()):.6f} (rtol={rtol}, atol={atol})",
|
||||
file=sys.stderr,
|
||||
)
|
||||
return False
|
||||
return True
|
||||
except Exception:
|
||||
pass
|
||||
|
|
@ -1,22 +0,0 @@
|
|||
---
|
||||
sectionTitle: "题目描述"
|
||||
type: "Text"
|
||||
---
|
||||
你需要实现 FlashInfer paged KV cache prefill 的 CUDA C++ 前向算子。
|
||||
|
||||
本题输入采用 FlashInfer `BatchPrefillWithPagedKVCacheWrapper` 的 paged `NHD` 布局。每个 batch 中有 `seq_len` 个 query token,KV cache 也有 `seq_len` 个 token,并按 page 存储。
|
||||
|
||||
评测程序会调用你提交代码中的 `run_kernel` 函数。你需要根据 `qo_indptr`、`kv_indptr`、`kv_indices` 和 `last_page_len` 读取 paged KV cache,并将结果写入 `output`。
|
||||
|
||||
baseline 使用 FlashInfer paged prefill 的 Python API:
|
||||
|
||||
```python
|
||||
wrapper = flashinfer.BatchPrefillWithPagedKVCacheWrapper(workspace, kv_layout="NHD", backend="auto")
|
||||
wrapper.plan(qo_indptr, kv_indptr, kv_indices, last_page_len,
|
||||
num_qo_heads, num_kv_heads, head_dim, page_block_size,
|
||||
causal=bool(causal),
|
||||
q_data_type=torch.bfloat16, kv_data_type=torch.bfloat16)
|
||||
wrapper.run(q, kv_data, out=output)
|
||||
```
|
||||
|
||||
如何提交代码详见[评测指南](/d/2)。
|
||||
|
|
@ -1,42 +0,0 @@
|
|||
---
|
||||
sectionTitle: "接口约定"
|
||||
type: "codeSample"
|
||||
lang: "cuda"
|
||||
---
|
||||
你必须在提交的 CUDA 源码中提供如下 **C 符号**,函数名、参数类型、顺序必须完全一致,并使用 `extern "C"` 防止 name mangling:
|
||||
|
||||
```cpp
|
||||
#include <stdint.h>
|
||||
#include <cuda_bf16.h>
|
||||
|
||||
extern "C" void run_kernel(
|
||||
const __nv_bfloat16* q,
|
||||
const __nv_bfloat16* kv_data,
|
||||
__nv_bfloat16* output,
|
||||
const int32_t* qo_indptr,
|
||||
const int32_t* kv_indptr,
|
||||
const int32_t* kv_indices,
|
||||
const int32_t* last_page_len,
|
||||
int64_t batch_size,
|
||||
int64_t seq_len,
|
||||
int64_t num_qo_heads,
|
||||
int64_t num_kv_heads,
|
||||
int64_t head_dim,
|
||||
int64_t page_block_size,
|
||||
int64_t causal
|
||||
);
|
||||
```
|
||||
|
||||
### 参数说明
|
||||
|
||||
* `q`:query tensor,shape `(batch_size * seq_len, num_qo_heads, head_dim)`,连续 `bf16`
|
||||
* `kv_data`:paged KV cache,shape `(num_blocks, 2, page_block_size, num_kv_heads, head_dim)`,连续 `bf16`,其中 `kv_data[:, 0]` 为 key,`kv_data[:, 1]` 为 value
|
||||
* `output`:输出缓冲区,shape `(batch_size * seq_len, num_qo_heads, head_dim)`,连续 `bf16`
|
||||
* `qo_indptr`:query/output indptr,shape `(batch_size + 1)`,连续 `int32`
|
||||
* `kv_indptr`:paged KV indptr,shape `(batch_size + 1)`,连续 `int32`
|
||||
* `kv_indices`:page index,shape `(num_blocks)`,连续 `int32`
|
||||
* `last_page_len`:每个 batch 最后一个 page 的有效 token 数,shape `(batch_size)`,连续 `int32`
|
||||
* `page_block_size`:page size,评测中固定为 `16`
|
||||
* `causal`:是否启用 causal mask,本题按 benchmark case 固定为 `0`
|
||||
|
||||
`run_kernel` 内部需要自行计算合适的 launch 配置并启动 CUDA kernel。为保证计时准确,不建议在 `run_kernel` 内部做 `cudaDeviceSynchronize()` 或显式同步。
|
||||
|
|
@ -1,55 +0,0 @@
|
|||
---
|
||||
sectionTitle: "接口约定"
|
||||
type: "codeSample"
|
||||
lang: "tilelang"
|
||||
---
|
||||
你必须在提交的 Python 代码中提供 `run_kernel` 函数,函数名、参数顺序、类型必须完全一致:
|
||||
|
||||
```python
|
||||
import tilelang
|
||||
import tilelang.language as T
|
||||
from tilelang import jit
|
||||
|
||||
real_kernel = None
|
||||
|
||||
@jit
|
||||
def build_kernel(*args):
|
||||
@T.prim_func
|
||||
def kernel(*args):
|
||||
...
|
||||
return kernel
|
||||
|
||||
def run_kernel(
|
||||
q, # Tensor[bf16], shape (batch_size * seq_len, num_qo_heads, head_dim)
|
||||
kv_data, # Tensor[bf16], shape (num_blocks, 2, page_block_size, num_kv_heads, head_dim)
|
||||
output, # Tensor[bf16], shape (batch_size * seq_len, num_qo_heads, head_dim)
|
||||
qo_indptr, # Tensor[int32], shape (batch_size + 1)
|
||||
kv_indptr, # Tensor[int32], shape (batch_size + 1)
|
||||
kv_indices, # Tensor[int32], shape (num_blocks)
|
||||
last_page_len, # Tensor[int32], shape (batch_size)
|
||||
batch_size, # int64
|
||||
seq_len, # int64
|
||||
num_qo_heads, # int64
|
||||
num_kv_heads, # int64
|
||||
head_dim, # int64
|
||||
page_block_size, # int64
|
||||
causal, # int64
|
||||
):
|
||||
global real_kernel
|
||||
if real_kernel is None:
|
||||
real_kernel = build_kernel(...)
|
||||
real_kernel(q, kv_data, output, qo_indptr, kv_indptr, kv_indices, last_page_len,
|
||||
batch_size, seq_len, num_qo_heads, num_kv_heads,
|
||||
head_dim, page_block_size, causal)
|
||||
```
|
||||
|
||||
### 参数说明
|
||||
|
||||
* `q`:query tensor,连续 `bfloat16`
|
||||
* `kv_data`:paged KV cache,连续 `bfloat16`
|
||||
* `output`:输出缓冲区,连续 `bfloat16`,需要写入结果
|
||||
* `qo_indptr/kv_indptr/kv_indices/last_page_len`:paged KV metadata,连续 `int32`
|
||||
* `page_block_size`:评测中固定为 `16`
|
||||
* `causal`:评测中固定为 `0`
|
||||
|
||||
`run_kernel` 内部需要自行计算合适的 grid/block,并 launch 你实现的 TileLang kernel。
|
||||
|
|
@ -1,44 +0,0 @@
|
|||
---
|
||||
sectionTitle: "接口约定"
|
||||
type: "codeSample"
|
||||
lang: "triton"
|
||||
---
|
||||
你必须在提交的 Python 代码中提供 `run_kernel` 函数,函数名、参数顺序、类型必须完全一致:
|
||||
|
||||
```python
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
@triton.jit
|
||||
def your_kernel(...):
|
||||
...
|
||||
|
||||
def run_kernel(
|
||||
q, # Tensor[bf16], shape (batch_size * seq_len, num_qo_heads, head_dim)
|
||||
kv_data, # Tensor[bf16], shape (num_blocks, 2, page_block_size, num_kv_heads, head_dim)
|
||||
output, # Tensor[bf16], shape (batch_size * seq_len, num_qo_heads, head_dim)
|
||||
qo_indptr, # Tensor[int32], shape (batch_size + 1)
|
||||
kv_indptr, # Tensor[int32], shape (batch_size + 1)
|
||||
kv_indices, # Tensor[int32], shape (num_blocks)
|
||||
last_page_len, # Tensor[int32], shape (batch_size)
|
||||
batch_size, # int64
|
||||
seq_len, # int64
|
||||
num_qo_heads, # int64
|
||||
num_kv_heads, # int64
|
||||
head_dim, # int64
|
||||
page_block_size, # int64
|
||||
causal, # int64
|
||||
):
|
||||
...
|
||||
```
|
||||
|
||||
### 参数说明
|
||||
|
||||
* `q`:query tensor,连续 `bfloat16`
|
||||
* `kv_data`:paged KV cache,连续 `bfloat16`
|
||||
* `output`:输出缓冲区,连续 `bfloat16`,需要写入结果
|
||||
* `qo_indptr/kv_indptr/kv_indices/last_page_len`:paged KV metadata,连续 `int32`
|
||||
* `page_block_size`:评测中固定为 `16`
|
||||
* `causal`:评测中固定为 `0`
|
||||
|
||||
`run_kernel` 内部需要自行计算合适的 grid/block,并 launch 你实现的 Triton kernel。
|
||||
|
|
@ -1,26 +0,0 @@
|
|||
---
|
||||
sectionTitle: "输入格式"
|
||||
type: "Text"
|
||||
---
|
||||
## 数据范围
|
||||
|
||||
- 数据类型:`q/kv_data/output` 均为 `bfloat16`
|
||||
- KV layout:`NHD`
|
||||
- `num_qo_heads = 32`
|
||||
- `num_kv_heads = 4`
|
||||
- `page_block_size = 16`
|
||||
- `causal = 0`
|
||||
- `head_dim` 固定为 `128`
|
||||
- `batch_size` 取值为 `1, 4, 16`
|
||||
- `seq_len` 取值为 `1024, 4096, 8192, 16384`
|
||||
|
||||
测试点顺序与 `McFlashInfer/benchmarks/bench_batch_prefill_paged.py` 中的 cases 一致,即:
|
||||
|
||||
```python
|
||||
for head_dim in [128]:
|
||||
for batch_size in [1, 4, 16]:
|
||||
for seq_len in [1024, 4096, 8192, 16384]:
|
||||
...
|
||||
```
|
||||
|
||||
输出参与 `torch.allclose` 校验,容差为 `rtol=1e-2, atol=1e-2`。
|
||||
Some files were not shown because too many files have changed in this diff Show More
Loading…
Reference in New Issue