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
86 changes: 43 additions & 43 deletions docs/docs/developer_guide/layer_boundary.mdx

Large diffs are not rendered by default.

14 changes: 7 additions & 7 deletions python/sglang/srt/batch_overlap/two_batch_overlap.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,7 @@
tbo_split_moves,
)
from sglang.srt.layers.layer_boundary.residual import batch as residual_batch
from sglang.srt.layers.layer_boundary.residual.access import finish_layer_stack
from sglang.srt.layers.layer_boundary.residual.access import export_output
from sglang.srt.layers.layer_boundary.residual.stream import ResidualStream
from sglang.srt.layers.moe import (
get_deepep_mode,
Expand Down Expand Up @@ -954,17 +954,17 @@ def model_forward_stages(
if not enable_tbo:
return execute_operations(inputs, strategy.operations)["hidden_states"]

stream = residual_batch.current(forward_batch)
stream = residual_batch.stream_of(forward_batch)
pending = stream.pending
hidden_states, residual = stream.finish(hidden_states)
hidden_states, residual = stream.export(hidden_states)
inputs["hidden_states"] = hidden_states
parts = _model_forward_tbo_split_inputs(
**inputs,
residual=residual,
layer_input_rows=layers[0].attn_boundary.incoming_residual_rows,
)
for part in parts:
part["hidden_states"], child_stream = ResidualStream.arrive(
part["hidden_states"], child_stream = ResidualStream.from_handoff(
part["hidden_states"],
part.pop("residual"),
pending.update if pending is not None else None,
Expand All @@ -987,7 +987,7 @@ def model_forward_stages(
delta_stages=[0, strategy.tbo_delta_stages],
)
for output in outputs:
output["residual"] = residual_batch.current(output["forward_batch"])
output["residual"] = residual_batch.stream_of(output["forward_batch"])
hidden_states, forward_batch.residual_stream = _model_forward_tbo_merge_outputs(
*outputs, original_len
)
Expand Down Expand Up @@ -1110,7 +1110,7 @@ def _model_forward_tbo_merge_outputs(output_a, output_b, original_len):
assert pending_a.update is pending_b.update
update = pending_a.update
for output in (output_a, output_b):
output["hidden_states"], output["residual"] = finish_layer_stack(
output["hidden_states"], output["residual"] = export_output(
output["hidden_states"], output["residual"], output["forward_batch"]
)

Expand All @@ -1133,7 +1133,7 @@ def _handle_key(name):

hidden, residual = _handle_key("hidden_states"), _handle_key("residual")
return (
ResidualStream.arrive(hidden, residual, update)
ResidualStream.from_handoff(hidden, residual, update)
if has_stream
else (hidden, residual)
)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -35,7 +35,7 @@ def forward_mha_prepare_npu(
hidden_states: torch.Tensor,
forward_batch: "ForwardBatch",
zero_allocator: "BumpAllocator",
input_on_attention_tp_slices: bool,
input_on_attn_tp_slices: bool,
):
if m.q_lora_rank is not None:
q, latent_cache = (
Expand Down Expand Up @@ -63,7 +63,7 @@ def forward_mha_prepare_npu(

else:
q = m.q_a_layernorm(q)
if _use_ag_after_qlora and input_on_attention_tp_slices:
if _use_ag_after_qlora and input_on_attn_tp_slices:
q = scattered_to_tp_attn_full(q, forward_batch)
latent_cache = scattered_to_tp_attn_full(latent_cache, forward_batch)
q = m.q_b_proj(q)[0].view(-1, m.num_local_heads, m.qk_head_dim)
Expand Down Expand Up @@ -157,7 +157,7 @@ def forward_mla_prepare_npu(
hidden_states: torch.Tensor,
forward_batch: "ForwardBatch",
zero_allocator: "BumpAllocator",
input_on_attention_tp_slices: bool,
input_on_attn_tp_slices: bool,
):
if is_mla_preprocess_enabled():
if not hasattr(m, "mla_preprocess"):
Expand Down Expand Up @@ -190,7 +190,7 @@ def forward_mla_prepare_npu(
q_lora = None
if m.q_lora_rank is not None:
qkv_latent = get_attn_tp_context().fetch_qkv_latent()
if _use_ag_after_qlora and input_on_attention_tp_slices:
if _use_ag_after_qlora and input_on_attn_tp_slices:
q, latent_cache = qkv_latent.split(
[m.q_lora_rank, m.kv_lora_rank + m.qk_rope_head_dim],
dim=-1,
Expand Down Expand Up @@ -362,7 +362,7 @@ def forward_dsa_prepare_npu(
hidden_states: torch.Tensor,
forward_batch: "ForwardBatch",
zero_allocator: "BumpAllocator",
input_on_attention_tp_slices: bool,
input_on_attn_tp_slices: bool,
prev_topk_indices: torch.Tensor = None,
):
dynamic_scale = None
Expand Down Expand Up @@ -396,7 +396,7 @@ def forward_dsa_prepare_npu(
)
# overlap qk norm
q = m.q_a_layernorm(q)
if _use_ag_after_qlora and input_on_attention_tp_slices:
if _use_ag_after_qlora and input_on_attn_tp_slices:
q = scattered_to_tp_attn_full(q, forward_batch)
latent_cache = scattered_to_tp_attn_full(latent_cache, forward_batch)
q_lora = q.clone() # required for topk_indices
Expand Down Expand Up @@ -485,7 +485,7 @@ def forward_dsa_prepare_npu(
positions,
forward_batch,
m.layer_id,
input_on_attention_tp_slices,
input_on_attn_tp_slices,
dynamic_scale,
)
else:
Expand Down
6 changes: 3 additions & 3 deletions python/sglang/srt/layers/attention/dsa/dsa_npu_indexer.py
Original file line number Diff line number Diff line change
Expand Up @@ -60,7 +60,7 @@ def forward_npu(
positions: torch.Tensor,
forward_batch: ForwardBatch,
layer_id: int,
input_on_attention_tp_slices: bool = False,
input_on_attn_tp_slices: bool = False,
dynamic_scale: torch.Tensor = None,
) -> torch.Tensor:
if get_attn_backend().forward_metadata.seq_lens_cpu_int is None:
Expand Down Expand Up @@ -140,7 +140,7 @@ def forward_npu(

k_proj = self.wk(x)[0] # [b, s, 7168] @ [7168, 128] = [b, s, 128]
k = self.k_norm(k_proj)
if _use_ag_after_qlora and input_on_attention_tp_slices:
if _use_ag_after_qlora and input_on_attn_tp_slices:
k = scattered_to_tp_attn_full(k, forward_batch)
k_pe, k_nope = torch.split(
k,
Expand Down Expand Up @@ -284,7 +284,7 @@ def forward_npu(
torch.npu.current_stream().wait_event(q_rope_event)
if envs.SGLANG_NPU_USE_MULTI_STREAM.get():
torch.npu.current_stream().wait_event(weights_event)
if _use_ag_after_qlora and input_on_attention_tp_slices:
if _use_ag_after_qlora and input_on_attn_tp_slices:
weights = scattered_to_tp_attn_full(weights, forward_batch)
block_table = get_attn_backend().forward_metadata.block_tables
if (
Expand Down
2 changes: 1 addition & 1 deletion python/sglang/srt/layers/flashinfer_comm_fusion.py
Original file line number Diff line number Diff line change
Expand Up @@ -941,7 +941,7 @@ def can_use_flashinfer_allreduce(

# Size checks stay last: they read the token dim, which is symbolic under
# Dynamo, so statically-off configs must short-circuit before reaching them
# (same ordering rule as apply_flashinfer_allreduce_fusion).
# (same ordering rule as flashinfer_ar_fusion_applies).
token_num, hidden_dim = input_.shape

# MNNVL hard-fails instead of falling back when the width is not float4-aligned
Expand Down
86 changes: 43 additions & 43 deletions python/sglang/srt/layers/layer_boundary/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,27 +18,27 @@
get_attn_tp_context,
)
from sglang.srt.layers.layer_boundary.boundary import (
make_boundary,
make_output_boundary,
bind_entry,
bind_exit,
tbo_split_moves,
)
from sglang.srt.layers.layer_boundary.construction import (
BatchVariant,
StageEdges,
VariantEdges,
)
from sglang.srt.layers.layer_boundary.contracts import (
EdgeDecl,
FusedMlpInput,
HandoffRows,
EdgeContract,
EntryPath,
ExitRows,
FfnInputFusion,
InputContract,
OutputContract,
ProducerReduction,
StageDecl,
StageEntry,
StageInput,
StageContract,
StageKind,
StageOutput,
StageSteps,
StagePath,
)
from sglang.srt.layers.layer_boundary.exit import FfnCompletion, FfnExit, MixerExit
from sglang.srt.layers.layer_boundary.exit import ExitDecision, FfnExit, MixerExit
from sglang.srt.layers.layer_boundary.factories import (
declare_attn,
declare_ffn,
Expand All @@ -50,27 +50,27 @@
Layout,
SumGroup,
TokenAxis,
enable_moe_dense_fully_dp,
moe_cp_gathers_sparse_moe_input,
sparse_moe_gathers_over_moe_cp,
batch_gathers_over_moe_cp,
is_dense_ffn_fully_dp,
moe_gathers_over_moe_cp,
token_axis_sizes,
)
from sglang.srt.layers.layer_boundary.ops import (
move_rows,
tp_reduce_scatter,
)
from sglang.srt.layers.layer_boundary.output import (
HandoffOutput,
DeferredFinalize,
UnreducedOutput,
reduce_output,
complete_owed,
)
from sglang.srt.layers.layer_boundary.residual import LayerResidual
from sglang.srt.layers.layer_boundary.residual import LayerResidualOps
from sglang.srt.layers.layer_boundary.residual.add_norm import (
ADD,
FUSE_ALLREDUCE_MAX_BATCH_SIZE,
NORM_QUANT_READ,
NORM_READ,
PLAIN_RESIDUAL,
NORM_QUANT_READOUT,
NORM_READOUT,
PLAIN_ADD,
PLAIN_RESIDUAL_OPS,
)
from sglang.srt.layers.layer_boundary.residual.mhc import (
MHCState,
Expand All @@ -82,40 +82,40 @@
"make_attn_stage",
"make_ffn_stage",
"make_stages",
"ADD",
"PLAIN_ADD",
"AttentionInputs",
"StageSteps",
"EdgeDecl",
"HandoffRows",
"StagePath",
"EdgeContract",
"ExitRows",
"ProducerReduction",
"FUSE_ALLREDUCE_MAX_BATCH_SIZE",
"FfnCompletion",
"ExitDecision",
"FfnExit",
"FusedMlpInput",
"HandoffOutput",
"LayerResidual",
"FfnInputFusion",
"DeferredFinalize",
"LayerResidualOps",
"Layout",
"MHCState",
"MixerExit",
"NORM_QUANT_READ",
"NORM_READ",
"PLAIN_RESIDUAL",
"StageDecl",
"StageEntry",
"StageInput",
"NORM_QUANT_READOUT",
"NORM_READOUT",
"PLAIN_RESIDUAL_OPS",
"StageContract",
"EntryPath",
"InputContract",
"StageKind",
"StageOutput",
"OutputContract",
"SumGroup",
"TokenAxis",
"UnreducedOutput",
"enable_moe_dense_fully_dp",
"is_dense_ffn_fully_dp",
"get_attn_tp_context",
"make_boundary",
"make_output_boundary",
"moe_cp_gathers_sparse_moe_input",
"bind_entry",
"bind_exit",
"batch_gathers_over_moe_cp",
"move_rows",
"reduce_output",
"sparse_moe_gathers_over_moe_cp",
"complete_owed",
"moe_gathers_over_moe_cp",
"tbo_split_moves",
"token_axis_sizes",
"tp_reduce_scatter",
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -26,7 +26,7 @@
is_dp_attention_enabled,
)
from sglang.srt.layers.layer_boundary.layout import (
enable_moe_dense_fully_dp,
is_dense_ffn_fully_dp,
)
from sglang.srt.layers.moe import get_moe_a2a_backend
from sglang.srt.model_executor.cuda_graph_config import (
Expand Down Expand Up @@ -107,7 +107,7 @@ def init_context(self, q_lora_rank, is_dsa, is_mhc=False):
and get_parallel().tp_size > 1
and not is_dp_attention_enabled()
and get_moe_a2a_backend().is_none()
and not enable_moe_dense_fully_dp()
and not is_dense_ffn_fully_dp()
and not check_cuda_graph_backend(Phase.PREFILL, Backend.TC_PIECEWISE)
and get_spec().speculative_algorithm != "EAGLE3"
)
Expand Down Expand Up @@ -168,14 +168,14 @@ def get_attn_tp_context():
return ATTN_TP_CONTEXT


def _redistribute_from_attn_tp_shards(tensor: torch.Tensor) -> torch.Tensor:
def attn_tp_gather(tensor: torch.Tensor) -> torch.Tensor:
gathered = get_local_dp_buffer(
get_parallel().attn_tp_group, hidden_size=tensor.shape[-1]
)
attn_tp_all_gather_into_tensor(gathered, tensor)
return gathered


def _redistribute_to_attn_tp_shards(tensor: torch.Tensor) -> torch.Tensor:
def attn_tp_slice(tensor: torch.Tensor) -> torch.Tensor:
parallel = get_parallel()
return tensor.tensor_split(parallel.attn_tp_size)[parallel.attn_tp_rank]
14 changes: 7 additions & 7 deletions python/sglang/srt/layers/layer_boundary/adapters/branch.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,8 +37,8 @@ def branch_input(
"""The FFN input and residual that ``source``'s boundary read for its
own FFN, for this layer's FFN, which branches from the same input: moved
to the rows this FFN needs and its residual's rows."""
rows, residual_rows, _ = source._branch_rows(forward_batch)
to, residual_to, _ = plan._branch_rows(forward_batch)
rows, residual_rows, _ = source.branch_rows(forward_batch)
to, residual_to, _ = plan.branch_rows(forward_batch)
if stream.pending is not None:
raise RuntimeError("a branch must start from a prepared stage input")
residual = stream.residual
Expand All @@ -56,7 +56,7 @@ def branch_output(
"""This layer's complete FFN output as a branch's contribution, which
adds to the layer's output without writing the residual: moved to the
rows the layer hands on."""
rows, _, to = plan._branch_rows(forward_batch)
rows, _, to = plan.branch_rows(forward_batch)
return move_rows(hidden_states, rows, to, forward_batch)


Expand All @@ -71,15 +71,15 @@ def merge_branch(
"""A contribution from ``branch_output`` summed with what ``source``'s
layer hands on, ``hidden_states`` and ``residual``, moved to the rows
this layer hands on."""
_, _, rows = source._branch_rows(forward_batch)
_, _, to = plan._branch_rows(forward_batch)
_, _, rows = source.branch_rows(forward_batch)
_, _, to = plan.branch_rows(forward_batch)
stream.check(hidden_states)
if stream.pending is None:
raise RuntimeError("branch merge requires a pending producer contribution")
update = stream.pending.update
hidden_states, residual = stream.finish(hidden_states)
hidden_states, residual = stream.export(hidden_states)
hidden_states = move_rows(hidden_states, rows, to, forward_batch)
residual = move_rows(residual, rows, to, forward_batch)
output = contribution + hidden_states
stream.write(residual)
return stream.leave(output, update), stream
return stream.record(output, update), stream
Original file line number Diff line number Diff line change
Expand Up @@ -49,7 +49,7 @@ def maybe_prefetch_next_full_attention_kv(
prefetch_kv_buffer(next_full_attention_layer_id)


def dsa_cp_gather_hidden_states(hidden_states: torch.Tensor):
def attn_cp_gather(hidden_states: torch.Tensor):
attn_dp_size = get_parallel().attn_dp_size
attn_tp_size = get_parallel().attn_tp_size
assert attn_dp_size == 1 and attn_tp_size == 1
Expand All @@ -61,7 +61,7 @@ def dsa_cp_gather_hidden_states(hidden_states: torch.Tensor):
return hidden_states


def dsa_cp_reduce_scatter_hidden_states(hidden_states: torch.Tensor):
def attn_cp_reduce_scatter(hidden_states: torch.Tensor):
attn_dp_size = get_parallel().attn_dp_size
attn_tp_size = get_parallel().attn_tp_size
assert attn_dp_size == 1 and attn_tp_size == 1
Expand Down
Loading
Loading