Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
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
26 changes: 26 additions & 0 deletions python/sglang/srt/model_loader/auto_loader.py
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,7 @@
from torch.nn import Parameter

from sglang.srt.layers.utils.common import get_layer_id
from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.models.utils import AutoWeightsLoader, WeightsMapper

__all__ = [
Expand All @@ -48,6 +49,7 @@
"STANDARD_GATE_UP_MAPPING",
"STANDARD_STACKED_MAPPING",
"LLAMA_STACKED_MAPPING",
"load_with_stacked_dispatch",
"filter_pp_weights",
"register_weight_remap",
"get_weight_remap",
Expand Down Expand Up @@ -148,6 +150,30 @@ def try_load(
)


def load_with_stacked_dispatch(
module: nn.Module,
weights: Iterable[tuple[str, torch.Tensor]],
mapping: StackedParamsDispatch,
*,
ignore_unexpected_suffixes: tuple[str, ...] = (".bias", ".kv_scale"),
) -> set[str]:
"""Load submodule weights via stacked dispatch, then direct param loaders."""
loaded: set[str] = set()
params_dict = dict(module.named_parameters())
for name, tensor in weights:
target = mapping.try_load(name, tensor, params_dict)
if target is not None:
loaded.add(target)
continue
if name in params_dict:
wl = getattr(params_dict[name], "weight_loader", default_weight_loader)
wl(params_dict[name], tensor)
loaded.add(name)
elif not any(name.endswith(suffix) for suffix in ignore_unexpected_suffixes):
pass
return loaded


# ---------------------------------------------------------------------------
# Pipeline Parallel Weight Filter
# ---------------------------------------------------------------------------
Expand Down
44 changes: 44 additions & 0 deletions python/sglang/srt/models/gemma2.py
Original file line number Diff line number Diff line change
Expand Up @@ -92,6 +92,14 @@ def forward(self, x: torch.Tensor) -> torch.Tensor:
x, _ = self.down_proj(x)
return x

def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]) -> set[str]:
from sglang.srt.model_loader.auto_loader import (
STANDARD_GATE_UP_MAPPING,
load_with_stacked_dispatch,
)

return load_with_stacked_dispatch(self, weights, STANDARD_GATE_UP_MAPPING)


class Gemma2Attention(nn.Module):
def __init__(
Expand Down Expand Up @@ -200,6 +208,14 @@ def forward(
output, _ = self.o_proj(attn_output)
return output

def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]) -> set[str]:
from sglang.srt.model_loader.auto_loader import (
STANDARD_QKV_MAPPING,
load_with_stacked_dispatch,
)

return load_with_stacked_dispatch(self, weights, STANDARD_QKV_MAPPING)


