Skip to content
Open
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
Original file line number Diff line number Diff line change
Expand Up @@ -3,8 +3,9 @@
import torch.distributed as dist
from ..transformer_layer_infer import TransformerLayerInfer
from ...infer_struct import InferStateInfo
from lightllm.distributed import all_reduce
from lightllm.distributed import all_reduce, all_reduce_fused_add_rmsnorm
from typing import Tuple
from lightllm.utils.envs_utils import get_env_start_args
from lightllm.utils.tensor_utils import tensor_to_no_ref_tensor


Expand All @@ -21,8 +22,48 @@ def __init__(self, layer_num, network_config):
self.tp_o_head_num_ = -1
self.head_dim_ = -1
self.embed_dim_ = -1
# Subclasses that use a plain RMSNorm after the post-attention residual add
# (gamma tensor at ``layer_weight.ffn_norm_weight_.weight``) can set this True
# to fuse all_reduce + residual add + RMSNorm via FlashInfer. See
# _should_fuse_ar_add_norm / _reduce_add_ffn_norm.
self._enable_fused_ar_add_norm = False
return

def _should_fuse_ar_add_norm(self, infer_state: InferStateInfo) -> bool:
# Only the plain-TP all-reduce path is fusible: skip under tpsp mix mode
# (reduce-scatter) and when the FlashInfer all-reduce backend is absent.
if not self._enable_fused_ar_add_norm or self.tp_world_size_ <= 1:
return False
args = get_env_start_args()
if args.enable_tpsp_mix_mode or args.disable_fused_allreduce_norm:
return False
return getattr(infer_state.dist_group, "flashinfer_reduce", None) is not None

def _reduce_add_ffn_norm(self, o, input_embdings, infer_state: InferStateInfo, layer_weight) -> torch.Tensor:
# o is the pre-reduce partial o_proj output (reduce was deferred in _get_o).
# Fuse all_reduce + (input_embdings += o) + ffn_norm, else do them separately.
# ponytail: only the post-attention reduce+add+norm is fused. The post-ffn
# reduce+add is followed by the *next* layer's att_norm, which lives in a
# different token_forward call, so it is left unfused. Fusing it too would
# need threading a separate residual across the layer loop (vLLM-style) --
# a much larger, riskier refactor for the second half of the win.
o = o.view(-1, self.embed_dim_)
flashinfer_reduce = getattr(infer_state.dist_group, "flashinfer_reduce", None)
if flashinfer_reduce is None or not flashinfer_reduce.should_use(o):
all_reduce(o, group=infer_state.dist_group)
input_embdings.add_(o)
return self._ffn_norm(input_embdings, infer_state, layer_weight)

norm_out = self.alloc_tensor(o.shape, o.dtype)
fused = all_reduce_fused_add_rmsnorm(
o, input_embdings, layer_weight.ffn_norm_weight_.weight, self.eps_, norm_out, group=infer_state.dist_group
)
if fused:
return norm_out
all_reduce(o, group=infer_state.dist_group)
input_embdings.add_(o)
return self._ffn_norm(input_embdings, infer_state, layer_weight)

def _att_norm(self, input, infer_state: InferStateInfo, layer_weight) -> torch.Tensor:
raise Exception("need to impl")

