[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:
co-authored by
github-actions[bot]
Stefan He
parent
e039ff382c
commit
0296f1cdad
@@ -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:
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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(
|
||||||
|
|||||||
@@ -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(
|
||||||
|
|||||||
@@ -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.",
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user