chore: import upstream snapshot with attribution
Lint / lint (push) Has been cancelled
Build Docs / Deploy Docs (push) Has been cancelled
Windows CI / Windows (push) Has been cancelled

This commit is contained in:
wehub-resource-sync
2026-07-13 13:23:58 +08:00
commit 770d92cb1f
694 changed files with 114634 additions and 0 deletions
@@ -0,0 +1,48 @@
"""
This file specifies how MLC's Gemma2 parameter maps from other formats, for example HuggingFace
PyTorch, HuggingFace safetensors.
"""
import functools
from mlc_llm.loader import ExternMapping
from mlc_llm.loader.standard_loader import make_standard_hf_loader
from mlc_llm.quantization import Quantization
from .gemma2_model import Gemma2Config, Gemma2ForCausalLM
def huggingface(model_config: Gemma2Config, quantization: Quantization) -> ExternMapping:
"""Create HF weight mapping for Gemma2."""
model = Gemma2ForCausalLM(model_config)
if quantization is not None:
model.to(quantization.model_dtype)
_, _named_params, _ = model.export_tvm(
spec=model.get_default_spec(),
allow_extern=True,
)
named_parameters = dict(_named_params)
base_loader = make_standard_hf_loader(
model_cls=Gemma2ForCausalLM,
)
mapping = base_loader(model_config, quantization)
def add_one(name: str) -> None:
mlc_param = named_parameters[name]
mapping.add_mapping(
name,
[name],
functools.partial(
lambda x, dtype: (x + 1).astype(dtype),
dtype=mlc_param.dtype,
),
)
for i in range(model_config.num_hidden_layers):
add_one(f"model.layers.{i}.input_layernorm.weight")
add_one(f"model.layers.{i}.post_attention_layernorm.weight")
add_one(f"model.layers.{i}.pre_feedforward_layernorm.weight")
add_one(f"model.layers.{i}.post_feedforward_layernorm.weight")
add_one("model.norm.weight")
return mapping
+122
View File
@@ -0,0 +1,122 @@
"""Implementation for Gemma2 architecture."""
import dataclasses
from tvm.relax.frontend import nn
from tvm.relax.frontend.nn import Tensor, op
from mlc_llm.model.gemma.gemma_model import (
GemmaAttention,
GemmaConfig,
GemmaForCausalLM,
GemmaMLP,
GemmaModel,
)
from mlc_llm.nn import PagedKVCache
from mlc_llm.support import logging
from mlc_llm.support import tensor_parallel as tp
logger = logging.getLogger(__name__)
@dataclasses.dataclass
class Gemma2Config(GemmaConfig):
"""Configuration of the Gemma2 model, in addition to the Gemma model"""
# NOTE: We ignore attn_logit_softcapping in the gemma2 implementation for now.
# The Gemma 2 team observed minor differences when soft-capping is removed during inference,
# according to https://huggingface.co/blog/gemma2.
# The soft-capping is also not supported by HuggingFace transformers `Gemma2SdpaAttention`.
attn_logit_softcapping: float = None
final_logit_softcapping: float = None
query_pre_attn_scalar: int = None
sliding_window: int = None
def __post_init__(self):
super().__post_init__()
# NOTE: override the context window size with the Gemma2 sliding window size,
# as the sliding window attention every other layer is yet to be supported.
self.context_window_size = self.sliding_window
class Gemma2Attention(GemmaAttention):
def __init__(self, config: Gemma2Config):
super().__init__(config)
self.scaling_factor = (config.head_dim / config.query_pre_attn_scalar) ** 0.5
class Gemma2DecoderLayer(nn.Module):
def __init__(self, config: Gemma2Config):
rms_norm_eps = config.rms_norm_eps
self.self_attn = Gemma2Attention(config)
self.mlp = GemmaMLP(config)
# Gemma RMSNorm adds 1 to the weights. It is already fused in the loader
self.input_layernorm = nn.RMSNorm(config.hidden_size, -1, rms_norm_eps, bias=False)
self.post_attention_layernorm = nn.RMSNorm(config.hidden_size, -1, rms_norm_eps, bias=False)
self.pre_feedforward_layernorm = nn.RMSNorm(
config.hidden_size, -1, rms_norm_eps, bias=False
)
self.post_feedforward_layernorm = nn.RMSNorm(
config.hidden_size, -1, rms_norm_eps, bias=False
)
def _set_tp():
def _set(layer, hint):
layer.weight.attrs["shard_strategy"] = hint
hd = config.head_dim
q = self.self_attn.num_q_heads * hd
k = self.self_attn.num_kv_heads * hd
v = self.self_attn.num_kv_heads * hd
i = self.mlp.intermediate_size
_set(
self.self_attn.qkv_proj,
tp.ShardSingleDim("_shard_qkv", segs=[q, k, v], dim=0),
)
_set(self.self_attn.o_proj, tp.ShardSingleDim("_shard_o", dim=1))
_set(
self.mlp.gate_up_proj,
tp.ShardSingleDim("_shard_mlp_up", segs=[i, i], dim=0),
)
_set(self.mlp.down_proj, tp.ShardSingleDim("_shard_mlp_down", dim=1))
self.tensor_parallel_shards = config.tensor_parallel_shards
_set_tp()
def forward(self, hidden_states: Tensor, paged_kv_cache: PagedKVCache, layer_id: int):
out = self.self_attn(self.input_layernorm(hidden_states), paged_kv_cache, layer_id)
out = self._apply_post_matmul_norm(out, norm=self.post_attention_layernorm)
hidden_states = out + hidden_states
out = self.pre_feedforward_layernorm(hidden_states)
out = self.mlp(out)
out = self._apply_post_matmul_norm(out, norm=self.post_feedforward_layernorm)
hidden_states = out + hidden_states
return hidden_states
def _apply_post_matmul_norm(self, out: Tensor, norm: nn.Tensor):
if self.tensor_parallel_shards > 1:
return norm(op.ccl_allreduce(out, "sum"))
return norm(out)
class Gemma2Model(GemmaModel):
def __init__(self, config: Gemma2Config):
super().__init__(config)
self.layers = nn.ModuleList(
[Gemma2DecoderLayer(config) for _ in range(config.num_hidden_layers)]
)
class Gemma2ForCausalLM(GemmaForCausalLM):
def __init__(self, config: Gemma2Config):
super().__init__(config)
self.model = Gemma2Model(config)
self.final_logit_softcapping = config.final_logit_softcapping
def get_logits(self, hidden_states: Tensor):
logits = super().get_logits(hidden_states)
if self.final_logit_softcapping is not None:
logits = op.tanh(logits / self.final_logit_softcapping) * self.final_logit_softcapping
return logits