diff --git a/python/sglang/srt/layers/activation.py b/python/sglang/srt/layers/activation.py index 6340c4701..6db49012d 100644 --- a/python/sglang/srt/layers/activation.py +++ b/python/sglang/srt/layers/activation.py @@ -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: diff --git a/python/sglang/srt/layers/logits_processor.py b/python/sglang/srt/layers/logits_processor.py index d9eab9d25..e2c7d2ab6 100644 --- a/python/sglang/srt/layers/logits_processor.py +++ b/python/sglang/srt/layers/logits_processor.py @@ -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() diff --git a/python/sglang/srt/layers/rotary_embedding.py b/python/sglang/srt/layers/rotary_embedding.py index 5d0469991..a2e8e60b2 100644 --- a/python/sglang/srt/layers/rotary_embedding.py +++ b/python/sglang/srt/layers/rotary_embedding.py @@ -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 diff --git a/python/sglang/srt/layers/sampler.py b/python/sglang/srt/layers/sampler.py index 584e6c8c9..59a0f3bb9 100644 --- a/python/sglang/srt/layers/sampler.py +++ b/python/sglang/srt/layers/sampler.py @@ -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 diff --git a/python/sglang/srt/models/qwen2.py b/python/sglang/srt/models/qwen2.py index 75ceff821..a7dbadec6 100644 --- a/python/sglang/srt/models/qwen2.py +++ b/python/sglang/srt/models/qwen2.py @@ -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( diff --git a/python/sglang/srt/models/qwen3.py b/python/sglang/srt/models/qwen3.py index b031d6e03..9a9ac4da8 100644 --- a/python/sglang/srt/models/qwen3.py +++ b/python/sglang/srt/models/qwen3.py @@ -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( diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index d025c1b73..6d4e77c61 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -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.", )