Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 5 additions & 1 deletion lightllm/common/basemodel/attention/fa3/fp.py
Original file line number Diff line number Diff line change
Expand Up @@ -158,6 +158,7 @@ class Fa3DecodeAttState(BaseDecodeAttState):
b_att_seq_len: torch.Tensor = None
# 在是否开启mtp 的不同模式下,其设置不同的值,可以加速算子的运行。
decode_max_q_seq_len: int = None
decode_max_kv_seq_len: int = None
causal: bool = None

def init_state(self):
Expand Down Expand Up @@ -225,6 +226,8 @@ def _init_page_table(self, b_att_req_idx: torch.Tensor):
att_batch_size = b_att_req_idx.shape[0]
model = self.backend.model
actual_max_kv_len = self.infer_state.max_kv_seq_len
# Graph 捕获会将 infer_state.max_kv_seq_len 改为容量上限,提前保存真实长度用于 FA3 配置查找。
self.decode_max_kv_seq_len = actual_max_kv_len
page_table_width = actual_max_kv_len
if model.graph is not None and model.graph.can_run(
batch_size=self.infer_state.batch_size,
Expand Down Expand Up @@ -295,8 +298,9 @@ def _normal_decode_att(
page_table=self.page_table,
cache_seqlens=self.b_att_seq_len,
cu_seqlens_q=self.cu_seqlens_q,
cu_seqlens_k_new=self.cu_seqlens_k,
cu_seqlens_k_new=None, # KV 已提前写入缓存,此处不追加新的 K/V。
max_seqlen_q=self.decode_max_q_seq_len,
max_seqlen_k=self.decode_max_kv_seq_len,
softmax_scale=sm_scale,
causal=self.causal,
window_size=window_size,
Expand Down
5 changes: 5 additions & 0 deletions lightllm/common/basemodel/attention/triton/fp.py
Original file line number Diff line number Diff line change
Expand Up @@ -95,8 +95,11 @@ def _nomarl_prefill_att(
@dataclasses.dataclass
class TritonDecodeAttState(BaseDecodeAttState):
b_mark_mtp_shared_group: torch.Tensor = None
decode_max_kv_seq_len: int = None

def init_state(self):
# Graph 捕获会改写 infer_state 的长度上限,提前保存真实长度用于 GQA decode 配置查找。
self.decode_max_kv_seq_len = self.infer_state.max_kv_seq_len
draft_step = self.backend.model.mtp_manager.get_decode_draft_step(self.backend.model.is_mtp_draft_model)
if draft_step > 0:
self.b_mark_mtp_shared_group = build_mtp_shared_group_markers(
Expand Down Expand Up @@ -212,6 +215,7 @@ def _normal_decode_gqa_flash_decoding_att(
infer_state=self.infer_state,
cache_k=k,
cache_v=v,
max_len_in_batch=self.decode_max_kv_seq_len,
out=out,
alloc_tensor_func=alloc_func,
sliding_window=sliding_window,
Expand All @@ -238,6 +242,7 @@ def _spec_decode_gqa_att(
B_req_idx=self.infer_state.b_req_idx,
b_seq_len=self.infer_state.b_seq_len,
b_mark_shared_group=self.b_mark_mtp_shared_group,
max_kv_len=self.decode_max_kv_seq_len,
alloc_tensor_func=alloc_func,
)

Expand Down
6 changes: 5 additions & 1 deletion lightllm/common/basemodel/attention/triton/int4kv.py
Original file line number Diff line number Diff line change
Expand Up @@ -115,8 +115,11 @@ def _groupsize_quant_prefill_att(

@dataclasses.dataclass
class Int4kvTritonDecodeAttState(BaseDecodeAttState):
decode_max_kv_seq_len: int = None

def init_state(self):
pass
# Graph 捕获会改写 infer_state 的长度上限,提前保存真实长度用于配置查找。
self.decode_max_kv_seq_len = self.infer_state.max_kv_seq_len

def copy_for_decode_cuda_graph(self, new_state: "Int4kvTritonDecodeAttState"):
super().copy_for_decode_cuda_graph(new_state)
Expand Down Expand Up @@ -166,5 +169,6 @@ def ppl_int4kv_decode_att(
cache_k_scale=k_scale,
cache_v=v,
cache_v_scale=v_scale,
max_kv_seq_len=self.decode_max_kv_seq_len,
alloc_tensor_func=alloc_func,
)
4 changes: 4 additions & 0 deletions lightllm/common/basemodel/attention/triton/int8kv.py
Original file line number Diff line number Diff line change
Expand Up @@ -118,8 +118,11 @@ def _groupsize_quant_prefill_att(
class Int8kvTritonDecodeAttState(BaseDecodeAttState):
b_shared_seq_len: torch.Tensor = None
b_mark_shared_group: torch.Tensor = None
decode_max_kv_seq_len: int = None

def init_state(self):
# Graph 捕获会改写 infer_state 的长度上限,提前保存真实长度用于普通 decode 配置查找。
self.decode_max_kv_seq_len = self.infer_state.max_kv_seq_len
if enable_diverse_mode_gqa_decode_fast_kernel():
self.b_mark_shared_group = build_diverse_shared_group_markers(
b_shared_radix_node_id=self.infer_state.b_shared_radix_node_id,
Expand Down Expand Up @@ -203,5 +206,6 @@ def normal_decode_att(
cache_k_scale=k_scale,
cache_v=v,
cache_v_scale=v_scale,
max_len_in_batch=self.decode_max_kv_seq_len,
alloc_tensor_func=alloc_func,
)
64 changes: 3 additions & 61 deletions lightllm/common/basemodel/basemodel.py
Original file line number Diff line number Diff line change
Expand Up @@ -38,11 +38,8 @@
)
from lightllm.common.basemodel.mtp_manager import MtpManager
from lightllm.utils.custom_kernel_utis import pad2dim_tensor_to_new_batch
from lightllm.utils.envs_utils import (
set_model_init_status,
enable_full_att_decode_tune,
)
from lightllm.common.triton_utils.autotuner import Autotuner
from lightllm.utils.envs_utils import set_model_init_status
from lightllm.common.triton_utils.autotuner import Autotuner, AutotuneKernelType
from lightllm.utils.infer_utils import post_empty_cache
from lightllm.utils.torch_memory_saver_utils import (
TorchMemorySaverWrapper,
Expand Down Expand Up @@ -140,7 +137,6 @@ def __init__(self, kvargs):

self._init_hidden_collector()
self._autotune_warmup()
self._full_att_decode_autotune()
self._init_padded_req()
self._init_cudagraph()
self._init_prefill_cuda_graph()
Expand Down Expand Up @@ -308,60 +304,6 @@ def _init_prefill_cuda_graph(self):
else:
self.prefill_graph.warmup(self)

@final
@torch.no_grad()
@post_empty_cache
def _full_att_decode_autotune(self):
"""
Warm up / autotune FA3 full-attention decode ``num_splits`` before CUDA Graph capture.

Runs only when all of the following hold:
- CUDA Graph is enabled (``disable_cudagraph`` is False)
- this is the main model (MTP draft models are skipped)
- ``ENABLE_FULL_ATT_DECODE_TUNE`` is set to 1/ON/TRUE (default off)
- decode attention backend is ``Fa3AttBackend``

Candidate batch sizes follow the same schedule as CUDA Graph capture.
Actual benchmarking is delegated to ``fa3_decode_autotune`` in ``sgl_utils``.
"""
if self.disable_cudagraph:
return
# Only tune on the main model; MTP draft models skip this path.
if self.is_mtp_draft_model:
return

# Opt-in switch for FA3 full-attention decode num_splits tuning.
# Set ENABLE_FULL_ATT_DECODE_TUNE=1/ON/TRUE to enable; default is off.
if not enable_full_att_decode_tune():
return

# Only Fa3AttBackend decode path needs this num_splits warmup.
decode_backends = [
self.decode_att_backend,
self.decode_att_backend1,
]
if not any(
backend is not None and backend.__class__.__name__ == "Fa3AttBackend" for backend in decode_backends
):
return

from lightllm.utils.sgl_utils import fa3_decode_autotune

decode_batch_multiplier = self.mtp_manager.get_decode_batch_multiplier(self.is_mtp_draft_model)
cuda_graph_grow_step_size = self.mtp_manager.get_decode_cuda_graph_grow_step_size(self.is_mtp_draft_model)
cuda_graph_batch_sizes = CudaGraph.gen_cuda_graph_batch_sizes(
batch_step_size_before_split=cuda_graph_grow_step_size,
split_batch_size=self.args.graph_split_batch_size * decode_batch_multiplier,
batch_step_size_after_split=self.args.graph_grow_step_size * cuda_graph_grow_step_size,
max_batch_size=self.graph_max_batch_size,
tp_world_size=self.tp_world_size_,
)
cuda_graph_batch_sizes = [
batch_size for batch_size in cuda_graph_batch_sizes if batch_size % decode_batch_multiplier == 0
]
fa3_decode_autotune(self, cuda_graph_batch_sizes, batch_multiplier=decode_batch_multiplier)
return

def _init_custom(self):
pass

Expand Down Expand Up @@ -1155,7 +1097,7 @@ def autotune_layers(self):
@torch.no_grad()
@post_empty_cache
def _autotune_warmup(self):
Autotuner.start_autotune_warmup()
Autotuner.start_autotune_warmup(AutotuneKernelType.GENERAL)
torch.distributed.barrier()

warmup_lengths = [1, 4, 8, 16, 32, 64, 128, 256, 1024, 2048, 4096]
Expand Down
8 changes: 6 additions & 2 deletions lightllm/common/basemodel/cuda_graph.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@
from lightllm.utils.envs_utils import get_env_start_args
from lightllm.distributed import dist_group_manager
from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput
from lightllm.common.triton_utils.autotuner import Autotuner, AutotuneKernelType
from lightllm.utils.torch_memory_saver_utils import (
TorchMemorySaverWrapper,
MemoryTag,
Expand Down Expand Up @@ -117,7 +118,9 @@ def _capture_decode(self, decode_func, infer_state: InferStateInfo):
# 记录原始存在的变量
pure_para_set = set(vars(infer_state).keys())
torch.cuda.synchronize()
decode_func(copy.copy(infer_state))
# 在正式捕获前调优 decode attention,退出作用域后再捕获选定的配置。
with Autotuner.autotune_warmup(AutotuneKernelType.DECODE_ATTENTION):
decode_func(copy.copy(infer_state))
torch.cuda.synchronize()
for param_name in set(vars(infer_state).keys()):
if param_name not in pure_para_set:
Expand Down Expand Up @@ -149,7 +152,8 @@ def _capture_decode_overlap(
pure_para_set = set(vars(infer_state).keys())
pure_para_set1 = set(vars(infer_state1).keys())
torch.cuda.synchronize()
decode_func(copy.copy(infer_state), copy.copy(infer_state1))
with Autotuner.autotune_warmup(AutotuneKernelType.DECODE_ATTENTION):
decode_func(copy.copy(infer_state), copy.copy(infer_state1))
torch.cuda.synchronize()
for para_name in set(vars(infer_state).keys()):
if para_name not in pure_para_set:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,7 @@
quantize_fused_experts_input,
)
from lightllm.common.basemodel.triton_kernel.fused_moe.moe_silu_and_mul import silu_and_mul_fwd
from lightllm.common.triton_utils.autotuner import Autotuner
from lightllm.common.triton_utils.autotuner import Autotuner, AutotuneKernelType
from lightllm.common.basemodel.triton_kernel.redundancy_topk_ids_repair import redundancy_topk_ids_repair


Expand Down Expand Up @@ -250,7 +250,7 @@ def prefilled_group_gemm(
# A rank may receive no tokens during autotune warmup. Run one dummy token through
# silu_and_mul_fwd so the empty rank matches the first kernel call made by non-empty ranks.
# This branch does not synchronize additional calls caused by different positive chunk counts.
if Autotuner.is_autotune_warmup():
if Autotuner.is_kernel_autotune_warmup(AutotuneKernelType.GENERAL):
N = w13_weight.shape[1]
_gemm_out_a = torch.zeros((1, N), device=recv_x[0].device, dtype=hidden_dtype)
_silu_out = torch.zeros((1, N // 2), device=recv_x[0].device, dtype=hidden_dtype)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@ def gqa_token_decode_attention_flash_decoding(
infer_state,
cache_k: torch.Tensor,
cache_v: torch.Tensor,
max_len_in_batch: int,
out=None,
alloc_tensor_func=torch.empty,
sliding_window=(-1, -1),
Expand Down Expand Up @@ -41,7 +42,7 @@ def gqa_token_decode_attention_flash_decoding(
Req_to_tokens=infer_state.req_manager.req_to_token_indexs,
B_req_idx=infer_state.b_req_idx,
B_Seqlen=infer_state.b_seq_len,
max_len_in_batch=infer_state.max_kv_seq_len,
max_len_in_batch=max_len_in_batch,
mid_out=mid_o,
mid_out_logsumexp=mid_o_logexpsum,
block_seq=BLOCK_SEQ,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,8 @@
import triton
import triton.language as tl
from typing import Optional
from lightllm.common.triton_utils.autotuner import autotune, Autotuner
from lightllm.common.triton_utils.autotuner import autotune, Autotuner, AutotuneKernelType, AutotuneLevel
from lightllm.utils.envs_utils import get_decode_attn_autotune_seq_len, get_triton_autotune_level


@triton.jit
Expand Down Expand Up @@ -138,27 +139,92 @@ def get_test_configs():
return configs


def get_static_key(q, k, block_seq):
def get_static_key(q, k, block_seq, sliding_window):
key_params = {
"gqa_group_size": int(q.shape[1] // k.shape[1]),
"q_head_dim": int(q.shape[2]),
"block_seq": block_seq,
"sliding_window": tuple(sliding_window),
"out_dtype": str(q.dtype),
}
return key_params


def get_run_key(q, max_len_in_batch):
batch_size = q.shape[0]
return batch_size * 1000 * 1000 * 1000 + max_len_in_batch
# 正常执行使用调用方在 CPU 上保存的真实 KV 长度,不读取 GPU 长度张量或 Graph 的容量上限。
max_kv_len = int(max_len_in_batch)
if Autotuner.is_kernel_autotune_warmup(AutotuneKernelType.DECODE_ATTENTION) and get_triton_autotune_level() in [
AutotuneLevel.ADAPTIVE_AUTOTUNE,
AutotuneLevel.FORCE_AUTOTUNE,
]:
max_kv_len = get_decode_attn_autotune_seq_len()
# 调优和正常查找统一按 512 token 向上分桶,同一区间复用配置匹配结果。
max_kv_len = (max_kv_len + 511) // 512 * 512
return batch_size * 1000 * 1000 * 1000 + max_kv_len


def rebuild_inputs(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
Req_to_tokens: torch.Tensor,
B_req_idx: torch.Tensor,
B_Seqlen: torch.Tensor,
max_len_in_batch: int,
mid_out: torch.Tensor,
mid_out_logsumexp: torch.Tensor,
block_seq: int,
sliding_window=(-1, -1),
**kwargs,
):
# Graph 初始化时真实请求很短,Req_to_tokens 的宽度则是容量上限,都不代表期望调优的长度。
# 仅在实际搜索配置前重建一次输入,构造开销不计入 benchmark;正常执行和 Graph 捕获使用原输入。
batch_size = q.shape[0]
# 与调优时的 run key 共用该环境变量,默认 32768 token;实际计算保持精确长度,不做 512 分桶。
max_len_in_batch = get_decode_attn_autotune_seq_len()
assert k.shape[0] == v.shape[0], "K/V caches must have the same number of tokens"
num_tokens = k.shape[0]
if num_tokens == 0:
raise ValueError("GQA decode autotuning requires a non-empty KV cache")

# 新建每个请求到物理 token 的映射,不能直接扩展 B_Seqlen 后读取原映射中未初始化的条目。
# 物理 token 充足时各请求使用不同位置,不足时取模循环复用,保证所有索引均落在 K/V 缓存内。
# 复用已有 K/V 可以避免分配完整的长请求缓存,但可能提高 GPU 缓存命中率,影响调优的访存特征。
Req_to_tokens = torch.arange(batch_size * max_len_in_batch, dtype=Req_to_tokens.dtype, device=Req_to_tokens.device)
Req_to_tokens = Req_to_tokens.remainder_(num_tokens).view(batch_size, max_len_in_batch)
# 新映射只有 batch_size 行,请求索引也必须重建,避免继续使用原全局请求表中的行号。
B_req_idx = torch.arange(batch_size, dtype=B_req_idx.dtype, device=B_req_idx.device)
B_Seqlen = torch.full_like(B_Seqlen, max_len_in_batch)

# 保留 Q、滑窗语义、BLOCK_SEQ 和中间缓冲区布局;一个 program 可循环处理多个 KV 块,
# 无需按调优长度扩容 mid_out。调优结束后 stage1 使用原始输入重新覆盖有效中间块,
# stage2 仍按相同 BLOCK_SEQ 和缓冲区中的 block_num 归约。
return (
q,
k,
v,
Req_to_tokens,
B_req_idx,
B_Seqlen,
max_len_in_batch,
mid_out,
mid_out_logsumexp,
block_seq,
sliding_window,
), kwargs


@autotune(
kernel_name="_fwd_kernel_gqa_flash_decode_stage1:v3",
kernel_name="_fwd_kernel_gqa_flash_decode_stage1:v4",
kernel_type=AutotuneKernelType.DECODE_ATTENTION,
configs_gen_func=get_test_configs,
static_key_func=get_static_key,
run_key_func=get_run_key,
mutates_args=["mid_out", "mid_out_logsumexp"],
rebuild_input_func=rebuild_inputs,
# stage1 对有效中间块执行覆盖写,候选配置不会读取已有输出,正式执行也会重新覆盖真实请求的有效块。
# 不标记这两个大缓冲区,避免每个候选配置 benchmark 时反复 clone,增加显存峰值和拷贝开销。
# mutates_args=["mid_out", "mid_out_logsumexp"],
)
@torch.no_grad()
def flash_decode_stage1(
Expand Down Expand Up @@ -246,8 +312,6 @@ def flash_decode_stage1(


if __name__ == "__main__":
from lightllm.utils.envs_utils import get_triton_autotune_level

if get_triton_autotune_level() != 2:
raise Exception("you need set env LIGHTLLM_TRITON_AUTOTUNE_LEVEL=2 to start program.")

Expand All @@ -258,11 +322,11 @@ def flash_decode_stage1(
out_dtype = torch.bfloat16

batch_sizes = [1, 8, 16, 32, 64, 128]
decode_lengths = [1024, 2048, 8192, 16384]
decode_lengths = [get_decode_attn_autotune_seq_len()]

q_head_num = gqa_group_size

Autotuner.start_autotune_warmup()
Autotuner.start_autotune_warmup(AutotuneKernelType.DECODE_ATTENTION)
# autotuing kernel
for batch_size in batch_sizes:
for length in decode_lengths:
Expand Down
Loading
Loading