class Gemma2DecoderLayer(nn.Module):
def __init__(
Expand Down Expand Up @@ -458,6 +474,13 @@ def get_attention_sliding_window_size(self):
return get_attention_sliding_window_size(self.config)

def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
from sglang.srt.environ import envs

if envs.SGLANG_ENABLE_WEIGHT_LOADER_V2.get():
return self._load_weights_v2(weights)
return self._legacy_load_weights(weights)

def _legacy_load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
stacked_params_mapping = [
# (param_name, shard_name, shard_id)
("qkv_proj", "q_proj", "q"),
Expand Down Expand Up @@ -498,5 +521,26 @@ def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
weight_loader(param, loaded_weight)
loaded_params.add(name)

def _load_weights_v2(self, weights: Iterable[Tuple[str, torch.Tensor]]) -> set[str]:
from sglang.srt.model_loader.auto_loader import AutoWeightsLoader

params_dict = dict(self.named_parameters())

def _prepare(
src: Iterable[Tuple[str, torch.Tensor]],
) -> Iterable[Tuple[str, torch.Tensor]]:
for name, loaded_weight in src:
remapped = maybe_remap_kv_scale_name(name, params_dict)
if remapped is None:
continue
yield remapped, loaded_weight

loader = AutoWeightsLoader(
self,
skip_prefixes=["lm_head."],
ignore_unexpected_suffixes=[".bias", ".kv_scale"],
)
return loader.load_weights(_prepare(weights))


EntryClass = Gemma2ForCausalLM
58 changes: 58 additions & 0 deletions python/sglang/srt/models/glm4.py
Original file line number Diff line number Diff line change
Expand Up @@ -97,6 +97,14 @@ def forward(
x, _ = self.down_proj(x)
return x

def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]) -> set[str]:
from sglang.srt.model_loader.auto_loader import (
STANDARD_GATE_UP_MAPPING,
load_with_stacked_dispatch,
)

return load_with_stacked_dispatch(self, weights, STANDARD_GATE_UP_MAPPING)


class Glm4Attention(nn.Module):
def __init__(
Expand Down Expand Up @@ -192,6 +200,14 @@ def forward(
output, _ = self.o_proj(attn_output)
return output

def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]) -> set[str]:
from sglang.srt.model_loader.auto_loader import (
STANDARD_QKV_MAPPING,
load_with_stacked_dispatch,
)

return load_with_stacked_dispatch(self, weights, STANDARD_QKV_MAPPING)


class Glm4DecoderLayer(nn.Module):
"""A single transformer layer.
Expand Down Expand Up @@ -549,6 +565,13 @@ def end_layer(self):
return self.model.end_layer

def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
from sglang.srt.environ import envs

if envs.SGLANG_ENABLE_WEIGHT_LOADER_V2.get():
return self._load_weights_v2(weights)
return self._legacy_load_weights(weights)

def _legacy_load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
stacked_params_mapping = [
# (param_name, shard_name, shard_id)
(".qkv_proj", ".q_proj", "q"),
Expand Down Expand Up @@ -611,6 +634,41 @@ def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
else:
logger.warning(f"Parameter {name} not found in params_dict")

def _load_weights_v2(self, weights: Iterable[Tuple[str, torch.Tensor]]) -> set[str]:
from sglang.srt.model_loader.auto_loader import (
AutoWeightsLoader,
filter_pp_weights,
)

if hasattr(self.model, "start_layer"):
weights = filter_pp_weights(
weights, self.model.start_layer, self.model.end_layer
)

skip_prefixes = []
if self.config.tie_word_embeddings:
skip_prefixes.append("lm_head.")

loader = AutoWeightsLoader(
self,
skip_prefixes=skip_prefixes,
skip_substrs=["projector"],
ignore_unexpected_suffixes=[".bias", ".kv_scale"],
)
loaded = loader.load_weights(weights)

if self.config.tie_word_embeddings:
params_dict = dict(self.named_parameters())
if "lm_head.weight" in params_dict:
embed = dict(self.model.named_parameters()).get("embed_tokens.weight")
if embed is not None:
lm_head = params_dict["lm_head.weight"]
wl = getattr(lm_head, "weight_loader", default_weight_loader)
wl(lm_head, embed.data)
loaded.add("lm_head.weight")

return loaded

def get_embed_and_head(self):
return self.model.embed_tokens.weight, self.lm_head.weight

Expand Down
50 changes: 50 additions & 0 deletions python/sglang/srt/models/granite.py
Original file line number Diff line number Diff line change
Expand Up @@ -119,6 +119,14 @@ def forward(self, x):
x, _ = self.down_proj(x)
return x

def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]) -> set[str]:
from sglang.srt.model_loader.auto_loader import (
STANDARD_GATE_UP_MAPPING,
load_with_stacked_dispatch,
)

return load_with_stacked_dispatch(self, weights, STANDARD_GATE_UP_MAPPING)


class GraniteAttention(nn.Module):
def __init__(
Expand Down Expand Up @@ -220,6 +228,14 @@ def forward(
output, _ = self.o_proj(attn_output)
return output

def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]) -> set[str]:
from sglang.srt.model_loader.auto_loader import (
STANDARD_QKV_MAPPING,
load_with_stacked_dispatch,
)

return load_with_stacked_dispatch(self, weights, STANDARD_QKV_MAPPING)


class GraniteDecoderLayer(nn.Module):
def __init__(
Expand Down Expand Up @@ -420,6 +436,13 @@ def get_num_params(self):
return len(params_dict)

def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
from sglang.srt.environ import envs

if envs.SGLANG_ENABLE_WEIGHT_LOADER_V2.get():
return self._load_weights_v2(weights)
return self._legacy_load_weights(weights)

def _legacy_load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
stacked_params_mapping = [
# (param_name, shard_name, shard_id)
(".qkv_proj", ".q_proj", "q"),
Expand Down Expand Up @@ -471,6 +494,33 @@ def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
weight_loader = getattr(param, "weight_loader", default_weight_loader)
weight_loader(param, loaded_weight)

def _load_weights_v2(self, weights: Iterable[Tuple[str, torch.Tensor]]) -> set[str]:
from sglang.srt.model_loader.auto_loader import AutoWeightsLoader

skip_prefixes = []
if self.config.tie_word_embeddings:
skip_prefixes.append("lm_head.")

loader = AutoWeightsLoader(
self,
skip_prefixes=skip_prefixes,
skip_substrs=["projector", "model.vision_tower"],
ignore_unexpected_suffixes=[".bias", ".kv_scale"],
)
loaded = loader.load_weights(weights)

if self.config.tie_word_embeddings:
params_dict = dict(self.named_parameters())
if "lm_head.weight" in params_dict:
embed = dict(self.model.named_parameters()).get("embed_tokens.weight")
if embed is not None:
lm_head = params_dict["lm_head.weight"]
wl = getattr(lm_head, "weight_loader", default_weight_loader)
wl(lm_head, embed.data)
loaded.add("lm_head.weight")

return loaded

def get_weights_by_name(
self, name: str, truncate_size: int = 100, tp_size: int = 1
) -> Optional[torch.Tensor]:
Expand Down
61 changes: 61 additions & 0 deletions python/sglang/srt/models/internlm2.py
Original file line number Diff line number Diff line change
Expand Up @@ -80,10 +80,25 @@ def forward(self, x):
x, _ = self.w2(x)
return x

def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]) -> set[str]:
from sglang.srt.model_loader.auto_loader import (
StackedParamsDispatch,
load_with_stacked_dispatch,
)

mapping = StackedParamsDispatch(
mappings=(
("gate_up_proj", "w1", 0),
("gate_up_proj", "w3", 1),
)
)
return load_with_stacked_dispatch(self, weights, mapping)


class InternLM2Attention(nn.Module):
def __init__(
self,
config: PretrainedConfig,
hidden_size: int,
num_heads: int,
num_kv_heads: int,
Expand All @@ -95,6 +110,7 @@ def __init__(
prefix: str = "",
) -> None:
super().__init__()
self.config = config
self.hidden_size = hidden_size
tp_size = get_parallel().tp_size
self.total_num_heads = num_heads
Expand Down Expand Up @@ -164,6 +180,34 @@ def forward(
output, _ = self.wo(attn_output)
return output

def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]) -> set[str]:
loaded: set[str] = set()
params_dict = dict(self.named_parameters())
config = self.config
for name, tensor in weights:
if "wqkv" in name and name in params_dict:
param = params_dict[name]
kv_groups = config.num_attention_heads // config.num_key_value_heads
head_dim = config.hidden_size // config.num_attention_heads
loaded_weight = tensor.view(
-1, 2 + kv_groups, head_dim, tensor.shape[-1]
)
wq, wk, wv = torch.split(loaded_weight, [kv_groups, 1, 1], dim=1)
wq = wq.reshape(-1, wq.shape[-1])
wk = wk.reshape(-1, wk.shape[-1])
wv = wv.reshape(-1, wv.shape[-1])
weight_loader = param.weight_loader
weight_loader(param, wq, "q")
weight_loader(param, wk, "k")
weight_loader(param, wv, "v")
loaded.add(name)
continue
if name in params_dict:
wl = getattr(params_dict[name], "weight_loader", default_weight_loader)
wl(params_dict[name], tensor)
loaded.add(name)
return loaded


class InternLMDecoderLayer(nn.Module):
def __init__(
Expand All @@ -179,6 +223,7 @@ def __init__(
rope_scaling = getattr(config, "rope_scaling", None)
max_position_embeddings = getattr(config, "max_position_embeddings", 8192)
self.attention = InternLM2Attention(
config=config,
hidden_size=self.hidden_size,
num_heads=config.num_attention_heads,
num_kv_heads=config.num_key_value_heads,
Expand Down Expand Up @@ -309,6 +354,13 @@ def forward(
)

def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
from sglang.srt.environ import envs

if envs.SGLANG_ENABLE_WEIGHT_LOADER_V2.get():
return self._load_weights_v2(weights)
return self._legacy_load_weights(weights)

def _legacy_load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
stacked_params_mapping = [
# (param_name, shard_name, shard_id)
("gate_up_proj", "w1", 0),
Expand Down Expand Up @@ -355,5 +407,14 @@ def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
)
weight_loader(param, loaded_weight)

def _load_weights_v2(self, weights: Iterable[Tuple[str, torch.Tensor]]) -> set[str]:
from sglang.srt.model_loader.auto_loader import AutoWeightsLoader

loader = AutoWeightsLoader(
self,
ignore_unexpected_suffixes=[".bias", ".kv_scale"],
)
return loader.load_weights(weights)


EntryClass = InternLM2ForCausalLM
Loading
Loading