Skip to content
Closed
Changes from 1 commit
Commits
Show all changes
22 commits
Select commit Hold shift + click to select a range
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
35 changes: 34 additions & 1 deletion python/sglang/srt/models/qwen3_5.py
Original file line number Diff line number Diff line change
Expand Up @@ -76,6 +76,7 @@
from sglang.srt.models.qwen3_vl import Qwen3VLForConditionalGeneration

# Utils
from sglang.srt.layers.utils import get_layer_id
from sglang.srt.utils import add_prefix, is_cuda, is_npu, make_layers, set_weight_attrs
from sglang.srt.utils.hf_transformers_utils import get_processor

Expand Down Expand Up @@ -672,6 +673,8 @@ def __init__(
org_num_embeddings=config.vocab_size,
enable_tp=not is_dp_attention_enabled(),
)
else:
self.embed_tokens = PPMissingLayer()

# Decoder layers
def get_layer(idx: int, prefix: str):
Expand All @@ -689,9 +692,11 @@ def get_layer(idx: int, prefix: str):
alt_stream=alt_stream,
)

self.layers = make_layers(
self.layers, self.start_layer, self.end_layer= make_layers(
config.num_hidden_layers,
get_layer,
pp_rank=self.pp_group.rank_in_group,
pp_size=self.pp_group.world_size,
prefix=f"{prefix}.layers",
)

Expand Down Expand Up @@ -789,6 +794,13 @@ def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
name = name.replace(r"model.language_model.", r"model.")
if ".self_attn." in name:
name = name.replace(".self_attn", "")
layer_id = get_layer_id(name)
if (
layer_id is not None
and hasattr(self, "start_layer")
and (layer_id < self.start_layer or layer_id >= self.end_layer)
):
continue
Comment on lines +807 to +813

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

medium

This logic to skip loading weights for layers not on the current pipeline parallel rank is duplicated in three other load_weights methods in this file (at lines 925, 1091, and 1241). To improve maintainability and reduce code duplication, consider extracting this logic into a single helper function.


for param_name, weight_name, shard_id in stacked_params_mapping:
if weight_name not in name:
Expand Down Expand Up @@ -910,6 +922,13 @@ def load_fused_expert_weights(
name = name.replace(r"model.language_model.", r"model.")
if ".self_attn." in name:
name = name.replace(".self_attn", "")
layer_id = get_layer_id(name)
if (
layer_id is not None
and hasattr(self, "start_layer")
and (layer_id < self.start_layer or layer_id >= self.end_layer)
):
continue

for param_name, weight_name, shard_id in stacked_params_mapping:
if "experts.gate_up_proj" in name or "experts.down_proj" in name:
Expand Down Expand Up @@ -1069,6 +1088,13 @@ def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
name = name.replace(r"model.language_model.", r"model.")
if ".self_attn." in name:
name = name.replace(".self_attn", "")
layer_id = get_layer_id(name)
if (
layer_id is not None
and hasattr(self.model, "start_layer")
and (layer_id < self.model.start_layer or layer_id >= self.model.end_layer)
):
continue

for param_name, weight_name, shard_id in stacked_params_mapping:
if weight_name not in name:
Expand Down Expand Up @@ -1212,6 +1238,13 @@ def load_fused_expert_weights(
name = name.replace(r"model.language_model.", r"model.")
if ".self_attn." in name:
name = name.replace(".self_attn", "")
layer_id = get_layer_id(name)
if (
layer_id is not None
and hasattr(self.model, "start_layer")
and (layer_id < self.model.start_layer or layer_id >= self.model.end_layer)
):
continue

for param_name, weight_name, shard_id in stacked_params_mapping:
if name.endswith("experts.gate_up_proj") or name.endswith(
Expand Down
Loading