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
28 changes: 28 additions & 0 deletions scripts/models/gemma4-26B-A4B.sh
Original file line number Diff line number Diff line change
@@ -0,0 +1,28 @@
MODEL_ARGS=(
--spec "vime_plugins.models.gemma4" "get_gemma4_spec"
--custom-model-provider-path "vime_plugins.models.gemma4_provider.model_provider"
--num-layers 30
--hidden-size 2816
--ffn-hidden-size 2112
--num-attention-heads 16
--group-query-attention
--num-query-groups 8
--kv-channels 256
--use-rotary-position-embeddings
--disable-bias-linear
--normalization "RMSNorm"
--norm-epsilon 1e-6
--rotary-base 10000
--rotary-percent 1.0
--vocab-size 262144
--qk-layernorm
--num-experts 128
--moe-ffn-hidden-size 704
--moe-router-topk 8
--moe-router-dtype fp32
--moe-router-score-function softmax
--moe-router-load-balancing-type none
--moe-aux-loss-coeff 0.0
--moe-token-dispatcher-type alltoall
--moe-grouped-gemm
)
1 change: 1 addition & 0 deletions tests/test_hf_to_megatron.py
Original file line number Diff line number Diff line change
Expand Up @@ -350,6 +350,7 @@ def test_loader_scope_stays_explicit():
assert set(_LOADERS) == {
"deepseek_v3",
"deepseek_v32",
"gemma4_text",
"glm4",
"glm4_moe",
"glm4_moe_lite",
Expand Down
234 changes: 234 additions & 0 deletions tests/utils/test_gemma4.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,234 @@
import importlib.util
from pathlib import Path
from types import SimpleNamespace

import pytest
import torch

from vime.backends.megatron_utils.hf_to_megatron import _LOADERS
from vime.backends.megatron_utils.hf_to_megatron.gemma4 import gemma4_hf_tensor
from vime.backends.megatron_utils.megatron_to_hf.gemma4 import _config_cache, _expert_buffers, convert_gemma4_to_hf

NUM_GPUS = 0

try:
_has_megatron = importlib.util.find_spec("megatron.core") is not None
except ModuleNotFoundError:
_has_megatron = False

requires_megatron = pytest.mark.skipif(not _has_megatron, reason="requires the Megatron runtime")


class Reader(dict):
def get_tensor(self, name):
return self[name]


def conversion_config(config):
return {
"global_attn_layers": {
index for index, layer_type in enumerate(config.layer_types) if layer_type == "full_attention"
},
"local_head_dim": config.head_dim,
"global_head_dim": config.global_head_dim,
"num_attention_heads": config.num_attention_heads,
"local_num_kv_heads": config.num_key_value_heads,
"global_num_kv_heads": config.num_global_key_value_heads,
"hidden_size": config.hidden_size,
"num_experts": config.num_experts,
}


@pytest.fixture
def config():
return SimpleNamespace(
model_type="gemma4_text",
enable_moe_block=True,
layer_types=["sliding_attention", "full_attention"],
num_attention_heads=4,
num_key_value_heads=2,
num_global_key_value_heads=1,
head_dim=2,
global_head_dim=4,
hidden_size=3,
num_experts=2,
)


@pytest.mark.unit
def test_registers_native_hf_loader():
assert _LOADERS["gemma4_text"] is gemma4_hf_tensor


@pytest.mark.unit
def test_native_model_does_not_import_bridge():
source = (Path(__file__).resolve().parents[2] / "vime_plugins/models/gemma4.py").read_text()

assert "megatron.bridge" not in source
assert "mbridge" not in source


@pytest.mark.unit
@requires_megatron
def test_native_model_extends_transformer_config():
from megatron.core.transformer.transformer_config import TransformerConfig

from vime_plugins.models.gemma4 import Gemma4TransformerConfig

assert issubclass(Gemma4TransformerConfig, TransformerConfig)


@pytest.mark.unit
@requires_megatron
def test_marks_experts_for_direct_hf_loading():
from vime_plugins.models.gemma4_provider import _mark_expert_weights_for_direct_hf_loading

fc1 = torch.nn.Parameter(torch.zeros(1))
fc2 = torch.nn.Parameter(torch.zeros(1))
dense = torch.nn.Parameter(torch.zeros(1))
model = SimpleNamespace(
named_parameters=lambda: iter(
[
("decoder.layers.0.mlp.experts.linear_fc1.weight0", fc1),
("decoder.layers.0.mlp.experts.linear_fc2.weight0", fc2),
("decoder.layers.0.self_attention.linear_proj.weight", dense),
]
)
)

_mark_expert_weights_for_direct_hf_loading(model)

assert (fc1.tensor_model_parallel, fc1.partition_dim, fc1.partition_stride) == (True, 0, 1)
assert (fc2.tensor_model_parallel, fc2.partition_dim, fc2.partition_stride) == (True, 1, 1)
assert not hasattr(dense, "tensor_model_parallel")


@pytest.mark.unit
@requires_megatron
def test_promotes_layer_scalars_for_direct_hf_loading():
from vime_plugins.models.gemma4_provider import _promote_layer_scalars

model = torch.nn.Module()
model.layer = torch.nn.Module()
model.layer.register_buffer("layer_scalar", torch.tensor([0.5]))

_promote_layer_scalars(model)

assert torch.equal(dict(model.named_parameters())["layer.layer_scalar"], torch.tensor([0.5]))
assert not model.layer.layer_scalar.requires_grad


@pytest.mark.unit
def test_tied_output_weight_is_not_exported(config):
args = SimpleNamespace(hf_checkpoint="checkpoint")
_config_cache[args.hf_checkpoint] = conversion_config(config)

assert convert_gemma4_to_hf(args, "module.module.output_layer.weight", torch.empty(1)) == []


@pytest.mark.unit
@pytest.mark.parametrize("layer", [0, 1])
def test_qkv_round_trip(config, layer):
attention_config = (
config
if layer == 0
else SimpleNamespace(
num_attention_heads=config.num_attention_heads,
num_key_value_heads=config.num_global_key_value_heads,
head_dim=config.global_head_dim,
hidden_size=config.hidden_size,
)
)
q = torch.arange(
attention_config.num_attention_heads * attention_config.head_dim * config.hidden_size,
dtype=torch.float32,
).reshape(-1, config.hidden_size)
k = torch.arange(
attention_config.num_key_value_heads * attention_config.head_dim * config.hidden_size,
dtype=torch.float32,
).reshape(-1, config.hidden_size)
v = k if layer == 1 else k + 1000
prefix = f"model.language_model.layers.{layer}.self_attn"
reader = Reader(
{
f"{prefix}.q_proj.weight": q,
f"{prefix}.k_proj.weight": k,
**({} if layer == 1 else {f"{prefix}.v_proj.weight": v}),
}
)

megatron = gemma4_hf_tensor(
f"decoder.layers.{layer}.self_attention.linear_qkv.weight",
reader,
config,
)
args = SimpleNamespace(hf_checkpoint="checkpoint")
_config_cache[args.hf_checkpoint] = conversion_config(config)
converted = dict(
convert_gemma4_to_hf(
args,
f"module.module.decoder.layers.{layer}.self_attention.linear_qkv.weight",
megatron,
)
)

assert torch.equal(converted[f"{prefix}.q_proj.weight"], q)
assert torch.equal(converted[f"{prefix}.k_proj.weight"], k)
if layer == 0:
assert torch.equal(converted[f"{prefix}.v_proj.weight"], v)
else:
assert f"{prefix}.v_proj.weight" not in converted


@pytest.mark.unit
def test_expert_weights_round_trip_as_packed_hf_tensor(config):
args = SimpleNamespace(hf_checkpoint="checkpoint")
_config_cache[args.hf_checkpoint] = conversion_config(config)
_expert_buffers.clear()
packed = torch.arange(2 * 6 * 3, dtype=torch.float32).reshape(2, 6, 3)
reader = Reader({"model.language_model.layers.0.experts.gate_up_proj": packed})

first = gemma4_hf_tensor("decoder.layers.0.mlp.experts.linear_fc1.weight0", reader, config)
second = gemma4_hf_tensor("decoder.layers.0.mlp.experts.linear_fc1.weight1", reader, config)
assert (
convert_gemma4_to_hf(
args,
"module.module.decoder.layers.0.mlp.experts.linear_fc1.weight0",
first,
)
== []
)
converted = convert_gemma4_to_hf(
args,
"module.module.decoder.layers.0.mlp.experts.linear_fc1.weight1",
second,
)

assert converted[0][0] == "model.language_model.layers.0.experts.gate_up_proj"
assert torch.equal(converted[0][1], packed)
assert not _expert_buffers


@pytest.mark.unit
def test_dense_mlp_and_router_scale(config):
prefix = "model.language_model.layers.0"
gate = torch.arange(18, dtype=torch.float32).reshape(6, 3)
up = gate + 100
router_scale = torch.arange(3, dtype=torch.float32)
reader = Reader(
{
f"{prefix}.mlp.gate_proj.weight": gate,
f"{prefix}.mlp.up_proj.weight": up,
f"{prefix}.router.scale": router_scale,
}
)

assert torch.equal(
gemma4_hf_tensor("decoder.layers.0.dense_mlp.linear_fc1.weight", reader, config),
torch.cat((gate, up)),
)
assert torch.equal(gemma4_hf_tensor("decoder.layers.0.mlp.router.scale", reader, config), router_scale)


if __name__ == "__main__":
raise SystemExit(pytest.main([__file__]))
5 changes: 5 additions & 0 deletions vime/backends/megatron_utils/hf_to_megatron/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@

from .common import load_model_hf_weights
from .deepseek import deepseek_hf_tensor
from .gemma4 import gemma4_hf_tensor
from .glm import glm4_hf_tensor, glm4_moe_hf_tensor
from .qwen import mimo_hf_tensor, minimax_m2_hf_tensor, qwen_hf_tensor, qwen_moe_hf_tensor
from .qwen3_5 import qwen3_5_hf_tensor
Expand All @@ -12,6 +13,7 @@
_LOADERS = {
"deepseek_v3": deepseek_hf_tensor,
"deepseek_v32": deepseek_hf_tensor,
"gemma4_text": gemma4_hf_tensor,
"glm4": glm4_hf_tensor,
"glm4_moe": glm4_moe_hf_tensor,
"glm4_moe_lite": deepseek_hf_tensor,
Expand All @@ -32,11 +34,14 @@

def supports_hf_weight_loading(path: str | Path) -> bool:
config = AutoConfig.from_pretrained(path, trust_remote_code=True)
config = getattr(config, "text_config", config)
return config.model_type in _LOADERS


def load_hf_weights(args, model, path: str | Path) -> None:
config = AutoConfig.from_pretrained(path, trust_remote_code=True)
config = getattr(config, "text_config", config)

try:
get_hf_tensor = _LOADERS[config.model_type]
except KeyError as exc:
Expand Down
91 changes: 91 additions & 0 deletions vime/backends/megatron_utils/hf_to_megatron/gemma4.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,91 @@
import re
from types import SimpleNamespace

import torch

from .common import SafetensorReader, merge_gate_up, merge_qkv, strip_mcore_wrappers, text_config


def _layer(name: str) -> tuple[int, str]:
match = re.fullmatch(r"decoder\.layers\.(\d+)\.(.+)", name)
if not match:
raise KeyError(f"Unsupported Gemma4 Megatron parameter {name!r}")
return int(match.group(1)), match.group(2)


def _attention_config(config, layer: int):
config = text_config(config)
if hasattr(config, "per_layer_config"):
config = config.per_layer_config[layer]
return SimpleNamespace(
num_attention_heads=config.num_attention_heads,
num_key_value_heads=config.num_key_value_heads,
head_dim=config.head_dim,
hidden_size=config.hidden_size,
)
if config.layer_types[layer] != "full_attention":
return config
return SimpleNamespace(
num_attention_heads=config.num_attention_heads,
num_key_value_heads=config.num_global_key_value_heads,
head_dim=config.global_head_dim,
hidden_size=config.hidden_size,
)


def gemma4_hf_tensor(name: str, reader: SafetensorReader, config) -> torch.Tensor:
name = strip_mcore_wrappers(name)
direct = {
"embedding.word_embeddings.weight": "model.language_model.embed_tokens.weight",
"decoder.final_layernorm.weight": "model.language_model.norm.weight",
"output_layer.weight": "model.language_model.embed_tokens.weight",
}
if name in direct:
return reader.get_tensor(direct[name])

layer, rest = _layer(name)
prefix = f"model.language_model.layers.{layer}"
mapping = {
"self_attention.linear_qkv.layer_norm_weight": "input_layernorm.weight",
"self_attention.q_layernorm.weight": "self_attn.q_norm.weight",
"self_attention.k_layernorm.weight": "self_attn.k_norm.weight",
"self_attention.linear_proj.weight": "self_attn.o_proj.weight",
"pre_mlp_layernorm.weight": "pre_feedforward_layernorm.weight",
"dense_mlp.linear_fc1.layer_norm_weight": "pre_feedforward_layernorm.weight",
"post_attention_layernorm.weight": "post_attention_layernorm.weight",
"post_feedforward_layernorm.weight": "post_feedforward_layernorm.weight",
"mlp.pre_feedforward_layernorm_2.weight": "pre_feedforward_layernorm_2.weight",
"pre_feedforward_layernorm_2.weight": "pre_feedforward_layernorm_2.weight",
"post_feedforward_layernorm_1.weight": "post_feedforward_layernorm_1.weight",
"post_feedforward_layernorm_2.weight": "post_feedforward_layernorm_2.weight",
"mlp.router.proj.weight": "router.proj.weight",
"mlp.router.scale": "router.scale",
"mlp.router.per_expert_scale": "router.per_expert_scale",
"layer_scalar": "layer_scalar",
}
if rest in mapping:
return reader.get_tensor(f"{prefix}.{mapping[rest]}")

if rest == "self_attention.linear_qkv.weight":
attention_config = _attention_config(config, layer)
q = reader.get_tensor(f"{prefix}.self_attn.q_proj.weight")
k = reader.get_tensor(f"{prefix}.self_attn.k_proj.weight")
v_name = f"{prefix}.self_attn.v_proj.weight"
v = reader.get_tensor(v_name) if v_name in reader else k
return merge_qkv(q, k, v, attention_config)

if rest in {"mlp.linear_fc1.weight", "dense_mlp.linear_fc1.weight"}:
return merge_gate_up(
reader.get_tensor(f"{prefix}.mlp.gate_proj.weight"),
reader.get_tensor(f"{prefix}.mlp.up_proj.weight"),
)
if rest in {"mlp.linear_fc2.weight", "dense_mlp.linear_fc2.weight"}:
return reader.get_tensor(f"{prefix}.mlp.down_proj.weight")

match = re.fullmatch(r"mlp\.experts\.linear_fc([12])\.weight(\d+)", rest)
if match:
projection, expert = map(int, match.groups())
packed_name = "gate_up_proj" if projection == 1 else "down_proj"
return reader.get_tensor(f"{prefix}.experts.{packed_name}")[expert].contiguous()

raise KeyError(f"Unsupported Gemma4 Megatron parameter {name!r}")
Loading