Support Nemotron DP attention and MTP (#24955)
Co-authored-by: Jiajun Li <48857426+guapisolo@users.noreply.github.com> Co-authored-by: Zhichenzzz <northwesterniemsteaching@gmail.com>
This commit is contained in:
co-authored by
Jiajun Li
Zhichenzzz
parent
bb33594c1a
commit
6e0fa5afe1
@@ -762,6 +762,8 @@ class TboForwardBatchPreparer:
|
|||||||
tbo_parent_token_range=(start_token_index, end_token_index),
|
tbo_parent_token_range=(start_token_index, end_token_index),
|
||||||
tbo_children=None,
|
tbo_children=None,
|
||||||
original_global_num_tokens_cpu=None,
|
original_global_num_tokens_cpu=None,
|
||||||
|
_original_batch_size=None,
|
||||||
|
_original_forward_mode=None,
|
||||||
global_num_tokens_gpu=None,
|
global_num_tokens_gpu=None,
|
||||||
global_num_tokens_cpu=None,
|
global_num_tokens_cpu=None,
|
||||||
global_dp_buffer_len=global_dp_buffer_len,
|
global_dp_buffer_len=global_dp_buffer_len,
|
||||||
|
|||||||
@@ -24,7 +24,10 @@ from sglang.srt.layers.attention.utils import (
|
|||||||
assert_buffer_fits,
|
assert_buffer_fits,
|
||||||
create_flashinfer_kv_indices_triton,
|
create_flashinfer_kv_indices_triton,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.dp_attention import get_attention_tp_size
|
from sglang.srt.layers.dp_attention import (
|
||||||
|
get_attention_cp_size,
|
||||||
|
get_attention_tp_size,
|
||||||
|
)
|
||||||
from sglang.srt.layers.radix_attention import AttentionType
|
from sglang.srt.layers.radix_attention import AttentionType
|
||||||
from sglang.srt.mem_cache.base_swa_memory_pool import BaseSWAKVPool
|
from sglang.srt.mem_cache.base_swa_memory_pool import BaseSWAKVPool
|
||||||
from sglang.srt.mem_cache.memory_pool import KVWriteLoc
|
from sglang.srt.mem_cache.memory_pool import KVWriteLoc
|
||||||
@@ -48,6 +51,7 @@ from sglang.srt.utils import (
|
|||||||
is_flashinfer_available,
|
is_flashinfer_available,
|
||||||
is_sm100_supported,
|
is_sm100_supported,
|
||||||
next_power_of_2,
|
next_power_of_2,
|
||||||
|
require_gathered_buffer,
|
||||||
)
|
)
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -56,6 +60,19 @@ if TYPE_CHECKING:
|
|||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
def _cuda_graph_capture_max_bs(server_args, max_bs: int) -> int:
|
||||||
|
"""Pad max_bs to the alignment cuda-graph capture uses (see get_batch_sizes_to_capture)."""
|
||||||
|
mul_base = 1
|
||||||
|
if server_args.enable_two_batch_overlap:
|
||||||
|
mul_base *= 2
|
||||||
|
if require_gathered_buffer(server_args):
|
||||||
|
mul_base *= get_attention_tp_size()
|
||||||
|
if mul_base % get_attention_cp_size() != 0:
|
||||||
|
mul_base *= get_attention_cp_size()
|
||||||
|
return (max_bs + mul_base - 1) // mul_base * mul_base
|
||||||
|
|
||||||
|
|
||||||
if envs.SGLANG_ENABLE_TORCH_COMPILE.get():
|
if envs.SGLANG_ENABLE_TORCH_COMPILE.get():
|
||||||
torch._logging.set_logs(dynamo=logging.ERROR)
|
torch._logging.set_logs(dynamo=logging.ERROR)
|
||||||
torch._dynamo.config.suppress_errors = True
|
torch._dynamo.config.suppress_errors = True
|
||||||
@@ -266,7 +283,9 @@ class FlashInferAttnBackend(AttentionBackend):
|
|||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
self.workspace_buffer = global_workspace_buffer
|
self.workspace_buffer = global_workspace_buffer
|
||||||
max_bs = model_runner.req_to_token_pool.size
|
max_bs = _cuda_graph_capture_max_bs(
|
||||||
|
model_runner.server_args, model_runner.req_to_token_pool.size
|
||||||
|
)
|
||||||
if kv_indptr_buf is None:
|
if kv_indptr_buf is None:
|
||||||
self.kv_indptr = [
|
self.kv_indptr = [
|
||||||
torch.zeros(
|
torch.zeros(
|
||||||
@@ -1180,7 +1199,7 @@ class FlashInferIndicesUpdaterDecode:
|
|||||||
fixed_split_size: Optional[int] = None,
|
fixed_split_size: Optional[int] = None,
|
||||||
disable_split_kv: Optional[bool] = None,
|
disable_split_kv: Optional[bool] = None,
|
||||||
):
|
):
|
||||||
if spec_info is None:
|
if spec_info is None or getattr(spec_info, "kv_indptr", None) is None:
|
||||||
bs = len(req_pool_indices)
|
bs = len(req_pool_indices)
|
||||||
kv_indptr[1 : bs + 1] = torch.cumsum(paged_kernel_lens, dim=0)
|
kv_indptr[1 : bs + 1] = torch.cumsum(paged_kernel_lens, dim=0)
|
||||||
kv_indptr = kv_indptr[: bs + 1]
|
kv_indptr = kv_indptr[: bs + 1]
|
||||||
@@ -1646,7 +1665,9 @@ class FlashInferMultiStepDraftBackend:
|
|||||||
self.generate_draft_decode_kv_indices = generate_draft_decode_kv_indices
|
self.generate_draft_decode_kv_indices = generate_draft_decode_kv_indices
|
||||||
self.page_size = model_runner.page_size
|
self.page_size = model_runner.page_size
|
||||||
|
|
||||||
max_bs = model_runner.req_to_token_pool.size * self.topk
|
max_bs = _cuda_graph_capture_max_bs(
|
||||||
|
model_runner.server_args, model_runner.req_to_token_pool.size * self.topk
|
||||||
|
)
|
||||||
self.kv_indptr = torch.zeros(
|
self.kv_indptr = torch.zeros(
|
||||||
(
|
(
|
||||||
self.speculative_num_steps,
|
self.speculative_num_steps,
|
||||||
|
|||||||
@@ -140,9 +140,6 @@ class HybridAttnBackend(AttentionBackend):
|
|||||||
return backend.get_indexer_metadata(layer_id, forward_batch)
|
return backend.get_indexer_metadata(layer_id, forward_batch)
|
||||||
|
|
||||||
def update_mamba_state_after_mtp_verify(self, *args, **kwargs):
|
def update_mamba_state_after_mtp_verify(self, *args, **kwargs):
|
||||||
# Forward to whichever sub-backend handled target_verify, since its inner
|
|
||||||
# linear_attn_backend.forward_metadata holds the mamba_cache_indices the
|
|
||||||
# method consumes. Mirrors _select_backend's target_verify branch.
|
|
||||||
if self.model_runner.server_args.speculative_attention_mode == "decode":
|
if self.model_runner.server_args.speculative_attention_mode == "decode":
|
||||||
backend = self.decode_backend
|
backend = self.decode_backend
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -46,7 +46,12 @@ class MambaAttnBackendBase(AttentionBackend):
|
|||||||
|
|
||||||
def _execute_deferred_mamba_cow_and_clear(self, forward_batch: ForwardBatch):
|
def _execute_deferred_mamba_cow_and_clear(self, forward_batch: ForwardBatch):
|
||||||
"""Run deferred clear/COW ops on the forward stream to avoid races."""
|
"""Run deferred clear/COW ops on the forward stream to avoid races."""
|
||||||
if not forward_batch.forward_mode.is_extend() or self.is_draft_worker:
|
if (
|
||||||
|
not forward_batch.forward_mode.is_extend()
|
||||||
|
or forward_batch.forward_mode.is_target_verify()
|
||||||
|
or forward_batch.forward_mode.is_draft_extend(include_v2=True)
|
||||||
|
or self.is_draft_worker
|
||||||
|
):
|
||||||
return
|
return
|
||||||
if (
|
if (
|
||||||
forward_batch.mamba_clear_indices is not None
|
forward_batch.mamba_clear_indices is not None
|
||||||
@@ -81,6 +86,10 @@ class MambaAttnBackendBase(AttentionBackend):
|
|||||||
mamba_cache_indices = self.req_to_token_pool.get_mamba_indices(
|
mamba_cache_indices = self.req_to_token_pool.get_mamba_indices(
|
||||||
forward_batch.req_pool_indices
|
forward_batch.req_pool_indices
|
||||||
)
|
)
|
||||||
|
_real_bs = getattr(forward_batch, "_original_batch_size", None)
|
||||||
|
if _real_bs is not None and _real_bs < mamba_cache_indices.shape[0]:
|
||||||
|
mamba_cache_indices = mamba_cache_indices.clone()
|
||||||
|
mamba_cache_indices[_real_bs:] = -1
|
||||||
|
|
||||||
if forward_batch.forward_mode.is_decode_or_idle():
|
if forward_batch.forward_mode.is_decode_or_idle():
|
||||||
query_start_loc = torch.arange(
|
query_start_loc = torch.arange(
|
||||||
|
|||||||
@@ -19,6 +19,11 @@ from sglang.srt.layers.attention.mamba.ops import (
|
|||||||
mamba_chunk_scan_combined,
|
mamba_chunk_scan_combined,
|
||||||
selective_state_update,
|
selective_state_update,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.layers.dp_attention import (
|
||||||
|
get_attention_tp_rank,
|
||||||
|
get_attention_tp_size,
|
||||||
|
is_dp_attention_enabled,
|
||||||
|
)
|
||||||
from sglang.srt.layers.linear import (
|
from sglang.srt.layers.linear import (
|
||||||
ColumnParallelLinear,
|
ColumnParallelLinear,
|
||||||
MergedColumnParallelLinear,
|
MergedColumnParallelLinear,
|
||||||
@@ -226,8 +231,12 @@ class MambaMixer2(torch.nn.Module):
|
|||||||
# may be replicated to follow the head shard.
|
# may be replicated to follow the head shard.
|
||||||
# - NOTE: currently for the world size DOES NOT divide groups
|
# - NOTE: currently for the world size DOES NOT divide groups
|
||||||
# case, we only support the case when n_groups == 1
|
# case, we only support the case when n_groups == 1
|
||||||
self.tp_size = get_tensor_model_parallel_world_size()
|
if is_dp_attention_enabled():
|
||||||
self.tp_rank = get_tensor_model_parallel_rank()
|
self.tp_size = get_attention_tp_size()
|
||||||
|
self.tp_rank = get_attention_tp_rank()
|
||||||
|
else:
|
||||||
|
self.tp_size = get_tensor_model_parallel_world_size()
|
||||||
|
self.tp_rank = get_tensor_model_parallel_rank()
|
||||||
|
|
||||||
self.num_heads = num_heads = cache_params.shape.num_heads
|
self.num_heads = num_heads = cache_params.shape.num_heads
|
||||||
self.head_dim = cache_params.shape.head_dim
|
self.head_dim = cache_params.shape.head_dim
|
||||||
@@ -276,6 +285,8 @@ class MambaMixer2(torch.nn.Module):
|
|||||||
bias=use_conv_bias,
|
bias=use_conv_bias,
|
||||||
quant_config=None,
|
quant_config=None,
|
||||||
prefix=f"{prefix}.conv1d",
|
prefix=f"{prefix}.conv1d",
|
||||||
|
tp_rank=self.tp_rank,
|
||||||
|
tp_size=self.tp_size,
|
||||||
)
|
)
|
||||||
|
|
||||||
self.in_proj = MergedColumnParallelLinear(
|
self.in_proj = MergedColumnParallelLinear(
|
||||||
@@ -290,6 +301,8 @@ class MambaMixer2(torch.nn.Module):
|
|||||||
bias=use_bias,
|
bias=use_bias,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=f"{prefix}.in_proj",
|
prefix=f"{prefix}.in_proj",
|
||||||
|
tp_rank=self.tp_rank,
|
||||||
|
tp_size=self.tp_size,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
# This is the n_groups == 1 case,
|
# This is the n_groups == 1 case,
|
||||||
@@ -301,6 +314,8 @@ class MambaMixer2(torch.nn.Module):
|
|||||||
bias=use_conv_bias,
|
bias=use_conv_bias,
|
||||||
quant_config=None,
|
quant_config=None,
|
||||||
prefix=f"{prefix}.conv1d",
|
prefix=f"{prefix}.conv1d",
|
||||||
|
tp_rank=self.tp_rank,
|
||||||
|
tp_size=self.tp_size,
|
||||||
)
|
)
|
||||||
|
|
||||||
self.in_proj = ColumnParallelLinear(
|
self.in_proj = ColumnParallelLinear(
|
||||||
@@ -309,6 +324,8 @@ class MambaMixer2(torch.nn.Module):
|
|||||||
bias=use_bias,
|
bias=use_bias,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=f"{prefix}.in_proj",
|
prefix=f"{prefix}.in_proj",
|
||||||
|
tp_rank=self.tp_rank,
|
||||||
|
tp_size=self.tp_size,
|
||||||
)
|
)
|
||||||
|
|
||||||
# - because in_proj is a concatenation of 3 weights, we
|
# - because in_proj is a concatenation of 3 weights, we
|
||||||
@@ -412,6 +429,9 @@ class MambaMixer2(torch.nn.Module):
|
|||||||
bias=use_bias,
|
bias=use_bias,
|
||||||
input_is_parallel=True,
|
input_is_parallel=True,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
|
tp_rank=self.tp_rank,
|
||||||
|
tp_size=self.tp_size,
|
||||||
|
reduce_results=not is_dp_attention_enabled(),
|
||||||
prefix=f"{prefix}.out_proj",
|
prefix=f"{prefix}.out_proj",
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -447,6 +467,8 @@ class MambaMixer2(torch.nn.Module):
|
|||||||
|
|
||||||
query_start_loc = metadata.query_start_loc
|
query_start_loc = metadata.query_start_loc
|
||||||
|
|
||||||
|
padded_num_tokens = hidden_states.shape[0]
|
||||||
|
|
||||||
# 1. Gated MLP's linear projection
|
# 1. Gated MLP's linear projection
|
||||||
projected_states, _ = self.in_proj(hidden_states)
|
projected_states, _ = self.in_proj(hidden_states)
|
||||||
|
|
||||||
@@ -488,7 +510,12 @@ class MambaMixer2(torch.nn.Module):
|
|||||||
has_prefill = num_prefills > 0
|
has_prefill = num_prefills > 0
|
||||||
has_decode = num_decodes > 0
|
has_decode = num_decodes > 0
|
||||||
num_actual_tokens = num_prefill_tokens + num_decode_tokens
|
num_actual_tokens = num_prefill_tokens + num_decode_tokens
|
||||||
assert num_actual_tokens == projected_states.shape[0]
|
assert num_actual_tokens <= projected_states.shape[0]
|
||||||
|
hidden_states_B_C = hidden_states_B_C[:num_actual_tokens]
|
||||||
|
dt = dt[:num_actual_tokens]
|
||||||
|
|
||||||
|
local_num_heads = self.num_heads // self.tp_size
|
||||||
|
local_num_groups = self.n_groups // self.tp_size
|
||||||
|
|
||||||
# NOTE: V0 put prefill before decode
|
# NOTE: V0 put prefill before decode
|
||||||
# Separate prefill and decode by splitting varlen input
|
# Separate prefill and decode by splitting varlen input
|
||||||
@@ -503,12 +530,10 @@ class MambaMixer2(torch.nn.Module):
|
|||||||
[num_prefill_tokens, num_decode_tokens],
|
[num_prefill_tokens, num_decode_tokens],
|
||||||
dim=0,
|
dim=0,
|
||||||
)
|
)
|
||||||
# Split along batch dimension
|
state_indices_tensor_p = state_indices_tensor[:num_prefills]
|
||||||
state_indices_tensor_p, state_indices_tensor_d = torch.split(
|
state_indices_tensor_d = state_indices_tensor[
|
||||||
state_indices_tensor,
|
num_prefills : num_prefills + num_decodes
|
||||||
[num_prefills, num_decodes],
|
]
|
||||||
dim=0,
|
|
||||||
)
|
|
||||||
query_start_loc_p = query_start_loc[: num_prefills + 1] if has_prefill else None
|
query_start_loc_p = query_start_loc[: num_prefills + 1] if has_prefill else None
|
||||||
|
|
||||||
# Preallocate output tensor to avoid memcpy cost for merging prefill
|
# Preallocate output tensor to avoid memcpy cost for merging prefill
|
||||||
@@ -522,8 +547,9 @@ class MambaMixer2(torch.nn.Module):
|
|||||||
dtype=hidden_states.dtype,
|
dtype=hidden_states.dtype,
|
||||||
device=hidden_states.device,
|
device=hidden_states.device,
|
||||||
)
|
)
|
||||||
|
preallocated_ssm_out_active = preallocated_ssm_out[:num_actual_tokens]
|
||||||
preallocated_ssm_out_p, preallocated_ssm_out_d = torch.split(
|
preallocated_ssm_out_p, preallocated_ssm_out_d = torch.split(
|
||||||
preallocated_ssm_out,
|
preallocated_ssm_out_active,
|
||||||
[num_prefill_tokens, num_decode_tokens],
|
[num_prefill_tokens, num_decode_tokens],
|
||||||
dim=0,
|
dim=0,
|
||||||
)
|
)
|
||||||
@@ -580,12 +606,12 @@ class MambaMixer2(torch.nn.Module):
|
|||||||
# NOTE: final output is an in-place update of out tensor
|
# NOTE: final output is an in-place update of out tensor
|
||||||
intermediate_states, varlen_state = mamba_chunk_scan_combined(
|
intermediate_states, varlen_state = mamba_chunk_scan_combined(
|
||||||
hidden_states_p.view(
|
hidden_states_p.view(
|
||||||
1, num_prefill_tokens, self.num_heads // self.tp_size, self.head_dim
|
1, num_prefill_tokens, local_num_heads, self.head_dim
|
||||||
),
|
),
|
||||||
dt_p.unsqueeze(0),
|
dt_p.unsqueeze(0),
|
||||||
self.A,
|
self.A,
|
||||||
B_p.view(1, num_prefill_tokens, self.n_groups // self.tp_size, -1),
|
B_p.view(1, num_prefill_tokens, local_num_groups, -1),
|
||||||
C_p.view(1, num_prefill_tokens, self.n_groups // self.tp_size, -1),
|
C_p.view(1, num_prefill_tokens, local_num_groups, -1),
|
||||||
chunk_size=mixed_metadata.chunk_size,
|
chunk_size=mixed_metadata.chunk_size,
|
||||||
D=self.D,
|
D=self.D,
|
||||||
z=None,
|
z=None,
|
||||||
@@ -608,7 +634,8 @@ class MambaMixer2(torch.nn.Module):
|
|||||||
|
|
||||||
# update ssm states
|
# update ssm states
|
||||||
# - varlen state is a (num_prefills, nheads, headdim, dstate) tensor
|
# - varlen state is a (num_prefills, nheads, headdim, dstate) tensor
|
||||||
ssm_state[state_indices_tensor_p] = varlen_state
|
if varlen_state is not None:
|
||||||
|
ssm_state[state_indices_tensor_p] = varlen_state
|
||||||
|
|
||||||
# Process decode requests
|
# Process decode requests
|
||||||
if has_decode:
|
if has_decode:
|
||||||
@@ -666,7 +693,7 @@ class MambaMixer2(torch.nn.Module):
|
|||||||
hidden_states_d, B_d, C_d = split_hidden_states_B_C_fn(hidden_states_B_C_d)
|
hidden_states_d, B_d, C_d = split_hidden_states_B_C_fn(hidden_states_B_C_d)
|
||||||
|
|
||||||
# 3. State Space Model sequence transformation
|
# 3. State Space Model sequence transformation
|
||||||
n_groups = self.n_groups // self.tp_size
|
n_groups = local_num_groups
|
||||||
A_d = (
|
A_d = (
|
||||||
self.A[:, None, ...][:, :, None]
|
self.A[:, None, ...][:, :, None]
|
||||||
.expand(-1, self.head_dim, self.ssm_state_size)
|
.expand(-1, self.head_dim, self.ssm_state_size)
|
||||||
@@ -677,9 +704,7 @@ class MambaMixer2(torch.nn.Module):
|
|||||||
D_d = self.D[:, None, ...].expand(-1, self.head_dim)
|
D_d = self.D[:, None, ...].expand(-1, self.head_dim)
|
||||||
B_d = B_d.view(-1, n_groups, B_d.shape[1] // n_groups)
|
B_d = B_d.view(-1, n_groups, B_d.shape[1] // n_groups)
|
||||||
C_d = C_d.view(-1, n_groups, C_d.shape[1] // n_groups)
|
C_d = C_d.view(-1, n_groups, C_d.shape[1] // n_groups)
|
||||||
hidden_states_d = hidden_states_d.view(
|
hidden_states_d = hidden_states_d.view(-1, local_num_heads, self.head_dim)
|
||||||
-1, self.num_heads // self.tp_size, self.head_dim
|
|
||||||
)
|
|
||||||
|
|
||||||
if is_target_verify:
|
if is_target_verify:
|
||||||
selective_state_update(
|
selective_state_update(
|
||||||
@@ -736,12 +761,11 @@ class MambaMixer2(torch.nn.Module):
|
|||||||
# GatedRMSNorm internally applying SiLU to the gate
|
# GatedRMSNorm internally applying SiLU to the gate
|
||||||
# SiLU is applied internally before normalization, unlike standard
|
# SiLU is applied internally before normalization, unlike standard
|
||||||
# norm usage
|
# norm usage
|
||||||
hidden_states = self.norm(preallocated_ssm_out, gate[:num_actual_tokens])
|
hidden_states = self.norm(preallocated_ssm_out, gate)
|
||||||
|
|
||||||
# 5. Final linear projection
|
|
||||||
mixer_out, _ = self.out_proj(hidden_states)
|
mixer_out, _ = self.out_proj(hidden_states)
|
||||||
if output is not None:
|
if output is not None:
|
||||||
output[:num_actual_tokens].copy_(mixer_out)
|
output[:padded_num_tokens].copy_(mixer_out)
|
||||||
|
|
||||||
return mixer_out, intermediate_states
|
return mixer_out, intermediate_states
|
||||||
|
|
||||||
|
|||||||
@@ -163,6 +163,7 @@ class Mamba2Metadata(ForwardMetadata):
|
|||||||
*,
|
*,
|
||||||
is_target_verify: bool,
|
is_target_verify: bool,
|
||||||
draft_token_num: int,
|
draft_token_num: int,
|
||||||
|
num_decodes: Optional[int] = None,
|
||||||
) -> "Mamba2Metadata":
|
) -> "Mamba2Metadata":
|
||||||
"""This path is run during CUDA graph capture, i.e. decode only, so `num_prefills` is 0"""
|
"""This path is run during CUDA graph capture, i.e. decode only, so `num_prefills` is 0"""
|
||||||
return Mamba2Metadata(
|
return Mamba2Metadata(
|
||||||
@@ -177,7 +178,7 @@ class Mamba2Metadata(ForwardMetadata):
|
|||||||
track_ssm_final_src=forward_metadata.track_ssm_final_src,
|
track_ssm_final_src=forward_metadata.track_ssm_final_src,
|
||||||
track_ssm_final_dst=forward_metadata.track_ssm_final_dst,
|
track_ssm_final_dst=forward_metadata.track_ssm_final_dst,
|
||||||
has_mamba_track_mask=forward_metadata.has_mamba_track_mask,
|
has_mamba_track_mask=forward_metadata.has_mamba_track_mask,
|
||||||
num_decodes=len(seq_lens),
|
num_decodes=len(seq_lens) if num_decodes is None else num_decodes,
|
||||||
num_prefills=0,
|
num_prefills=0,
|
||||||
num_prefill_tokens=0,
|
num_prefill_tokens=0,
|
||||||
is_target_verify=is_target_verify,
|
is_target_verify=is_target_verify,
|
||||||
@@ -198,28 +199,49 @@ class Mamba2Metadata(ForwardMetadata):
|
|||||||
if forward_batch.spec_info is not None
|
if forward_batch.spec_info is not None
|
||||||
else 1
|
else 1
|
||||||
)
|
)
|
||||||
|
num_decodes = getattr(forward_batch, "_original_batch_size", None)
|
||||||
|
if num_decodes is None:
|
||||||
|
num_decodes = len(forward_batch.seq_lens)
|
||||||
return cls.prepare_decode(
|
return cls.prepare_decode(
|
||||||
forward_metadata,
|
forward_metadata,
|
||||||
forward_batch.seq_lens,
|
forward_batch.seq_lens,
|
||||||
is_target_verify=forward_batch.forward_mode.is_target_verify(),
|
is_target_verify=forward_batch.forward_mode.is_target_verify(),
|
||||||
draft_token_num=draft_token_num,
|
draft_token_num=draft_token_num,
|
||||||
|
num_decodes=num_decodes,
|
||||||
)
|
)
|
||||||
num_prefills = len(forward_batch.extend_seq_lens)
|
extend_seq_lens_cpu = forward_batch.extend_seq_lens_cpu
|
||||||
num_prefill_tokens = forward_batch.extend_num_tokens
|
if extend_seq_lens_cpu is None:
|
||||||
num_decodes = len(forward_batch.seq_lens) - num_prefills
|
num_prefills = len(forward_batch.extend_seq_lens)
|
||||||
|
else:
|
||||||
|
num_prefills = len(extend_seq_lens_cpu)
|
||||||
|
if extend_seq_lens_cpu is not None:
|
||||||
|
num_prefill_tokens = int(sum(extend_seq_lens_cpu))
|
||||||
|
else:
|
||||||
|
num_prefill_tokens = int(forward_batch.extend_num_tokens)
|
||||||
|
batch_size = getattr(forward_batch, "_original_batch_size", None)
|
||||||
|
if batch_size is None:
|
||||||
|
batch_size = len(forward_batch.seq_lens)
|
||||||
|
num_decodes = batch_size - num_prefills
|
||||||
context_lens_tensor = forward_batch.extend_prefix_lens
|
context_lens_tensor = forward_batch.extend_prefix_lens
|
||||||
assert context_lens_tensor is not None
|
assert context_lens_tensor is not None
|
||||||
# precompute flag to avoid device syncs later
|
|
||||||
has_initial_states = context_lens_tensor > 0
|
has_initial_states = context_lens_tensor > 0
|
||||||
|
mamba_track_mask = getattr(forward_batch, "mamba_track_mask", None)
|
||||||
|
if mamba_track_mask is not None:
|
||||||
|
has_initial_states = (
|
||||||
|
has_initial_states & mamba_track_mask[: has_initial_states.shape[0]]
|
||||||
|
)
|
||||||
prep_initial_states = torch.any(has_initial_states[:num_prefills]).item()
|
prep_initial_states = torch.any(has_initial_states[:num_prefills]).item()
|
||||||
|
|
||||||
query_start_loc = forward_metadata.query_start_loc[: num_prefills + 1]
|
query_start_loc = forward_metadata.query_start_loc[: num_prefills + 1]
|
||||||
|
_seq_idx_output_size = (
|
||||||
|
num_prefill_tokens if extend_seq_lens_cpu is not None else None
|
||||||
|
)
|
||||||
seq_idx = torch.repeat_interleave(
|
seq_idx = torch.repeat_interleave(
|
||||||
torch.arange(
|
torch.arange(
|
||||||
num_prefills, dtype=torch.int32, device=query_start_loc.device
|
num_prefills, dtype=torch.int32, device=query_start_loc.device
|
||||||
),
|
),
|
||||||
query_start_loc.diff(),
|
query_start_loc.diff(),
|
||||||
output_size=num_prefill_tokens,
|
output_size=_seq_idx_output_size,
|
||||||
)
|
)
|
||||||
seq_idx.unsqueeze_(0)
|
seq_idx.unsqueeze_(0)
|
||||||
|
|
||||||
@@ -263,6 +285,6 @@ class Mamba2Metadata(ForwardMetadata):
|
|||||||
seq_idx=seq_idx,
|
seq_idx=seq_idx,
|
||||||
chunk_indices=chunk_indices,
|
chunk_indices=chunk_indices,
|
||||||
chunk_offsets=chunk_offsets,
|
chunk_offsets=chunk_offsets,
|
||||||
extend_seq_lens_cpu=forward_batch.extend_seq_lens_cpu,
|
extend_seq_lens_cpu=extend_seq_lens_cpu,
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -11,6 +11,13 @@ from sglang.srt.distributed.parallel_state import (
|
|||||||
get_tensor_model_parallel_world_size,
|
get_tensor_model_parallel_world_size,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.attention.fla.layernorm_gated import rms_norm_gated
|
from sglang.srt.layers.attention.fla.layernorm_gated import rms_norm_gated
|
||||||
|
from sglang.srt.layers.dp_attention import (
|
||||||
|
attn_tp_all_reduce,
|
||||||
|
get_attention_tp_group,
|
||||||
|
get_attention_tp_rank,
|
||||||
|
get_attention_tp_size,
|
||||||
|
is_dp_attention_enabled,
|
||||||
|
)
|
||||||
from sglang.srt.layers.utils import MultiPlatformOp
|
from sglang.srt.layers.utils import MultiPlatformOp
|
||||||
from sglang.srt.model_loader.weight_utils import sharded_weight_loader
|
from sglang.srt.model_loader.weight_utils import sharded_weight_loader
|
||||||
from sglang.srt.utils.common import set_weight_attrs
|
from sglang.srt.utils.common import set_weight_attrs
|
||||||
@@ -25,8 +32,13 @@ class Mixer2RMSNormGated(MultiPlatformOp):
|
|||||||
eps: float = 1e-6,
|
eps: float = 1e-6,
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.tp_size = get_tensor_model_parallel_world_size()
|
self.use_attn_tp_group = is_dp_attention_enabled()
|
||||||
self.tp_rank = get_tensor_model_parallel_rank()
|
if self.use_attn_tp_group:
|
||||||
|
self.tp_size = get_attention_tp_size()
|
||||||
|
self.tp_rank = get_attention_tp_rank()
|
||||||
|
else:
|
||||||
|
self.tp_size = get_tensor_model_parallel_world_size()
|
||||||
|
self.tp_rank = get_tensor_model_parallel_rank()
|
||||||
self.full_hidden_size = full_hidden_size
|
self.full_hidden_size = full_hidden_size
|
||||||
self.group_size = full_hidden_size // full_n_groups
|
self.group_size = full_hidden_size // full_n_groups
|
||||||
self.per_rank_hidden_size = full_hidden_size // self.tp_size
|
self.per_rank_hidden_size = full_hidden_size // self.tp_size
|
||||||
@@ -68,7 +80,10 @@ class Mixer2RMSNormGated(MultiPlatformOp):
|
|||||||
if self.tp_size > 1:
|
if self.tp_size > 1:
|
||||||
# Compute local sum and then reduce to obtain global sum
|
# Compute local sum and then reduce to obtain global sum
|
||||||
local_sums = x.pow(2).sum(dim=-1, keepdim=True)
|
local_sums = x.pow(2).sum(dim=-1, keepdim=True)
|
||||||
global_sums = tensor_model_parallel_all_reduce(local_sums)
|
if self.use_attn_tp_group:
|
||||||
|
global_sums = attn_tp_all_reduce(local_sums)
|
||||||
|
else:
|
||||||
|
global_sums = tensor_model_parallel_all_reduce(local_sums)
|
||||||
# Calculate the variance
|
# Calculate the variance
|
||||||
count = self.tp_size * x.shape[-1]
|
count = self.tp_size * x.shape[-1]
|
||||||
variance = global_sums / count
|
variance = global_sums / count
|
||||||
@@ -80,7 +95,12 @@ class Mixer2RMSNormGated(MultiPlatformOp):
|
|||||||
redundant_tp: bool = self.n_groups % self.tp_size != 0
|
redundant_tp: bool = self.n_groups % self.tp_size != 0
|
||||||
if redundant_tp:
|
if redundant_tp:
|
||||||
# To handle the general case, redundantly apply the variance
|
# To handle the general case, redundantly apply the variance
|
||||||
x = tensor_model_parallel_all_gather(x, -1)
|
if self.use_attn_tp_group:
|
||||||
|
parts = [torch.empty_like(x) for _ in range(self.tp_size)]
|
||||||
|
get_attention_tp_group().all_gather(x, output_tensor_list=parts)
|
||||||
|
x = torch.cat(parts, dim=-1)
|
||||||
|
else:
|
||||||
|
x = tensor_model_parallel_all_gather(x, -1)
|
||||||
|
|
||||||
*prefix_dims, hidden_dim = x.shape
|
*prefix_dims, hidden_dim = x.shape
|
||||||
group_count = hidden_dim // self.group_size
|
group_count = hidden_dim // self.group_size
|
||||||
|
|||||||
@@ -41,6 +41,7 @@ from sglang.srt.layers.dp_attention import (
|
|||||||
attn_tp_all_gather_into_tensor,
|
attn_tp_all_gather_into_tensor,
|
||||||
attn_tp_reduce_scatter_tensor,
|
attn_tp_reduce_scatter_tensor,
|
||||||
dp_gather_partial,
|
dp_gather_partial,
|
||||||
|
dp_gather_replicate,
|
||||||
dp_reduce_scatter_tensor,
|
dp_reduce_scatter_tensor,
|
||||||
dp_scatter,
|
dp_scatter,
|
||||||
get_attention_cp_rank,
|
get_attention_cp_rank,
|
||||||
@@ -448,6 +449,7 @@ class LayerCommunicator:
|
|||||||
allow_reduce_scatter: bool = False,
|
allow_reduce_scatter: bool = False,
|
||||||
is_last_layer: bool = False,
|
is_last_layer: bool = False,
|
||||||
qkv_latent_func: Optional[Callable] = None,
|
qkv_latent_func: Optional[Callable] = None,
|
||||||
|
force_layernorm_before_dp_gather: bool = False,
|
||||||
):
|
):
|
||||||
self.layer_scatter_modes = layer_scatter_modes
|
self.layer_scatter_modes = layer_scatter_modes
|
||||||
self.input_layernorm = input_layernorm
|
self.input_layernorm = input_layernorm
|
||||||
@@ -455,8 +457,12 @@ class LayerCommunicator:
|
|||||||
self.allow_reduce_scatter = allow_reduce_scatter
|
self.allow_reduce_scatter = allow_reduce_scatter
|
||||||
self.is_last_layer = is_last_layer
|
self.is_last_layer = is_last_layer
|
||||||
self.qkv_latent_func = qkv_latent_func
|
self.qkv_latent_func = qkv_latent_func
|
||||||
|
self.force_layernorm_before_dp_gather = force_layernorm_before_dp_gather
|
||||||
|
|
||||||
self._context = CommunicateContext.init_new()
|
self._context = CommunicateContext.init_new()
|
||||||
|
self._context.force_layernorm_before_dp_gather = (
|
||||||
|
force_layernorm_before_dp_gather
|
||||||
|
)
|
||||||
self._post_init_communicate()
|
self._post_init_communicate()
|
||||||
self._speculative_algo = SpeculativeAlgorithm.from_string(
|
self._speculative_algo = SpeculativeAlgorithm.from_string(
|
||||||
get_global_server_args().speculative_algorithm
|
get_global_server_args().speculative_algorithm
|
||||||
@@ -787,6 +793,7 @@ class CommunicateContext:
|
|||||||
tp_size: int
|
tp_size: int
|
||||||
cache = None
|
cache = None
|
||||||
tp_rank: int
|
tp_rank: int
|
||||||
|
force_layernorm_before_dp_gather: bool = False
|
||||||
|
|
||||||
def is_same_group_size(self, a: ScatterMode, b: ScatterMode):
|
def is_same_group_size(self, a: ScatterMode, b: ScatterMode):
|
||||||
return self.process_group_sizes[a] == self.process_group_sizes[b]
|
return self.process_group_sizes[a] == self.process_group_sizes[b]
|
||||||
@@ -1029,9 +1036,14 @@ class CommunicateWithAllReduceAndLayerNormFn:
|
|||||||
)
|
)
|
||||||
attn_tp_all_gather_into_tensor(residual, local_residual)
|
attn_tp_all_gather_into_tensor(residual, local_residual)
|
||||||
if context.attn_dp_size != 1:
|
if context.attn_dp_size != 1:
|
||||||
# Perform layernorm on smaller data before comm. Only valid when attn_tp_size is 1 (tp_size == dp_size)
|
use_layer_norm_before_gather = (
|
||||||
use_layer_norm_before_gather = context.attn_tp_size == 1
|
context.force_layernorm_before_dp_gather or context.attn_tp_size == 1
|
||||||
|
)
|
||||||
if use_layer_norm_before_gather and hidden_states.shape[0] != 0:
|
if use_layer_norm_before_gather and hidden_states.shape[0] != 0:
|
||||||
|
if context.attn_tp_size > 1:
|
||||||
|
hidden_states = attention_tensor_model_parallel_all_reduce(
|
||||||
|
hidden_states
|
||||||
|
)
|
||||||
with use_symmetric_memory(
|
with use_symmetric_memory(
|
||||||
get_tp_group(),
|
get_tp_group(),
|
||||||
disabled=not is_allocation_symmetric(),
|
disabled=not is_allocation_symmetric(),
|
||||||
@@ -1044,7 +1056,10 @@ class CommunicateWithAllReduceAndLayerNormFn:
|
|||||||
get_global_dp_buffer(get_tp_group()),
|
get_global_dp_buffer(get_tp_group()),
|
||||||
hidden_states,
|
hidden_states,
|
||||||
)
|
)
|
||||||
dp_gather_partial(hidden_states, local_hidden_states, forward_batch)
|
if use_layer_norm_before_gather:
|
||||||
|
dp_gather_replicate(hidden_states, local_hidden_states, forward_batch)
|
||||||
|
else:
|
||||||
|
dp_gather_partial(hidden_states, local_hidden_states, forward_batch)
|
||||||
|
|
||||||
if not use_layer_norm_before_gather:
|
if not use_layer_norm_before_gather:
|
||||||
dp_scatter(residual, hidden_states, forward_batch)
|
dp_scatter(residual, hidden_states, forward_batch)
|
||||||
|
|||||||
@@ -45,6 +45,7 @@ _ATTN_DP_SIZE: Optional[int] = None
|
|||||||
_LOCAL_ATTN_DP_SIZE: Optional[int] = None
|
_LOCAL_ATTN_DP_SIZE: Optional[int] = None
|
||||||
_LOCAL_ATTN_DP_RANK: Optional[int] = None
|
_LOCAL_ATTN_DP_RANK: Optional[int] = None
|
||||||
_ENABLE_DP_ATTENTION_FLAG: bool = False
|
_ENABLE_DP_ATTENTION_FLAG: bool = False
|
||||||
|
_DP_MAX_LEN_WITH_IDLE = False
|
||||||
|
|
||||||
_is_hip = is_hip()
|
_is_hip = is_hip()
|
||||||
_USE_ROCM700A_WA = _is_hip and get_bool_env_var("SGLANG_USE_ROCM700A")
|
_USE_ROCM700A_WA = _is_hip and get_bool_env_var("SGLANG_USE_ROCM700A")
|
||||||
@@ -74,6 +75,10 @@ class DpPaddingMode(IntEnum):
|
|||||||
# For dp_size=1, max_len equals sum_len, so prefer MAX_LEN mode
|
# For dp_size=1, max_len equals sum_len, so prefer MAX_LEN mode
|
||||||
# to enable symmetric memory optimization (needed for DSA CP, etc.).
|
# to enable symmetric memory optimization (needed for DSA CP, etc.).
|
||||||
if is_extend_in_batch and dp_size > 1:
|
if is_extend_in_batch and dp_size > 1:
|
||||||
|
# Hybrid-SSM models materialize idle ranks via the MAX_LEN
|
||||||
|
# fabricated-row conversion; other models keep mainline SUM_LEN.
|
||||||
|
if _DP_MAX_LEN_WITH_IDLE and min(global_num_tokens) == 0:
|
||||||
|
return DpPaddingMode.MAX_LEN
|
||||||
return DpPaddingMode.SUM_LEN
|
return DpPaddingMode.SUM_LEN
|
||||||
|
|
||||||
# we choose the mode that minimizes the communication cost
|
# we choose the mode that minimizes the communication cost
|
||||||
@@ -277,6 +282,10 @@ def initialize_dp_attention(
|
|||||||
):
|
):
|
||||||
global _ATTN_DP_RANK, _ATTN_DP_SIZE
|
global _ATTN_DP_RANK, _ATTN_DP_SIZE
|
||||||
global _LOCAL_ATTN_DP_SIZE, _LOCAL_ATTN_DP_RANK, _ENABLE_DP_ATTENTION_FLAG
|
global _LOCAL_ATTN_DP_SIZE, _LOCAL_ATTN_DP_RANK, _ENABLE_DP_ATTENTION_FLAG
|
||||||
|
global _DP_MAX_LEN_WITH_IDLE
|
||||||
|
_DP_MAX_LEN_WITH_IDLE = (
|
||||||
|
getattr(model_config.hf_config, "hybrid_override_pattern", None) is not None
|
||||||
|
)
|
||||||
enable_dp_attention = server_args.enable_dp_attention
|
enable_dp_attention = server_args.enable_dp_attention
|
||||||
dp_size = server_args.dp_size
|
dp_size = server_args.dp_size
|
||||||
moe_dense_tp_size = server_args.moe_dense_tp_size
|
moe_dense_tp_size = server_args.moe_dense_tp_size
|
||||||
|
|||||||
@@ -406,6 +406,8 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
|||||||
|
|
||||||
# For DP attention (MLP sync sizes)
|
# For DP attention (MLP sync sizes)
|
||||||
original_global_num_tokens_cpu: Optional[List[int]] = None
|
original_global_num_tokens_cpu: Optional[List[int]] = None
|
||||||
|
_original_batch_size: Optional[int] = None
|
||||||
|
_original_forward_mode: Optional[ForwardMode] = None
|
||||||
global_num_tokens_cpu: Optional[List[int]] = None
|
global_num_tokens_cpu: Optional[List[int]] = None
|
||||||
global_num_tokens_gpu: Optional[torch.Tensor] = None
|
global_num_tokens_gpu: Optional[torch.Tensor] = None
|
||||||
# Has to be None when cuda graph is captured.
|
# Has to be None when cuda graph is captured.
|
||||||
@@ -1082,6 +1084,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
|||||||
assert self.global_num_tokens_cpu is not None
|
assert self.global_num_tokens_cpu is not None
|
||||||
assert self.global_num_tokens_for_logprob_cpu is not None
|
assert self.global_num_tokens_for_logprob_cpu is not None
|
||||||
|
|
||||||
|
self._original_batch_size = self.batch_size
|
||||||
global_num_tokens = self.global_num_tokens_cpu
|
global_num_tokens = self.global_num_tokens_cpu
|
||||||
sync_group_size = len(global_num_tokens)
|
sync_group_size = len(global_num_tokens)
|
||||||
attn_tp_size = get_attention_tp_size()
|
attn_tp_size = get_attention_tp_size()
|
||||||
@@ -1135,20 +1138,68 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
|||||||
or self.forward_mode.is_draft_extend(include_v2=True)
|
or self.forward_mode.is_draft_extend(include_v2=True)
|
||||||
or self.forward_mode.is_idle()
|
or self.forward_mode.is_idle()
|
||||||
):
|
):
|
||||||
if self.is_extend_in_batch and dp_padding_mode.is_max_len():
|
# Mamba-hybrid families need the fabricated-row idle conversion
|
||||||
setattr(self, "_original_forward_mode", self.forward_mode)
|
# below; this includes their MTP draft workers, whose mamba-less
|
||||||
self.forward_mode = ForwardMode.EXTEND
|
# "*E" pattern makes mambaish_config return None.
|
||||||
self.extend_num_tokens = bs
|
hybrid_ssm = model_runner.mambaish_config is not None or (
|
||||||
self.extend_seq_lens = torch.full_like(self.seq_lens, 1)
|
model_runner.is_draft_worker
|
||||||
self.extend_prefix_lens = self.seq_lens - 1
|
and getattr(
|
||||||
self.extend_start_loc = torch.arange(
|
model_runner.model_config.hf_config,
|
||||||
bs, dtype=torch.int32, device=self.seq_lens.device
|
"mtp_hybrid_override_pattern",
|
||||||
|
None,
|
||||||
)
|
)
|
||||||
self.extend_prefix_lens_cpu = self.extend_prefix_lens.cpu()
|
is not None
|
||||||
self.extend_seq_lens_cpu = self.extend_seq_lens.cpu()
|
)
|
||||||
self.extend_logprob_start_lens_cpu = self.extend_prefix_lens_cpu
|
if (
|
||||||
|
hybrid_ssm
|
||||||
|
and self.spec_info is not None
|
||||||
|
and not self.spec_info.is_draft_input()
|
||||||
|
):
|
||||||
|
if self.forward_mode.is_idle():
|
||||||
|
self._original_forward_mode = self.forward_mode
|
||||||
|
self.forward_mode = ForwardMode.TARGET_VERIFY
|
||||||
|
bs = self.batch_size = num_tokens // self.spec_info.num_tokens_per_req
|
||||||
|
elif self.is_extend_in_batch and dp_padding_mode.is_max_len():
|
||||||
|
self._original_forward_mode = self.forward_mode
|
||||||
|
self.forward_mode = ForwardMode.EXTEND
|
||||||
|
if hybrid_ssm:
|
||||||
|
dev = self.seq_lens.device
|
||||||
|
assert (
|
||||||
|
self.seq_lens.shape[0] == 0
|
||||||
|
), "extend-idle conversion expects an empty rank"
|
||||||
|
self.extend_num_tokens = num_tokens
|
||||||
|
self.extend_seq_lens = torch.tensor(
|
||||||
|
[num_tokens], dtype=torch.int32, device=dev
|
||||||
|
)
|
||||||
|
self.extend_prefix_lens = torch.zeros(
|
||||||
|
1, dtype=self.seq_lens.dtype, device=dev
|
||||||
|
)
|
||||||
|
self.extend_start_loc = torch.zeros(
|
||||||
|
1, dtype=torch.int32, device=dev
|
||||||
|
)
|
||||||
|
self.seq_lens = torch.tensor(
|
||||||
|
[num_tokens], dtype=self.seq_lens.dtype, device=dev
|
||||||
|
)
|
||||||
|
self.seq_lens_sum = int(num_tokens)
|
||||||
|
if self.seq_lens_cpu is not None:
|
||||||
|
self.seq_lens_cpu = torch.tensor(
|
||||||
|
[num_tokens], dtype=self.seq_lens.dtype
|
||||||
|
)
|
||||||
|
self.extend_prefix_lens_cpu = [0]
|
||||||
|
self.extend_seq_lens_cpu = [int(num_tokens)]
|
||||||
|
self.extend_logprob_start_lens_cpu = [0]
|
||||||
|
bs = self.batch_size = 1
|
||||||
|
else:
|
||||||
|
self.extend_num_tokens = bs
|
||||||
|
self.extend_seq_lens = torch.full_like(self.seq_lens, 1)
|
||||||
|
self.extend_prefix_lens = self.seq_lens - 1
|
||||||
|
self.extend_start_loc = torch.arange(
|
||||||
|
bs, dtype=torch.int32, device=self.seq_lens.device
|
||||||
|
)
|
||||||
|
self.extend_prefix_lens_cpu = self.extend_prefix_lens.cpu()
|
||||||
|
self.extend_seq_lens_cpu = self.extend_seq_lens.cpu()
|
||||||
|
self.extend_logprob_start_lens_cpu = self.extend_prefix_lens_cpu
|
||||||
else:
|
else:
|
||||||
setattr(self, "_original_batch_size", self.batch_size)
|
|
||||||
if self.spec_info is not None:
|
if self.spec_info is not None:
|
||||||
bs = self.batch_size = (
|
bs = self.batch_size = (
|
||||||
num_tokens // self.spec_info.num_tokens_per_req
|
num_tokens // self.spec_info.num_tokens_per_req
|
||||||
@@ -1277,8 +1328,10 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
|||||||
self._pad_inputs_to_size(model_runner, tokens_padded, self.batch_size)
|
self._pad_inputs_to_size(model_runner, tokens_padded, self.batch_size)
|
||||||
|
|
||||||
def post_forward_mlp_sync_batch(self, logits_output: LogitsProcessorOutput):
|
def post_forward_mlp_sync_batch(self, logits_output: LogitsProcessorOutput):
|
||||||
self.forward_mode = getattr(self, "_original_forward_mode", self.forward_mode)
|
if self._original_forward_mode is not None:
|
||||||
self.batch_size = getattr(self, "_original_batch_size", self.batch_size)
|
self.forward_mode = self._original_forward_mode
|
||||||
|
if self._original_batch_size is not None:
|
||||||
|
self.batch_size = self._original_batch_size
|
||||||
bs = self.batch_size
|
bs = self.batch_size
|
||||||
|
|
||||||
if self.spec_info is not None:
|
if self.spec_info is not None:
|
||||||
|
|||||||
@@ -38,6 +38,12 @@ from sglang.srt.layers.attention.hybrid_linear_attn_backend import (
|
|||||||
Mamba2AttnBackend,
|
Mamba2AttnBackend,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.attention.mamba.mamba import MambaMixer2
|
from sglang.srt.layers.attention.mamba.mamba import MambaMixer2
|
||||||
|
from sglang.srt.layers.dp_attention import (
|
||||||
|
attn_tp_all_reduce,
|
||||||
|
get_attention_tp_rank,
|
||||||
|
get_attention_tp_size,
|
||||||
|
is_dp_attention_enabled,
|
||||||
|
)
|
||||||
from sglang.srt.layers.layernorm import RMSNorm
|
from sglang.srt.layers.layernorm import RMSNorm
|
||||||
from sglang.srt.layers.linear import (
|
from sglang.srt.layers.linear import (
|
||||||
ColumnParallelLinear,
|
ColumnParallelLinear,
|
||||||
@@ -74,6 +80,12 @@ from sglang.srt.model_loader.weight_utils import (
|
|||||||
replace_prefix,
|
replace_prefix,
|
||||||
replace_substrings,
|
replace_substrings,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.models.nemotron_h_utils import (
|
||||||
|
get_real_num_tokens,
|
||||||
|
is_attn_layer,
|
||||||
|
make_layer_communicator,
|
||||||
|
pad_to_original_num_tokens,
|
||||||
|
)
|
||||||
from sglang.srt.models.utils import WeightsMapper
|
from sglang.srt.models.utils import WeightsMapper
|
||||||
from sglang.srt.server_args import get_global_server_args
|
from sglang.srt.server_args import get_global_server_args
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
@@ -302,7 +314,37 @@ class NemotronHMoE(nn.Module):
|
|||||||
return final_hidden_states.view(num_tokens, hidden_dim)
|
return final_hidden_states.view(num_tokens, hidden_dim)
|
||||||
|
|
||||||
|
|
||||||
class NemotronHMLPDecoderLayer(nn.Module):
|
class NemotronHMLPLikeDecoderLayer(nn.Module):
|
||||||
|
"""Shared forward for the dense-MLP / MoE decoder layers."""
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
hidden_states: torch.Tensor,
|
||||||
|
residual: Optional[torch.Tensor],
|
||||||
|
forward_batch: ForwardBatch,
|
||||||
|
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||||
|
if is_dp_attention_enabled():
|
||||||
|
hidden_states, residual = self.layer_communicator.prepare_mlp(
|
||||||
|
hidden_states, residual, forward_batch
|
||||||
|
)
|
||||||
|
hidden_states = self.mixer.forward(hidden_states)
|
||||||
|
hidden_states, residual = self.layer_communicator.postprocess_layer(
|
||||||
|
hidden_states, residual, forward_batch
|
||||||
|
)
|
||||||
|
return hidden_states, residual
|
||||||
|
|
||||||
|
if residual is None:
|
||||||
|
residual = hidden_states
|
||||||
|
hidden_states = self.norm(hidden_states)
|
||||||
|
else:
|
||||||
|
hidden_states, residual = self.norm(hidden_states, residual)
|
||||||
|
|
||||||
|
hidden_states = self.mixer.forward(hidden_states)
|
||||||
|
return hidden_states, residual
|
||||||
|
|
||||||
|
|
||||||
|
class NemotronHMLPDecoderLayer(NemotronHMLPLikeDecoderLayer):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
config: NemotronHConfig,
|
config: NemotronHConfig,
|
||||||
@@ -315,6 +357,7 @@ class NemotronHMLPDecoderLayer(nn.Module):
|
|||||||
|
|
||||||
hybrid_override_pattern = config.hybrid_override_pattern
|
hybrid_override_pattern = config.hybrid_override_pattern
|
||||||
mlp_index = hybrid_override_pattern[: layer_idx + 1].count("-") - 1
|
mlp_index = hybrid_override_pattern[: layer_idx + 1].count("-") - 1
|
||||||
|
self.layer_idx = layer_idx
|
||||||
if isinstance(config.intermediate_size, list):
|
if isinstance(config.intermediate_size, list):
|
||||||
if len(config.intermediate_size) == 1:
|
if len(config.intermediate_size) == 1:
|
||||||
intermediate_size = config.intermediate_size[0]
|
intermediate_size = config.intermediate_size[0]
|
||||||
@@ -332,25 +375,10 @@ class NemotronHMLPDecoderLayer(nn.Module):
|
|||||||
)
|
)
|
||||||
|
|
||||||
self.norm = RMSNorm(config.hidden_size, eps=config.layer_norm_epsilon)
|
self.norm = RMSNorm(config.hidden_size, eps=config.layer_norm_epsilon)
|
||||||
|
self.layer_communicator = make_layer_communicator(self.norm, for_attn=False)
|
||||||
def forward(
|
|
||||||
self,
|
|
||||||
*,
|
|
||||||
hidden_states: torch.Tensor,
|
|
||||||
residual: Optional[torch.Tensor],
|
|
||||||
forward_batch: ForwardBatch,
|
|
||||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
|
||||||
if residual is None:
|
|
||||||
residual = hidden_states
|
|
||||||
hidden_states = self.norm(hidden_states)
|
|
||||||
else:
|
|
||||||
hidden_states, residual = self.norm(hidden_states, residual)
|
|
||||||
|
|
||||||
hidden_states = self.mixer.forward(hidden_states)
|
|
||||||
return hidden_states, residual
|
|
||||||
|
|
||||||
|
|
||||||
class NemotronHMoEDecoderLayer(nn.Module):
|
class NemotronHMoEDecoderLayer(NemotronHMLPLikeDecoderLayer):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
config: NemotronHConfig,
|
config: NemotronHConfig,
|
||||||
@@ -361,6 +389,7 @@ class NemotronHMoEDecoderLayer(nn.Module):
|
|||||||
super().__init__()
|
super().__init__()
|
||||||
layer_config = config.get_nemotron_h_config_for_layer(layer_idx)
|
layer_config = config.get_nemotron_h_config_for_layer(layer_idx)
|
||||||
|
|
||||||
|
self.layer_idx = layer_idx
|
||||||
self.mixer = NemotronHMoE(
|
self.mixer = NemotronHMoE(
|
||||||
layer_config,
|
layer_config,
|
||||||
layer_idx=layer_idx,
|
layer_idx=layer_idx,
|
||||||
@@ -369,25 +398,31 @@ class NemotronHMoEDecoderLayer(nn.Module):
|
|||||||
)
|
)
|
||||||
|
|
||||||
self.norm = RMSNorm(config.hidden_size, eps=config.layer_norm_epsilon)
|
self.norm = RMSNorm(config.hidden_size, eps=config.layer_norm_epsilon)
|
||||||
|
self.layer_communicator = make_layer_communicator(self.norm, for_attn=False)
|
||||||
|
|
||||||
def forward(
|
|
||||||
|
class NemotronHAttnLikeDecoderLayer(nn.Module):
|
||||||
|
"""Shared DP-attention input prep for the Mamba / full-attention layers."""
|
||||||
|
|
||||||
|
def _set_prev_layer_is_attn(self, config: NemotronHConfig, layer_idx: int) -> None:
|
||||||
|
self.prev_layer_is_attn = layer_idx > 0 and is_attn_layer(
|
||||||
|
config.hybrid_override_pattern[layer_idx - 1]
|
||||||
|
)
|
||||||
|
|
||||||
|
def _dp_attn_input(
|
||||||
self,
|
self,
|
||||||
*,
|
|
||||||
hidden_states: torch.Tensor,
|
hidden_states: torch.Tensor,
|
||||||
residual: Optional[torch.Tensor],
|
residual: Optional[torch.Tensor],
|
||||||
forward_batch: ForwardBatch,
|
forward_batch: ForwardBatch,
|
||||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
) -> tuple[torch.Tensor, Optional[torch.Tensor]]:
|
||||||
if residual is None:
|
if self.prev_layer_is_attn and residual is not None:
|
||||||
residual = hidden_states
|
hidden_states = attn_tp_all_reduce(hidden_states)
|
||||||
hidden_states = self.norm(hidden_states)
|
return self.layer_communicator.prepare_attn(
|
||||||
else:
|
hidden_states, residual, forward_batch
|
||||||
hidden_states, residual = self.norm(hidden_states, residual)
|
)
|
||||||
|
|
||||||
hidden_states = self.mixer.forward(hidden_states)
|
|
||||||
return hidden_states, residual
|
|
||||||
|
|
||||||
|
|
||||||
class NemotronHMambaDecoderLayer(nn.Module):
|
class NemotronHMambaDecoderLayer(NemotronHAttnLikeDecoderLayer):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
config: NemotronHConfig,
|
config: NemotronHConfig,
|
||||||
@@ -411,15 +446,22 @@ class NemotronHMambaDecoderLayer(nn.Module):
|
|||||||
)
|
)
|
||||||
|
|
||||||
self.norm = RMSNorm(config.hidden_size, eps=config.layer_norm_epsilon)
|
self.norm = RMSNorm(config.hidden_size, eps=config.layer_norm_epsilon)
|
||||||
|
self.layer_communicator = make_layer_communicator(self.norm, for_attn=True)
|
||||||
|
self._set_prev_layer_is_attn(config, layer_idx)
|
||||||
|
|
||||||
def _forward_mamba(
|
def _forward_mamba(
|
||||||
self, hidden_states: torch.Tensor, forward_batch: ForwardBatch
|
self, hidden_states: torch.Tensor, forward_batch: ForwardBatch
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
"""Core Mamba forward logic for the eager path; returns the result."""
|
"""Core Mamba forward logic, called directly or via split op."""
|
||||||
|
original_num_tokens = hidden_states.shape[0]
|
||||||
|
if forward_batch.forward_mode.is_extend():
|
||||||
|
real_num_tokens = get_real_num_tokens(hidden_states, forward_batch)
|
||||||
|
if real_num_tokens < original_num_tokens:
|
||||||
|
hidden_states = hidden_states[:real_num_tokens]
|
||||||
attn_backend = get_attn_backend()
|
attn_backend = get_attn_backend()
|
||||||
assert isinstance(attn_backend, HybridLinearAttnBackend)
|
assert isinstance(attn_backend, HybridLinearAttnBackend)
|
||||||
assert isinstance(attn_backend.linear_attn_backend, Mamba2AttnBackend)
|
assert isinstance(attn_backend.linear_attn_backend, Mamba2AttnBackend)
|
||||||
return attn_backend.linear_attn_backend.forward(
|
output = attn_backend.linear_attn_backend.forward(
|
||||||
mixer=self.mixer,
|
mixer=self.mixer,
|
||||||
layer_id=self.layer_id,
|
layer_id=self.layer_id,
|
||||||
hidden_states=hidden_states,
|
hidden_states=hidden_states,
|
||||||
@@ -427,6 +469,7 @@ class NemotronHMambaDecoderLayer(nn.Module):
|
|||||||
forward_batch=forward_batch,
|
forward_batch=forward_batch,
|
||||||
use_triton_causal_conv=True,
|
use_triton_causal_conv=True,
|
||||||
)
|
)
|
||||||
|
return pad_to_original_num_tokens(output, original_num_tokens)
|
||||||
|
|
||||||
def forward(
|
def forward(
|
||||||
self,
|
self,
|
||||||
@@ -435,6 +478,19 @@ class NemotronHMambaDecoderLayer(nn.Module):
|
|||||||
residual: Optional[torch.Tensor],
|
residual: Optional[torch.Tensor],
|
||||||
forward_batch: ForwardBatch,
|
forward_batch: ForwardBatch,
|
||||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||||
|
if is_dp_attention_enabled():
|
||||||
|
hidden_states, residual = self._dp_attn_input(
|
||||||
|
hidden_states, residual, forward_batch
|
||||||
|
)
|
||||||
|
if (
|
||||||
|
forward_batch.forward_mode.is_idle()
|
||||||
|
or get_real_num_tokens(hidden_states, forward_batch) == 0
|
||||||
|
):
|
||||||
|
return torch.zeros_like(hidden_states), residual
|
||||||
|
|
||||||
|
output = self._forward_mamba(hidden_states, forward_batch)
|
||||||
|
return output, residual
|
||||||
|
|
||||||
if residual is None:
|
if residual is None:
|
||||||
residual = hidden_states
|
residual = hidden_states
|
||||||
hidden_states = self.norm(hidden_states)
|
hidden_states = self.norm(hidden_states)
|
||||||
@@ -465,7 +521,8 @@ class NemotronHAttention(nn.Module):
|
|||||||
) -> None:
|
) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.hidden_size = config.hidden_size
|
self.hidden_size = config.hidden_size
|
||||||
tp_size = get_tensor_model_parallel_world_size()
|
tp_rank = get_attention_tp_rank()
|
||||||
|
tp_size = get_attention_tp_size()
|
||||||
self.total_num_heads = config.num_attention_heads
|
self.total_num_heads = config.num_attention_heads
|
||||||
assert self.total_num_heads % tp_size == 0
|
assert self.total_num_heads % tp_size == 0
|
||||||
self.num_heads = self.total_num_heads // tp_size
|
self.num_heads = self.total_num_heads // tp_size
|
||||||
@@ -494,6 +551,8 @@ class NemotronHAttention(nn.Module):
|
|||||||
self.total_num_kv_heads,
|
self.total_num_kv_heads,
|
||||||
bias=False,
|
bias=False,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
|
tp_rank=tp_rank,
|
||||||
|
tp_size=tp_size,
|
||||||
prefix=f"{prefix}.qkv_proj",
|
prefix=f"{prefix}.qkv_proj",
|
||||||
)
|
)
|
||||||
self.o_proj = RowParallelLinear(
|
self.o_proj = RowParallelLinear(
|
||||||
@@ -501,6 +560,9 @@ class NemotronHAttention(nn.Module):
|
|||||||
config.hidden_size,
|
config.hidden_size,
|
||||||
bias=False,
|
bias=False,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
|
tp_rank=tp_rank,
|
||||||
|
tp_size=tp_size,
|
||||||
|
reduce_results=not is_dp_attention_enabled(),
|
||||||
prefix=f"{prefix}.o_proj",
|
prefix=f"{prefix}.o_proj",
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -518,14 +580,43 @@ class NemotronHAttention(nn.Module):
|
|||||||
def forward(
|
def forward(
|
||||||
self, hidden_states: torch.Tensor, forward_batch: ForwardBatch
|
self, hidden_states: torch.Tensor, forward_batch: ForwardBatch
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
|
if not is_dp_attention_enabled():
|
||||||
|
qkv, _ = self.qkv_proj(hidden_states)
|
||||||
|
q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
|
||||||
|
attn_output = self.attn.forward(q, k, v, forward_batch)
|
||||||
|
output, _ = self.o_proj(attn_output)
|
||||||
|
return output
|
||||||
|
|
||||||
|
padded_shape = hidden_states.shape[0]
|
||||||
|
real_tokens = get_real_num_tokens(hidden_states, forward_batch)
|
||||||
|
has_padding = real_tokens < padded_shape
|
||||||
|
keep_q_padded = (
|
||||||
|
forward_batch.forward_mode.is_decode()
|
||||||
|
or forward_batch.forward_mode.is_target_verify()
|
||||||
|
or forward_batch.forward_mode.is_idle()
|
||||||
|
or forward_batch._original_forward_mode is not None
|
||||||
|
)
|
||||||
|
original_out_cache_loc = forward_batch.out_cache_loc
|
||||||
|
|
||||||
qkv, _ = self.qkv_proj(hidden_states)
|
qkv, _ = self.qkv_proj(hidden_states)
|
||||||
q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
|
q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
|
||||||
attn_output = self.attn.forward(q, k, v, forward_batch)
|
if has_padding and real_tokens > 0:
|
||||||
|
k, v = k[:real_tokens], v[:real_tokens]
|
||||||
|
if original_out_cache_loc is not None:
|
||||||
|
forward_batch.out_cache_loc = original_out_cache_loc[:real_tokens]
|
||||||
|
if not keep_q_padded:
|
||||||
|
q = q[:real_tokens]
|
||||||
|
attn_output = self.attn.forward(
|
||||||
|
q, k, v, forward_batch, save_kv_cache=real_tokens > 0
|
||||||
|
)
|
||||||
|
forward_batch.out_cache_loc = original_out_cache_loc
|
||||||
|
|
||||||
|
attn_output = pad_to_original_num_tokens(attn_output, padded_shape)
|
||||||
output, _ = self.o_proj(attn_output)
|
output, _ = self.o_proj(attn_output)
|
||||||
return output
|
return output
|
||||||
|
|
||||||
|
|
||||||
class NemotronHAttentionDecoderLayer(nn.Module):
|
class NemotronHAttentionDecoderLayer(NemotronHAttnLikeDecoderLayer):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
config: NemotronHConfig,
|
config: NemotronHConfig,
|
||||||
@@ -544,6 +635,8 @@ class NemotronHAttentionDecoderLayer(nn.Module):
|
|||||||
)
|
)
|
||||||
|
|
||||||
self.norm = RMSNorm(config.hidden_size, eps=config.layer_norm_epsilon)
|
self.norm = RMSNorm(config.hidden_size, eps=config.layer_norm_epsilon)
|
||||||
|
self.layer_communicator = make_layer_communicator(self.norm, for_attn=True)
|
||||||
|
self._set_prev_layer_is_attn(config, layer_idx)
|
||||||
|
|
||||||
def forward(
|
def forward(
|
||||||
self,
|
self,
|
||||||
@@ -552,6 +645,15 @@ class NemotronHAttentionDecoderLayer(nn.Module):
|
|||||||
residual: Optional[torch.Tensor],
|
residual: Optional[torch.Tensor],
|
||||||
forward_batch: ForwardBatch,
|
forward_batch: ForwardBatch,
|
||||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||||
|
if is_dp_attention_enabled():
|
||||||
|
hidden_states, residual = self._dp_attn_input(
|
||||||
|
hidden_states, residual, forward_batch
|
||||||
|
)
|
||||||
|
hidden_states = self.mixer.forward(
|
||||||
|
hidden_states=hidden_states, forward_batch=forward_batch
|
||||||
|
)
|
||||||
|
return hidden_states, residual
|
||||||
|
|
||||||
if residual is None:
|
if residual is None:
|
||||||
residual = hidden_states
|
residual = hidden_states
|
||||||
hidden_states = self.norm(hidden_states)
|
hidden_states = self.norm(hidden_states)
|
||||||
@@ -604,6 +706,7 @@ class NemotronHModel(nn.Module):
|
|||||||
self.vocab_size,
|
self.vocab_size,
|
||||||
config.hidden_size,
|
config.hidden_size,
|
||||||
org_num_embeddings=config.vocab_size,
|
org_num_embeddings=config.vocab_size,
|
||||||
|
use_attn_tp_group=is_dp_attention_enabled(),
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
self.embed_tokens = PPMissingLayer()
|
self.embed_tokens = PPMissingLayer()
|
||||||
@@ -732,6 +835,7 @@ class NemotronHForCausalLM(nn.Module):
|
|||||||
else lora_config.lora_vocab_padding_size
|
else lora_config.lora_vocab_padding_size
|
||||||
),
|
),
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
|
use_attn_tp_group=get_global_server_args().enable_dp_lm_head,
|
||||||
prefix=add_prefix("lm_head", prefix),
|
prefix=add_prefix("lm_head", prefix),
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -19,6 +19,13 @@ from torch import nn
|
|||||||
|
|
||||||
from sglang.srt.configs import NemotronHConfig
|
from sglang.srt.configs import NemotronHConfig
|
||||||
from sglang.srt.distributed import get_pp_group
|
from sglang.srt.distributed import get_pp_group
|
||||||
|
from sglang.srt.layers.dp_attention import (
|
||||||
|
attn_tp_all_reduce,
|
||||||
|
get_attention_tp_group,
|
||||||
|
get_attention_tp_rank,
|
||||||
|
get_attention_tp_size,
|
||||||
|
is_dp_attention_enabled,
|
||||||
|
)
|
||||||
from sglang.srt.layers.layernorm import RMSNorm
|
from sglang.srt.layers.layernorm import RMSNorm
|
||||||
from sglang.srt.layers.linear import ColumnParallelLinear
|
from sglang.srt.layers.linear import ColumnParallelLinear
|
||||||
from sglang.srt.layers.logits_processor import LogitsProcessor
|
from sglang.srt.layers.logits_processor import LogitsProcessor
|
||||||
@@ -33,6 +40,7 @@ from sglang.srt.models.nemotron_h import (
|
|||||||
NemotronHForCausalLM,
|
NemotronHForCausalLM,
|
||||||
NemotronHMoEDecoderLayer,
|
NemotronHMoEDecoderLayer,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.models.nemotron_h_utils import is_attn_layer
|
||||||
from sglang.srt.server_args import get_global_server_args
|
from sglang.srt.server_args import get_global_server_args
|
||||||
from sglang.srt.utils import add_prefix
|
from sglang.srt.utils import add_prefix
|
||||||
|
|
||||||
@@ -60,12 +68,14 @@ class NemotronHMTPAttentionDecoderLayer(NemotronHAttentionDecoderLayer):
|
|||||||
self.enorm = RMSNorm(config.hidden_size, eps=config.layer_norm_epsilon)
|
self.enorm = RMSNorm(config.hidden_size, eps=config.layer_norm_epsilon)
|
||||||
self.hnorm = RMSNorm(config.hidden_size, eps=config.layer_norm_epsilon)
|
self.hnorm = RMSNorm(config.hidden_size, eps=config.layer_norm_epsilon)
|
||||||
|
|
||||||
# Fusion layer to combine embeddings with target hidden states
|
_dp_attn = is_dp_attention_enabled()
|
||||||
self.eh_proj = ColumnParallelLinear(
|
self.eh_proj = ColumnParallelLinear(
|
||||||
input_size=config.hidden_size * 2,
|
input_size=config.hidden_size * 2,
|
||||||
output_size=config.hidden_size,
|
output_size=config.hidden_size,
|
||||||
bias=False,
|
bias=False,
|
||||||
gather_output=True,
|
gather_output=not _dp_attn,
|
||||||
|
tp_rank=get_attention_tp_rank() if _dp_attn else None,
|
||||||
|
tp_size=get_attention_tp_size() if _dp_attn else None,
|
||||||
params_dtype=(
|
params_dtype=(
|
||||||
config.dtype if hasattr(config, "dtype") else torch.bfloat16
|
config.dtype if hasattr(config, "dtype") else torch.bfloat16
|
||||||
),
|
),
|
||||||
@@ -95,6 +105,10 @@ class NemotronHMTPAttentionDecoderLayer(NemotronHAttentionDecoderLayer):
|
|||||||
[inputs_embeds_normed, previous_hidden_states_normed], dim=-1
|
[inputs_embeds_normed, previous_hidden_states_normed], dim=-1
|
||||||
)
|
)
|
||||||
hidden_states, _ = self.eh_proj(fused)
|
hidden_states, _ = self.eh_proj(fused)
|
||||||
|
if is_dp_attention_enabled():
|
||||||
|
hidden_states = get_attention_tp_group().all_gather(
|
||||||
|
hidden_states, dim=-1
|
||||||
|
)
|
||||||
|
|
||||||
hidden_states, residual = super().forward(
|
hidden_states, residual = super().forward(
|
||||||
hidden_states=hidden_states,
|
hidden_states=hidden_states,
|
||||||
@@ -130,16 +144,23 @@ class NemotronHMTPMoEDecoderLayer(NemotronHMoEDecoderLayer):
|
|||||||
)
|
)
|
||||||
self.has_start_projections = has_start_projections
|
self.has_start_projections = has_start_projections
|
||||||
self.has_end_norm = has_end_norm
|
self.has_end_norm = has_end_norm
|
||||||
|
_pat = config.mtp_hybrid_override_pattern
|
||||||
|
self.prev_layer_is_attn = layer_idx > 0 and is_attn_layer(
|
||||||
|
_pat[(layer_idx - 1) % len(_pat)]
|
||||||
|
)
|
||||||
|
|
||||||
if has_start_projections:
|
if has_start_projections:
|
||||||
self.enorm = RMSNorm(config.hidden_size, eps=config.layer_norm_epsilon)
|
self.enorm = RMSNorm(config.hidden_size, eps=config.layer_norm_epsilon)
|
||||||
self.hnorm = RMSNorm(config.hidden_size, eps=config.layer_norm_epsilon)
|
self.hnorm = RMSNorm(config.hidden_size, eps=config.layer_norm_epsilon)
|
||||||
|
|
||||||
|
_dp_attn = is_dp_attention_enabled()
|
||||||
self.eh_proj = ColumnParallelLinear(
|
self.eh_proj = ColumnParallelLinear(
|
||||||
input_size=config.hidden_size * 2,
|
input_size=config.hidden_size * 2,
|
||||||
output_size=config.hidden_size,
|
output_size=config.hidden_size,
|
||||||
bias=False,
|
bias=False,
|
||||||
gather_output=True,
|
gather_output=not _dp_attn,
|
||||||
|
tp_rank=get_attention_tp_rank() if _dp_attn else None,
|
||||||
|
tp_size=get_attention_tp_size() if _dp_attn else None,
|
||||||
params_dtype=(
|
params_dtype=(
|
||||||
config.dtype if hasattr(config, "dtype") else torch.bfloat16
|
config.dtype if hasattr(config, "dtype") else torch.bfloat16
|
||||||
),
|
),
|
||||||
@@ -169,6 +190,17 @@ class NemotronHMTPMoEDecoderLayer(NemotronHMoEDecoderLayer):
|
|||||||
[inputs_embeds_normed, previous_hidden_states_normed], dim=-1
|
[inputs_embeds_normed, previous_hidden_states_normed], dim=-1
|
||||||
)
|
)
|
||||||
hidden_states, _ = self.eh_proj(fused)
|
hidden_states, _ = self.eh_proj(fused)
|
||||||
|
if is_dp_attention_enabled():
|
||||||
|
hidden_states = get_attention_tp_group().all_gather(
|
||||||
|
hidden_states, dim=-1
|
||||||
|
)
|
||||||
|
|
||||||
|
if (
|
||||||
|
is_dp_attention_enabled()
|
||||||
|
and self.prev_layer_is_attn
|
||||||
|
and residual is not None
|
||||||
|
):
|
||||||
|
hidden_states = attn_tp_all_reduce(hidden_states)
|
||||||
|
|
||||||
hidden_states, residual = super().forward(
|
hidden_states, residual = super().forward(
|
||||||
hidden_states=hidden_states,
|
hidden_states=hidden_states,
|
||||||
@@ -212,6 +244,8 @@ class NemotronHMultiTokenPredictor(nn.Module):
|
|||||||
self.embed_tokens = VocabParallelEmbedding(
|
self.embed_tokens = VocabParallelEmbedding(
|
||||||
self.vocab_size,
|
self.vocab_size,
|
||||||
config.hidden_size,
|
config.hidden_size,
|
||||||
|
org_num_embeddings=config.vocab_size,
|
||||||
|
use_attn_tp_group=is_dp_attention_enabled(),
|
||||||
)
|
)
|
||||||
|
|
||||||
# Build flat list of layers
|
# Build flat list of layers
|
||||||
|
|||||||
@@ -0,0 +1,67 @@
|
|||||||
|
"""DP-attention helpers for the Nemotron-H model."""
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch import nn
|
||||||
|
|
||||||
|
from sglang.srt.configs.nemotron_h import ATTENTION, MAMBA
|
||||||
|
from sglang.srt.layers.communicator import (
|
||||||
|
LayerCommunicator,
|
||||||
|
LayerScatterModes,
|
||||||
|
ScatterMode,
|
||||||
|
)
|
||||||
|
from sglang.srt.layers.layernorm import RMSNorm
|
||||||
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
|
|
||||||
|
ATTN_LAYERS = (MAMBA, ATTENTION)
|
||||||
|
|
||||||
|
|
||||||
|
def is_attn_layer(layer_type: str) -> bool:
|
||||||
|
return layer_type in ATTN_LAYERS
|
||||||
|
|
||||||
|
|
||||||
|
def get_real_num_tokens(
|
||||||
|
hidden_states: torch.Tensor, forward_batch: ForwardBatch
|
||||||
|
) -> int:
|
||||||
|
"""Number of real (non DP-padding) rows in ``hidden_states``."""
|
||||||
|
real_tokens = hidden_states.shape[0]
|
||||||
|
num_token_non_padded_cpu = getattr(forward_batch, "num_token_non_padded_cpu", None)
|
||||||
|
if num_token_non_padded_cpu is not None:
|
||||||
|
real_tokens = min(real_tokens, int(num_token_non_padded_cpu))
|
||||||
|
if (
|
||||||
|
forward_batch.forward_mode.is_extend()
|
||||||
|
and not forward_batch.forward_mode.is_mixed()
|
||||||
|
and forward_batch.extend_seq_lens_cpu is not None
|
||||||
|
):
|
||||||
|
real_tokens = min(real_tokens, int(sum(forward_batch.extend_seq_lens_cpu)))
|
||||||
|
return real_tokens
|
||||||
|
|
||||||
|
|
||||||
|
def pad_to_original_num_tokens(
|
||||||
|
output: torch.Tensor, original_num_tokens: int
|
||||||
|
) -> torch.Tensor:
|
||||||
|
if output.shape[0] == original_num_tokens:
|
||||||
|
return output
|
||||||
|
padded = output.new_empty((original_num_tokens, *output.shape[1:]))
|
||||||
|
padded[: output.shape[0]] = output
|
||||||
|
return padded
|
||||||
|
|
||||||
|
|
||||||
|
def _build_layer_scatter_modes() -> LayerScatterModes:
|
||||||
|
return LayerScatterModes(
|
||||||
|
layer_input_mode=ScatterMode.TP_ATTN_FULL,
|
||||||
|
attn_mode=ScatterMode.TP_ATTN_FULL,
|
||||||
|
mlp_mode=ScatterMode.FULL,
|
||||||
|
middle_residual_mode=ScatterMode.TP_ATTN_FULL,
|
||||||
|
layer_output_mode=ScatterMode.TP_ATTN_FULL,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def make_layer_communicator(
|
||||||
|
layer_norm: RMSNorm, *, for_attn: bool
|
||||||
|
) -> LayerCommunicator:
|
||||||
|
return LayerCommunicator(
|
||||||
|
layer_scatter_modes=_build_layer_scatter_modes(),
|
||||||
|
input_layernorm=layer_norm if for_attn else nn.Identity(),
|
||||||
|
post_attention_layernorm=nn.Identity() if for_attn else layer_norm,
|
||||||
|
force_layernorm_before_dp_gather=True,
|
||||||
|
)
|
||||||
Reference in New Issue
Block a user