[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
parent e039ff382c
commit 0296f1cdad
7 changed files with 21 additions and 15 deletions

View File

@@ -62,7 +62,7 @@ logger = logging.getLogger(__name__)
class SiluAndMul(CustomOp):
def __init__(self, *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
def forward_native(self, x: torch.Tensor) -> torch.Tensor:

View File

@@ -824,7 +824,7 @@ class LogitsProcessor(nn.Module):
None, # bias
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
logits = torch.matmul(
hidden_states.bfloat16(), lm_head.weight.T.bfloat16()

View File

@@ -127,7 +127,7 @@ class RotaryEmbedding(CustomOp):
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._apply_rotary_emb_wrapped = torch.compile(dynamic=True)(
self._apply_rotary_emb_wrapped
@@ -140,7 +140,7 @@ class RotaryEmbedding(CustomOp):
# create the cache on GPU for faster initialization. This may cause
# a slight numerical difference between the HF implementation and ours.
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 / (
base
@@ -151,7 +151,7 @@ class RotaryEmbedding(CustomOp):
/ 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()
return inv_freq

View File

@@ -102,7 +102,7 @@ class Sampler(nn.Module):
if return_logprob and SGLANG_RETURN_ORIGINAL_LOGPROB:
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.bfloat16().div(sampling_info.temperatures).bfloat16()
)
@@ -156,7 +156,7 @@ class Sampler(nn.Module):
)
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
del logprobs_via_logsoftmax_kernel
# clamp to avoid -inf

View File

@@ -90,7 +90,7 @@ class Qwen2MLP(nn.Module):
self.act_fn = SiluAndMul()
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()
gate_up, _ = self.gate_up_proj(x)
@@ -281,7 +281,7 @@ class Qwen2Model(nn.Module):
prefix=add_prefix("embed_tokens", prefix),
params_dtype=(
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
),
)
@@ -311,7 +311,7 @@ class Qwen2Model(nn.Module):
override_orig_dtype=torch.float32,
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 {}
)
self.norm = RMSNorm(

View File

@@ -94,7 +94,7 @@ class Qwen3Attention(nn.Module):
weight_dtype=torch.float32,
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 {}
)
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,
forward_batch: ForwardBatch,
) -> 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()
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.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)
k = k.to(torch.bfloat16)
@@ -229,7 +229,7 @@ class Qwen3DecoderLayer(nn.Module):
override_orig_dtype=torch.float32,
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 {}
)
self.input_layernorm = RMSNorm(

View File

@@ -152,6 +152,8 @@ NSA_CHOICES = [
RADIX_EVICTION_POLICY_CHOICES = ["lru", "lfu"]
RL_ON_POLICY_TARGET_CHOICES = ["fsdp"]
MOE_RUNNER_BACKEND_CHOICES = [
"auto",
"deep_gemm",
@@ -204,6 +206,10 @@ def add_radix_eviction_policy_choices(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):
MAMBA_SSM_DTYPE_CHOICES.extend(choices)
@@ -3429,7 +3435,7 @@ class ServerArgs:
"--rl-on-policy-target",
type=str,
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.",
)