sgl-project/sglang

[Kernel] Optimize triton decoding kernels for long context

Open

#2,271 opened on Nov 30, 2024

View on GitHub
 (7 comments) (1 reaction) (1 assignee)Python (6,216 forks)auto 404
good first issuehelp wantedhigh priorityinactive

Repository metrics

Stars
 (28,442 stars)
PR merge metrics
 (Avg merge 2d 1h) (1,000 merged PRs in 30d)

Description

We noticed the current triton decoding kernel is very slow on long context. This is due to a missing flash decoding like optimization.

Reproduce

We test the decoding speed with a context length of 200 and 2,000.

triton backend: The decoding speed drops from 147.64 token/s to 126.41 token/s

$ python3 -m sglang.bench_offline_throughput --model meta-llama/Llama-3.1-8B-Instruct --dataset-name random --num-prompt 1 --random-input 128 --random-output 2048 --random-range 1 --attention-backend triton

[2024-11-30 05:10:04 TP0] Decode batch. #running-req: 1, #token: 234, token usage: 0.00, gen throughput (token/s): 147.64, #queue-req: 0
... 
[2024-11-30 05:10:18 TP0] Decode batch. #running-req: 1, #token: 2154, token usage: 0.00, gen throughput (token/s): 126.41, #queue-req: 0

flashinfer backend: The decoding speed only drops from 144.17 token/s to 143.35 token/s

$ python3 -m sglang.bench_offline_throughput --model meta-llama/Llama-3.1-8B-Instruct --dataset-name random --num-prompt 1 --random-input 128 --random-output 2048 --random-range 1

[2024-11-30 05:11:40 TP0] Decode batch. #running-req: 1, #token: 234, token usage: 0.00, gen throughput (token/s): 144.17, #queue-req: 0
...
[2024-11-30 05:11:54 TP0] Decode batch. #running-req: 1, #token: 2154, token usage: 0.00, gen throughput (token/s): 143.35, #queue-req: 0

Possible solutions

We can learn from the flash decoding triton kernel from lightllm and improve the current triton decoding kernel. Related links:

Contributor guide