[Auto Sync] Update activation.py, logits_processor.py, rota... (20251107) (#12853)

Co-authored-by: github-actions[bot] <github-actions[bot]@users.noreply.github.com>
Co-authored-by: Stefan He <hebiaobuaa@gmail.com>
This commit is contained in:
Lianmin Zheng
2025-11-07 22:07:51 -08:00
committed by GitHub
co-authored by github-actions[bot] Stefan He
parent e039ff382c
commit 0296f1cdad
7 changed files with 21 additions and 15 deletions
+1 -1
View File
@@ -62,7 +62,7 @@ logger = logging.getLogger(__name__)
class SiluAndMul(CustomOp): class SiluAndMul(CustomOp):
def __init__(self, *args, **kwargs): def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs) super().__init__(*args, **kwargs)
if get_global_server_args().rl_on_policy_target == "fsdp": if get_global_server_args().rl_on_policy_target is not None:
self._forward_method = self.forward_native self._forward_method = self.forward_native
def forward_native(self, x: torch.Tensor) -> torch.Tensor: def forward_native(self, x: torch.Tensor) -> torch.Tensor:
+1 -1
View File
@@ -824,7 +824,7 @@ class LogitsProcessor(nn.Module):
None, # bias None, # bias
True, # is_vnni True, # is_vnni
) )
elif get_global_server_args().rl_on_policy_target == "fsdp": elif get_global_server_args().rl_on_policy_target is not None:
# Due to tie-weight, we may not be able to change lm_head's weight dtype # Due to tie-weight, we may not be able to change lm_head's weight dtype
logits = torch.matmul( logits = torch.matmul(
hidden_states.bfloat16(), lm_head.weight.T.bfloat16() hidden_states.bfloat16(), lm_head.weight.T.bfloat16()
+3 -3
View File
@@ -127,7 +127,7 @@ class RotaryEmbedding(CustomOp):
self._apply_rotary_emb_wrapped = _apply_rotary_emb self._apply_rotary_emb_wrapped = _apply_rotary_emb
if get_global_server_args().rl_on_policy_target == "fsdp": if get_global_server_args().rl_on_policy_target is not None:
self._forward_method = self.forward_native self._forward_method = self.forward_native
self._apply_rotary_emb_wrapped = torch.compile(dynamic=True)( self._apply_rotary_emb_wrapped = torch.compile(dynamic=True)(
self._apply_rotary_emb_wrapped self._apply_rotary_emb_wrapped
@@ -140,7 +140,7 @@ class RotaryEmbedding(CustomOp):
# create the cache on GPU for faster initialization. This may cause # create the cache on GPU for faster initialization. This may cause
# a slight numerical difference between the HF implementation and ours. # a slight numerical difference between the HF implementation and ours.
init_device = ( init_device = (
"cpu" if get_global_server_args().rl_on_policy_target == "fsdp" else None "cpu" if get_global_server_args().rl_on_policy_target is not None else None
) )
inv_freq = 1.0 / ( inv_freq = 1.0 / (
base base
@@ -151,7 +151,7 @@ class RotaryEmbedding(CustomOp):
/ self.rotary_dim / self.rotary_dim
) )
) )
if get_global_server_args().rl_on_policy_target == "fsdp": if get_global_server_args().rl_on_policy_target is not None:
inv_freq = inv_freq.cuda() inv_freq = inv_freq.cuda()
return inv_freq return inv_freq
+2 -2
View File
@@ -102,7 +102,7 @@ class Sampler(nn.Module):
if return_logprob and SGLANG_RETURN_ORIGINAL_LOGPROB: if return_logprob and SGLANG_RETURN_ORIGINAL_LOGPROB:
probs_without_temp_scaling = torch.softmax(logits, dim=-1) probs_without_temp_scaling = torch.softmax(logits, dim=-1)
if get_global_server_args().rl_on_policy_target == "fsdp": if get_global_server_args().rl_on_policy_target is not None:
logits_div_temperature = ( logits_div_temperature = (
logits.bfloat16().div(sampling_info.temperatures).bfloat16() logits.bfloat16().div(sampling_info.temperatures).bfloat16()
) )
@@ -156,7 +156,7 @@ class Sampler(nn.Module):
) )
if return_logprob: if return_logprob:
if get_global_server_args().rl_on_policy_target == "fsdp": if get_global_server_args().rl_on_policy_target is not None:
logprobs = logprobs_via_logsoftmax_kernel logprobs = logprobs_via_logsoftmax_kernel
del logprobs_via_logsoftmax_kernel del logprobs_via_logsoftmax_kernel
# clamp to avoid -inf # clamp to avoid -inf
+3 -3
View File
@@ -90,7 +90,7 @@ class Qwen2MLP(nn.Module):
self.act_fn = SiluAndMul() self.act_fn = SiluAndMul()
def forward(self, x): def forward(self, x):
if get_global_server_args().rl_on_policy_target == "fsdp": if get_global_server_args().rl_on_policy_target is not None:
x = x.bfloat16() x = x.bfloat16()
gate_up, _ = self.gate_up_proj(x) gate_up, _ = self.gate_up_proj(x)
@@ -281,7 +281,7 @@ class Qwen2Model(nn.Module):
prefix=add_prefix("embed_tokens", prefix), prefix=add_prefix("embed_tokens", prefix),
params_dtype=( params_dtype=(
torch.float32 torch.float32
if get_global_server_args().rl_on_policy_target == "fsdp" if get_global_server_args().rl_on_policy_target is not None
else None else None
), ),
) )
@@ -311,7 +311,7 @@ class Qwen2Model(nn.Module):
override_orig_dtype=torch.float32, override_orig_dtype=torch.float32,
fp32_residual=True, fp32_residual=True,
) )
if get_global_server_args().rl_on_policy_target == "fsdp" if get_global_server_args().rl_on_policy_target is not None
else {} else {}
) )
self.norm = RMSNorm( self.norm = RMSNorm(
+4 -4
View File
@@ -94,7 +94,7 @@ class Qwen3Attention(nn.Module):
weight_dtype=torch.float32, weight_dtype=torch.float32,
cast_x_before_out_mul=True, cast_x_before_out_mul=True,
) )
if get_global_server_args().rl_on_policy_target == "fsdp" if get_global_server_args().rl_on_policy_target is not None
else {} else {}
) )
self.q_norm = RMSNorm(self.head_dim, eps=rms_norm_eps, **norm_kwargs) self.q_norm = RMSNorm(self.head_dim, eps=rms_norm_eps, **norm_kwargs)
@@ -167,7 +167,7 @@ class Qwen3Attention(nn.Module):
hidden_states: torch.Tensor, hidden_states: torch.Tensor,
forward_batch: ForwardBatch, forward_batch: ForwardBatch,
) -> torch.Tensor: ) -> torch.Tensor:
if get_global_server_args().rl_on_policy_target == "fsdp": if get_global_server_args().rl_on_policy_target is not None:
hidden_states = hidden_states.bfloat16() hidden_states = hidden_states.bfloat16()
qkv, _ = self.qkv_proj(hidden_states) qkv, _ = self.qkv_proj(hidden_states)
@@ -175,7 +175,7 @@ class Qwen3Attention(nn.Module):
q, k = self._apply_qk_norm(q, k) q, k = self._apply_qk_norm(q, k)
q, k = self.rotary_emb(positions, q, k) q, k = self.rotary_emb(positions, q, k)
if get_global_server_args().rl_on_policy_target == "fsdp": if get_global_server_args().rl_on_policy_target is not None:
q = q.to(torch.bfloat16) q = q.to(torch.bfloat16)
k = k.to(torch.bfloat16) k = k.to(torch.bfloat16)
@@ -229,7 +229,7 @@ class Qwen3DecoderLayer(nn.Module):
override_orig_dtype=torch.float32, override_orig_dtype=torch.float32,
fp32_residual=True, fp32_residual=True,
) )
if get_global_server_args().rl_on_policy_target == "fsdp" if get_global_server_args().rl_on_policy_target is not None
else {} else {}
) )
self.input_layernorm = RMSNorm( self.input_layernorm = RMSNorm(
+7 -1
View File
@@ -152,6 +152,8 @@ NSA_CHOICES = [
RADIX_EVICTION_POLICY_CHOICES = ["lru", "lfu"] RADIX_EVICTION_POLICY_CHOICES = ["lru", "lfu"]
RL_ON_POLICY_TARGET_CHOICES = ["fsdp"]
MOE_RUNNER_BACKEND_CHOICES = [ MOE_RUNNER_BACKEND_CHOICES = [
"auto", "auto",
"deep_gemm", "deep_gemm",
@@ -204,6 +206,10 @@ def add_radix_eviction_policy_choices(choices):
RADIX_EVICTION_POLICY_CHOICES.extend(choices) RADIX_EVICTION_POLICY_CHOICES.extend(choices)
def add_rl_on_policy_target_choices(choices):
RL_ON_POLICY_TARGET_CHOICES.extend(choices)
def add_mamba_ssm_dtype_choices(choices): def add_mamba_ssm_dtype_choices(choices):
MAMBA_SSM_DTYPE_CHOICES.extend(choices) MAMBA_SSM_DTYPE_CHOICES.extend(choices)
@@ -3429,7 +3435,7 @@ class ServerArgs:
"--rl-on-policy-target", "--rl-on-policy-target",
type=str, type=str,
default=ServerArgs.rl_on_policy_target, default=ServerArgs.rl_on_policy_target,
choices=["fsdp"], choices=RL_ON_POLICY_TARGET_CHOICES,
help="The training system that SGLang needs to match for true on-policy.", help="The training system that SGLang needs to match for true on-policy.",
) )