Skip to content
Draft
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
46 changes: 46 additions & 0 deletions python/sglang/srt/distributed/parallel_state.py
Original file line number Diff line number Diff line change
Expand Up @@ -872,6 +872,52 @@ def fused_allreduce_rmsnorm(
)
return fused_outputs

def fused_allreduce_mhc_post(
self,
input_: torch.Tensor,
residual: torch.Tensor,
post_layer_mix: torch.Tensor,
comb_res_mix: torch.Tensor,
) -> Optional[torch.Tensor]:
"""Fused all-reduce + mHC post via the aiter custom all-reduce: writes
``residual`` combined through ``comb_res_mix`` plus the reduced
``input_`` spread by ``post_layer_mix`` into new streams. ROCm/HIP only;
None when the communicator cannot take the input."""
ca_comm = self.ca_comm
if (
ca_comm is None
or getattr(ca_comm, "disabled", True)
or not hasattr(ca_comm, "_pool")
or not ca_comm.should_custom_ar(input_)
):
return None
try:
from aiter.ops.custom_all_reduce import fused_allreduce_mhc_post_only
except ImportError:
return None

registered = False
if ca_comm._IS_CAPTURING:
if not torch.cuda.is_current_stream_capturing():
return torch.zeros_like(residual)
registered = ca_comm.enable_register_for_capturing
reg = 0 if registered else ca_comm._pool["input"].data_ptr
reg_bytes = 0 if registered else ca_comm._pool["input"].max_size
out = torch.empty_like(residual)
fused_allreduce_mhc_post_only(
ca_comm._ptr,
input_,
out,
residual,
post_layer_mix,
comb_res_mix,
True,
False,
reg,
reg_bytes,
)
return out

def fused_allreduce_rmsnorm_quant_per_group(
self,
input_: torch.Tensor,
Expand Down
3 changes: 3 additions & 0 deletions python/sglang/srt/environ.py
Original file line number Diff line number Diff line change
Expand Up @@ -917,6 +917,9 @@ class Envs:
# Enable dual-stream MoE (shared experts vs routed experts) on the
# ROCm/AITER path. Requires GPU_MAX_HW_QUEUES>=5 to avoid HW-queue serialization.
SGLANG_ROCM_USE_MULTI_STREAM = EnvBool(False)
# Run an mHC layer's FFN all-reduce and hc_post as one aiter kernel on
# decode batches.
SGLANG_ROCM_FUSED_AR_MHC_POST = EnvBool(False)
# Fold the KDA [f_a|b] tail into the wide [q,k,v,g] projection so the whole
# in-proj is one GEMM. Decode is bandwidth bound there, so the 144 extra
# output columns ride along nearly free.
Expand Down
54 changes: 54 additions & 0 deletions python/sglang/srt/layers/layer_boundary/exit.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,7 @@
from sglang.srt.layers.layer_boundary.adapters.attention import get_attn_tp_context
from sglang.srt.layers.layer_boundary.layout import (
SumGroup,
_batch_size,
_ffn_has_tokens,
_sum_group,
)
Expand All @@ -39,6 +40,7 @@
all_reduce_to_dp_local,
dp_reduce_scatter,
dp_reduce_scatterv,
keep_output,
to_dp_local,
)
from sglang.srt.layers.layer_boundary.output import (
Expand All @@ -48,6 +50,7 @@
aiter_ar_fusion_applies,
flashinfer_ar_fusion_applies,
)
from sglang.srt.layers.layer_boundary.residual.mhc import _FfnUpdate as _MhcFfnUpdate
from sglang.srt.layers.layer_boundary.residual.stream import ResidualStream
from sglang.srt.layers.moe import (
can_merge_post_experts_all_reduce,
Expand All @@ -65,6 +68,11 @@
get_lora,
get_parallel,
)
from sglang.srt.utils import get_bool_env_var, is_hip

_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and is_hip()
# Decode batches, where the fused all-reduce + mHC post beats the two launches.
_FUSED_AR_MHC_POST_MAX_TOKENS = 16