Expand All @@ -47,52 +88,70 @@ def _context_attention_kernel(self, q, kv, infer_state: InferStateInfo, layer_we
def _token_attention_kernel(self, q, infer_state: InferStateInfo, layer_weight, out=None) -> torch.Tensor:
raise Exception("need to impl")

def _get_o(self, input, infer_state: InferStateInfo, layer_weight) -> torch.Tensor:
def _get_o(self, input, infer_state: InferStateInfo, layer_weight, defer_reduction=False) -> torch.Tensor:
raise Exception("need to impl")

def _ffn(self, input, infer_state: InferStateInfo, layer_weight) -> torch.Tensor:
raise Exception("need to impl")

def context_attention_forward(self, input_embdings, infer_state: InferStateInfo, layer_weight):
def context_attention_forward(
self, input_embdings, infer_state: InferStateInfo, layer_weight, defer_reduction=False
):
q, cache_kv = self._get_qkv(input_embdings, infer_state, layer_weight)
self._post_cache_kv(cache_kv, infer_state, layer_weight)
o = self._context_attention_wrapper_run(
q=q, cache_kv=cache_kv, infer_state=infer_state, layer_weight=layer_weight
)
q = None
o = self._get_o(o, infer_state, layer_weight)
if defer_reduction:
o = self._get_o(o, infer_state, layer_weight, defer_reduction=True)
else:
o = self._get_o(o, infer_state, layer_weight)

return o

def context_forward(self, input_embdings, infer_state: InferStateInfo, layer_weight):
input1 = self._att_norm(input_embdings, infer_state, layer_weight)
o = self.context_attention_forward(input1, infer_state, layer_weight)
input_embdings.add_(o.view(-1, self.embed_dim_))
use_fused_reduce = self._should_fuse_ar_add_norm(infer_state)
if use_fused_reduce:
o = self.context_attention_forward(input1, infer_state, layer_weight, defer_reduction=True)
input1 = self._reduce_add_ffn_norm(o, input_embdings, infer_state, layer_weight)
else:
o = self.context_attention_forward(input1, infer_state, layer_weight)
input_embdings.add_(o.view(-1, self.embed_dim_))
input1 = self._ffn_norm(input_embdings, infer_state, layer_weight)
o = None

input1 = self._ffn_norm(input_embdings, infer_state, layer_weight)
ffn_out = self._ffn(input1, infer_state, layer_weight)
input1 = None

input_embdings.add_(ffn_out.view(-1, self.embed_dim_))
return input_embdings

def token_attention_forward(self, input_embdings, infer_state: InferStateInfo, layer_weight):
def token_attention_forward(self, input_embdings, infer_state: InferStateInfo, layer_weight, defer_reduction=False):
q, cache_kv = self._get_qkv(input_embdings, infer_state, layer_weight)
self._post_cache_kv(cache_kv, infer_state, layer_weight)
o = self._token_attention_kernel(q, infer_state, layer_weight)
q = None
o = self._get_o(o, infer_state, layer_weight)
if defer_reduction:
o = self._get_o(o, infer_state, layer_weight, defer_reduction=True)
else:
o = self._get_o(o, infer_state, layer_weight)

return o

def token_forward(self, input_embdings, infer_state: InferStateInfo, layer_weight):
input1 = self._att_norm(input_embdings, infer_state, layer_weight)
o = self.token_attention_forward(input1, infer_state, layer_weight)
input_embdings.add_(o.view(-1, self.embed_dim_))
use_fused_reduce = self._should_fuse_ar_add_norm(infer_state)
if use_fused_reduce:
o = self.token_attention_forward(input1, infer_state, layer_weight, defer_reduction=True)
input1 = self._reduce_add_ffn_norm(o, input_embdings, infer_state, layer_weight)
else:
o = self.token_attention_forward(input1, infer_state, layer_weight)
input_embdings.add_(o.view(-1, self.embed_dim_))
input1 = self._ffn_norm(input_embdings, infer_state, layer_weight)
o = None

input1 = self._ffn_norm(input_embdings, infer_state, layer_weight)
ffn_out = self._ffn(input1, infer_state, layer_weight)

input_embdings.add_(ffn_out.view(-1, self.embed_dim_))
Expand Down
35 changes: 35 additions & 0 deletions lightllm/distributed/communication_op.py
Original file line number Diff line number Diff line change
Expand Up @@ -100,6 +100,22 @@ def all_reduce(self, input_: torch.Tensor) -> None:
return
return dist.all_reduce(input_, group=self.device_group)

def all_reduce_fused_add_rmsnorm(
self,
input_: torch.Tensor,
residual: torch.Tensor,
weight: torch.Tensor,
eps: float,
norm_out: torch.Tensor,
) -> bool:
# Fuse all_reduce + residual add + RMSNorm via FlashInfer when the message is
# small enough for its oneshot path. Returns True if fused (residual & norm_out
# filled), False if the caller must fall back to plain all_reduce + add + norm.
if self.flashinfer_reduce is not None and self.flashinfer_reduce.should_use(input_):
self.flashinfer_reduce.all_reduce_fused_add_rmsnorm(input_, residual, weight, eps, norm_out)
return True
return False

def all_gather_into_tensor(self, output_: torch.Tensor, input_: torch.Tensor, async_op: bool = False) -> None:
return dist.all_gather_into_tensor(output_, input_, group=self.device_group, async_op=async_op)

Expand Down Expand Up @@ -235,6 +251,25 @@ def all_reduce(
return dist.all_reduce(input_, op, group, async_op)


def all_reduce_fused_add_rmsnorm(
input_: torch.Tensor,
residual: torch.Tensor,
weight: torch.Tensor,
eps: float,
norm_out: torch.Tensor,
group: Optional[Union[ProcessGroup, CustomProcessGroup]] = None,
) -> bool:
"""Try the FlashInfer all_reduce + residual add + RMSNorm fusion.

Returns True on success (``residual`` and ``norm_out`` are filled). Returns
False when unavailable (no custom group, message too large, dtype/world-size
unsupported); the caller then does plain all_reduce + add + norm itself.
"""
if isinstance(group, CustomProcessGroup):
return group.all_reduce_fused_add_rmsnorm(input_, residual, weight, eps, norm_out)
return False


def all_gather_into_tensor(
output_: torch.Tensor,
input_: torch.Tensor,
Expand Down
25 changes: 25 additions & 0 deletions lightllm/distributed/flashinfer_all_reduce.py
Original file line number Diff line number Diff line change
Expand Up @@ -134,3 +134,28 @@ def all_reduce(self, inp: torch.Tensor) -> torch.Tensor:
pattern=flashinfer_comm.AllReduceFusionPattern.kAllReduce,
# launch_with_pdl=True, # TODO: learn pdl and ensure no other side effects.
)

def all_reduce_fused_add_rmsnorm(
self,
inp: torch.Tensor,
residual: torch.Tensor,
weight: torch.Tensor,
eps: float,
norm_out: torch.Tensor,
) -> None:
"""Fused: residual += all_reduce(inp); norm_out = rmsnorm(residual, weight).

``residual`` is updated in place (residual_in aliases residual_out). Same
size / dtype constraints as ``all_reduce`` -- gate with ``should_use``.
"""
flashinfer_comm.allreduce_fusion(
input=inp,
workspace=self._workspace,
pattern=flashinfer_comm.AllReduceFusionPattern.kARResidualRMSNorm,
residual_in=residual,
residual_out=residual,
norm_out=norm_out,
rms_gamma=weight,
rms_eps=eps,
)
return
Original file line number Diff line number Diff line change
Expand Up @@ -200,14 +200,20 @@ def _get_qkv(
return q, cache_kv

def _get_o(
self, input: torch.Tensor, infer_state: Deepseek2InferStateInfo, layer_weight: Deepseek2TransformerLayerWeight
self,
input: torch.Tensor,
infer_state: Deepseek2InferStateInfo,
layer_weight: Deepseek2TransformerLayerWeight,
defer_reduction=False,
) -> torch.Tensor:
if infer_state.need_dp_prefill_balance:
input = infer_state._all_to_all_balance_get(data=input)

if input.shape[2] == self.kv_lora_rank:
input = layer_weight.v_b_proj_.bmm(input.transpose(0, 1)).transpose(0, 1)
o_tensor = layer_weight.o_weight_.mm(input.reshape(-1, self.tp_q_head_num_ * self.v_head_dim))
if defer_reduction:
return o_tensor
o_tensor = self._tpsp_reduce(input=o_tensor, infer_state=infer_state)
return o_tensor

Expand Down
13 changes: 12 additions & 1 deletion lightllm/models/llama/layer_infer/transformer_layer_infer.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,6 +25,9 @@ def __init__(self, layer_num, network_config):
self.tp_o_head_num_ = self.tp_q_head_num_
self.head_dim_ = network_config["hidden_size"] // network_config["num_attention_heads"]
self.embed_dim_ = network_config["hidden_size"]
# Llama uses a plain RMSNorm after the post-attention residual add, so the
# all_reduce + add + ffn_norm sequence can be fused (see the template).
self._enable_fused_ar_add_norm = True
self._bind_func()
return

Expand Down Expand Up @@ -97,14 +100,22 @@ def _get_qkv(
return q, cache_kv

def _get_o(
self, input, infer_state: LlamaInferStateInfo, layer_weight: LlamaTransformerLayerWeight
self,
input,
infer_state: LlamaInferStateInfo,
layer_weight: LlamaTransformerLayerWeight,
defer_reduction=False,
) -> torch.Tensor:
if infer_state.need_dp_prefill_balance:
input = infer_state._all_to_all_balance_get(data=input)

input = input.view(-1, self.tp_o_head_num_ * self.head_dim_)
o_tensor = layer_weight.o_proj.mm(input)

# When fusing, defer the all-reduce to the fused reduce+add+ffn_norm op in
# the template; return the pre-reduce partial here.
if defer_reduction:
return o_tensor
o_tensor = self._tpsp_reduce(input=o_tensor, infer_state=infer_state)
return o_tensor

Expand Down
17 changes: 13 additions & 4 deletions lightllm/models/qwen3next/layer_infer/transformer_layer_infer.py
Original file line number Diff line number Diff line change
Expand Up @@ -167,6 +167,7 @@ def _get_o(
input,
infer_state: Qwen3NextInferStateInfo,
layer_weight: Qwen3NextTransformerLayerWeight,
defer_reduction=False,
) -> torch.Tensor:
"""Output projection with gating (in-place multiply to save one allocation)."""
if infer_state.need_dp_prefill_balance:
Expand All @@ -175,6 +176,8 @@ def _get_o(
sigmoid_mul_(input, infer_state.gate_logics_value)
infer_state.gate_logics_value = None
o_tensor = layer_weight.o_proj.mm(input)
if defer_reduction:
return o_tensor
o_tensor = self._tpsp_reduce(input=o_tensor, infer_state=infer_state)
return o_tensor

Expand Down Expand Up @@ -240,10 +243,13 @@ def context_attention_forward(
input_embdings,
infer_state: Qwen3NextInferStateInfo,
layer_weight: Qwen3NextTransformerLayerWeight,
defer_reduction=False,
):
# full attention layer
if not self.is_linear_attention_layer:
return super().context_attention_forward(input_embdings, infer_state, layer_weight)
return super().context_attention_forward(
input_embdings, infer_state, layer_weight, defer_reduction=defer_reduction
)

assert isinstance(infer_state.mem_manager, Qwen3NextMemManager)
mixed_qkvzba = self._linear_in_proj(input_embdings, layer_weight)
Expand All @@ -268,7 +274,7 @@ def context_attention_forward(

gdn_out = self._linear_post(core_attn_out, z, layer_weight)

if self.tp_world_size_ > 1:
if self.tp_world_size_ > 1 and not defer_reduction:
all_reduce(gdn_out, op=dist.ReduceOp.SUM, group=infer_state.dist_group, async_op=False)
return gdn_out

Expand All @@ -277,9 +283,12 @@ def token_attention_forward(
input_embdings,
infer_state: Qwen3NextInferStateInfo,
layer_weight: Qwen3NextTransformerLayerWeight,
defer_reduction=False,
):
if not self.is_linear_attention_layer:
return super().token_attention_forward(input_embdings, infer_state, layer_weight)
return super().token_attention_forward(
input_embdings, infer_state, layer_weight, defer_reduction=defer_reduction
)

assert isinstance(infer_state.mem_manager, Qwen3NextMemManager)
mixed_qkvzba = self._linear_in_proj(input_embdings, layer_weight)
Expand All @@ -299,7 +308,7 @@ def token_attention_forward(
)
gdn_out = self._linear_post(core_attn_out, z, layer_weight)

if self.tp_world_size_ > 1:
if self.tp_world_size_ > 1 and not defer_reduction:
all_reduce(gdn_out, op=dist.ReduceOp.SUM, group=infer_state.dist_group, async_op=False)
return gdn_out

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@
class StablelmTransformerLayerInfer(LlamaTransformerLayerInfer):
def __init__(self, layer_num, network_config):
super().__init__(layer_num, network_config)
self._enable_fused_ar_add_norm = False
self.partial_rotary_factor = self.network_config_.get("partial_rotary_factor", 1)
return

Expand Down Expand Up @@ -38,13 +39,19 @@ def _get_qkv(
return q, cache_kv

def _get_o(
self, input, infer_state: LlamaInferStateInfo, layer_weight: StablelmTransformerLayerWeight
self,
input,
infer_state: LlamaInferStateInfo,
layer_weight: StablelmTransformerLayerWeight,
defer_reduction=False,
) -> torch.Tensor:
if infer_state.need_dp_prefill_balance:
input = infer_state._all_to_all_balance_get(data=input)
o_tensor = layer_weight.o_proj.mm(
input.view(-1, self.tp_o_head_num_ * self.head_dim_),
)
if defer_reduction:
return o_tensor
o_tensor = self._tpsp_reduce(input=o_tensor, infer_state=infer_state)
return o_tensor

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@
class Starcoder2TransformerLayerInfer(LlamaTransformerLayerInfer):
def __init__(self, layer_num, network_config):
super().__init__(layer_num, network_config)
self._enable_fused_ar_add_norm = False

def _att_norm(
self, input, infer_state: LlamaInferStateInfo, layer_weight: Starcoder2TransformerLayerWeight
Expand Down
7 changes: 7 additions & 0 deletions lightllm/server/api_cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -357,6 +357,13 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser:
action="store_true",
help="Disable the default FlashInfer all-reduce fast path and fall back to SymmMem / NCCL.",
)
parser.add_argument(
"--disable_fused_allreduce_norm",
action="store_true",
help="Disable the FlashInfer fused all-reduce + residual add + RMSNorm op "
"(post-attention), using separate all-reduce + add + norm instead. "
"Requires the FlashInfer all-reduce path; no effect under tpsp mix mode.",
)
parser.add_argument(
"--enable_tpsp_mix_mode",
action="store_true",
Expand Down
1 change: 1 addition & 0 deletions lightllm/server/core/objs/start_args_type.py
Original file line number Diff line number Diff line change
Expand Up @@ -105,6 +105,7 @@ class StartArgs:
visual_use_proxy_mode: bool = field(default=False)
disable_symm_mem_allreduce: bool = field(default=False)
disable_flashinfer_allreduce: bool = field(default=False)
disable_fused_allreduce_norm: bool = field(default=False)
enable_tpsp_mix_mode: bool = field(default=False)
enable_dp_prefill_balance: bool = field(default=False)
enable_decode_microbatch_overlap: bool = field(default=False)
Expand Down
Loading
Loading