[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:
@@ -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:
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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.",
|
||||
)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user