def _select_dp_reduce_scatter(
Expand Down Expand Up @@ -193,6 +201,18 @@ def _decide(self, forward_batch: ForwardBatch, steps) -> ExitDecision:
self.plan.fusions is not None
and self.plan.fusions.can_defer_finalize(self.plan, forward_batch)
)
if (
not defer_moe_finalize
and dp_step is None
and not mlp_reduce_scatter
and self._fuses_all_reduce_into_mhc_post(forward_batch, steps)
):
return ExitDecision(
defer_moe_finalize=False,
fuse_mlp_allreduce=True,
mlp_reduce_scatter=False,
complete=partial(self._complete_with_mhc_post, steps=steps),
)
if not steps.output.may_defer_to_next and not (
self.plan.terminal and defer_moe_finalize
):
Expand Down Expand Up @@ -242,6 +262,40 @@ def _decide(self, forward_batch: ForwardBatch, steps) -> ExitDecision:
complete=complete,
)

def _fuses_all_reduce_into_mhc_post(
self, forward_batch: ForwardBatch, steps
) -> bool:
"""Whether this layer's FFN leaves its TP sum to aiter's all-reduce +
mHC post: an MHC write-back of a plain TP sum that stays on this rank's
rows, on a decode-sized batch."""
if not (_use_aiter and envs.SGLANG_ROCM_FUSED_AR_MHC_POST.get()):
return False
return (
isinstance(steps.output.update, _MhcFfnUpdate)
and steps.output.transform is None
and steps.output_move is keep_output
and not steps.returns_over_dp
and get_parallel().tp_size > 1
and not is_dp_attention_enabled()
and not is_enable_moe_cp_allgather()
and get_moe_a2a_backend().is_none()
and not get_attn_tp_context().input_scattered
and self._sum_deferral_allowed(steps)
and post_experts_sum_is_one_all_reduce()
and self.ffn_reduction_group(steps) is get_parallel().tp_group
and 0 < _batch_size(forward_batch) <= _FUSED_AR_MHC_POST_MAX_TOKENS
)

def _complete_with_mhc_post(
self, hidden_states: torch.Tensor, residual: torch.Tensor, *, steps
) -> Tuple[torch.Tensor, None]:
group = self.ffn_reduction_group(steps)
update = steps.output.update
written = update.reduce_and_update(group, hidden_states, residual)
if written is None:
written = update.update(group.all_reduce(hidden_states), residual)
return written, None

def _complete_now(
self,
hidden_states: torch.Tensor,
Expand Down
26 changes: 25 additions & 1 deletion python/sglang/srt/layers/layer_boundary/residual/mhc.py
Original file line number Diff line number Diff line change
Expand Up @@ -95,6 +95,20 @@ def update_and_read_ffn_input(
def apply_post(self, hidden_states, residual):
return self.hc_post(hidden_states, residual, self.h_res, self.h_post)

def reduce_and_apply_post(self, group, hidden_states, residual):
"""``apply_post`` of a partial sum, all-reduced over ``group`` in the
same kernel; None when the fused kernel declines."""
num_tokens, hidden_size = hidden_states.shape
out = group.fused_allreduce_mhc_post(
hidden_states,
residual.view(num_tokens, self.hc_mult, hidden_size),
self.h_post.view(num_tokens, self.hc_mult, 1),
self.h_res.view(num_tokens, self.hc_mult, self.hc_mult),
)
if out is None:
return None
return out.view(num_tokens, -1)

def clear_coefficients(self):
self.h_res = None
self.h_post = None
Expand Down Expand Up @@ -200,7 +214,17 @@ def __init__(self, state: MHCState):
self.state = state

def update(self, hidden_states, residual):
hidden_states = self.state.apply_post(hidden_states, residual)
return self._finish(self.state.apply_post(hidden_states, residual))

def reduce_and_update(self, group, hidden_states, residual):
"""``update`` of a partial sum, all-reduced over ``group`` in the same
kernel as hc_post; None when the fused kernel declines."""
hidden_states = self.state.reduce_and_apply_post(group, hidden_states, residual)
if hidden_states is None:
return None
return self._finish(hidden_states)

def _finish(self, hidden_states):
self.state.clear_coefficients()
if self.state.is_last_layer:
hidden_states = hc_contract(hidden_states, self.state.hc_mult)
Expand Down
Loading