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
81 changes: 81 additions & 0 deletions tests/generate/utils_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -564,6 +564,87 @@ def test_transfer_state_with_mappings_gemma4(self):
)
)

def test_transfer_state_with_mappings_gemma(self):
"""Test transfer_state_with_mappings for Gemma."""
from tunix.models.gemma import mapping_vllm_jax
from tunix.models.gemma import model as gemma_model

mapping_config = mapping_vllm_jax.VLLM_JAX_MAPPING
self.assertIn("to_hf_mappings", gemma_model.Gemma.mapping_for("vllm_jax"))
self.assertIsNotNone(gemma_model.Gemma.to_hf_mappings("vllm_jax"))

src_params = {
"embedder.input_embedding": MockParam(
jnp.arange(16 * 32, dtype=jnp.float32).reshape(16, 32)
),
"layers.0.pre_attention_norm.scale": MockParam(
jnp.arange(32, dtype=jnp.float32)
),
"layers.0.attn.q_einsum.w": MockParam(
jnp.arange(4 * 32 * 8, dtype=jnp.float32).reshape(4, 32, 8)
),
"layers.0.attn.kv_einsum.w": MockParam(
jnp.arange(2 * 2 * 32 * 8, dtype=jnp.float32).reshape(2, 2, 32, 8)
),
"layers.0.mlp.gate_proj.kernel": MockParam(
jnp.arange(32 * 64, dtype=jnp.float32).reshape(32, 64)
),
"layers.0.mlp.up_proj.kernel": MockParam(
jnp.arange(32 * 64, dtype=jnp.float32).reshape(32, 64)
),
"layers.0.mlp.down_proj.kernel": MockParam(
jnp.arange(64 * 32, dtype=jnp.float32).reshape(64, 32)
),
"final_norm.scale": MockParam(jnp.arange(32, dtype=jnp.float32)),
}
src_state = MockState(src_params)

dst_params = {
"model.embed_tokens.weight": MockParam(
jnp.zeros((16, 32), dtype=jnp.float32)
),
"model.layers.0.input_layernorm.weight": MockParam(
jnp.zeros(32, dtype=jnp.float32)
),
"model.layers.0.self_attn.qkv_proj.weight": MockParam(
jnp.zeros((32, 64), dtype=jnp.float32)
),
"model.layers.0.mlp.gate_up_proj.weight": MockParam(
jnp.zeros((32, 128), dtype=jnp.float32)
),
"model.layers.0.mlp.down_proj.weight": MockParam(
jnp.zeros((64, 32), dtype=jnp.float32)
),
"model.norm.weight": MockParam(jnp.zeros(32, dtype=jnp.float32)),
}
dst_state = MockState(dst_params)

if "preprocess_src_state" in mapping_config:
src_state = mapping_config["preprocess_src_state"](src_state)

key_mappings = mapping_config["to_hf_mappings"]
transpose_keys = mapping_config["to_hf_transpose_keys"]

new_tgt_state = utils.transfer_state_with_mappings(
src_state,
dst_state,
key_mappings=key_mappings,
transpose_keys=transpose_keys,
)

self.assertTrue(
jnp.array_equal(
new_tgt_state.params["model.embed_tokens.weight"],
src_params["embedder.input_embedding"].value,
)
)
self.assertTrue(
jnp.array_equal(
new_tgt_state.params["model.norm.weight"],
src_params["final_norm.scale"].value,
)
)

def test_verify_state_closeness(self):
"""Test verify_state_closeness function with various scenarios."""

Expand Down
31 changes: 31 additions & 0 deletions tunix/models/gemma/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,31 @@
# Copyright 2026 Google LLC
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# https://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.

"""Gemma API."""

from tunix.models.gemma import mapping_vllm_jax
from tunix.models.gemma import model
from tunix.models.gemma import params
from tunix.models.gemma import params_safetensors

BACKEND_MAPPINGS = {
'vllm_jax': mapping_vllm_jax.VLLM_JAX_MAPPING,
}

__all__ = [
'BACKEND_MAPPINGS',
'model',
'params',
'params_safetensors',
]
279 changes: 279 additions & 0 deletions tunix/models/gemma/mapping_vllm_jax.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,279 @@
# Copyright 2026 Google LLC
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# https://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.

"""vLLM JAX backend mappings for Gemma models."""

from __future__ import annotations

from typing import Any, Dict, Tuple

from flax import nnx
import jax.numpy as jnp

Sharding = Tuple[str | None, ...]
MappingEntry = Tuple[str, Sharding]


TO_HF_MAPPINGS: Dict[str, MappingEntry] = {
'embedder.input_embedding': ('model.embed_tokens.weight', ('model', None)),
'layers.*.pre_attention_norm.scale': (
'model.layers.*.input_layernorm.weight',
(None,),
),
'layers.*.attn.q_einsum.w': (
'model.layers.*.self_attn.q_proj.weight',
(None, 'model', None),
),
'layers.*.attn.k_einsum.w': (
'model.layers.*.self_attn.k_proj.weight',
(None, 'model', None),
),
'layers.*.attn.kv_einsum.w': (
'model.layers.*.self_attn.kv_proj.weight',
(None, 'model'),
),
'layers.*.attn.qkv_einsum.w': (
'model.layers.*.self_attn.qkv_proj.weight',
(None, 'model'),
),
'layers.*.attn._query_norm.scale': (
'model.layers.*.self_attn.q_norm.weight',
(None,),
),
'layers.*.attn._key_norm.scale': (
'model.layers.*.self_attn.k_norm.weight',
(None,),
),
'layers.*.attn.attn_vec_einsum.w': (
'model.layers.*.self_attn.o_proj.weight',
('model', None, None),
),
'layers.*.post_attn_norm.scale': (
'model.layers.*.post_attention_layernorm.weight',
(None,),
),
'layers.*.post_attention_norm.scale': (
'model.layers.*.post_attention_layernorm.weight',
(None,),
),
'layers.*.pre_ffw_norm.scale': (
'model.layers.*.pre_feedforward_layernorm.weight',
(None,),
),
'layers.*.mlp.gate_up_proj.kernel': (
'model.layers.*.mlp.gate_up_proj.weight',
(None, 'model'),
),
'layers.*.mlp.gate_proj.kernel': (
'model.layers.*.mlp.gate_proj.weight',
(None, 'model'),
),
'layers.*.mlp.up_proj.kernel': (
'model.layers.*.mlp.up_proj.weight',
(None, 'model'),
),
'layers.*.mlp.down_proj.kernel': (
'model.layers.*.mlp.down_proj.weight',
('model', None),
),
'layers.*.post_ffw_norm.scale': (
'model.layers.*.post_feedforward_layernorm.weight',
(None,),
),
'final_norm.scale': ('model.norm.weight', (None,)),
}


LORA_TO_HF_MAPPINGS: Dict[str, MappingEntry] = {
'layers.*.mlp.gate_proj.kernel_lora_a': (
'model.layers.*.mlp.gate_proj.weight_lora_a',
(None, None),
),
'layers.*.mlp.gate_proj.kernel_lora_b': (
'model.layers.*.mlp.gate_proj.weight_lora_b',
(None, 'model'),
),
'layers.*.mlp.up_proj.kernel_lora_a': (
'model.layers.*.mlp.up_proj.weight_lora_a',
(None, None),
),
'layers.*.mlp.up_proj.kernel_lora_b': (
'model.layers.*.mlp.up_proj.weight_lora_b',
(None, 'model'),
),
'layers.*.mlp.down_proj.kernel_lora_a': (
'model.layers.*.mlp.down_proj.weight_lora_a',
('model', None),
),
'layers.*.mlp.down_proj.kernel_lora_b': (
'model.layers.*.mlp.down_proj.weight_lora_b',
(None, None),
),
'layers.*.attn.q_proj.w_lora_a': (
'model.layers.*.self_attn.q_proj.weight_lora_a',
('model', None),
),
'layers.*.attn.q_proj.w_lora_b': (
'model.layers.*.self_attn.q_proj.weight_lora_b',
(None, None),
),
'layers.*.attn.k_proj.w_lora_a': (
'model.layers.*.self_attn.k_proj.weight_lora_a',
('model', None),
),
'layers.*.attn.k_proj.w_lora_b': (
'model.layers.*.self_attn.k_proj.weight_lora_b',
(None, None),
),
'layers.*.attn.v_proj.w_lora_a': (
'model.layers.*.self_attn.v_proj.weight_lora_a',
('model', None),
),
'layers.*.attn.v_proj.w_lora_b': (
'model.layers.*.self_attn.v_proj.weight_lora_b',
(None, None),
),
'layers.*.attn.o_proj.w_lora_a': (
'model.layers.*.self_attn.o_proj.weight_lora_a',
('model', None),
),
'layers.*.attn.o_proj.w_lora_b': (
'model.layers.*.self_attn.o_proj.weight_lora_b',
(None, None),
),
}

TO_HF_TRANSPOSE_KEYS = {
'layers.*.attn.q_einsum.w': (1, 0, 2),
'layers.*.attn.k_einsum.w': (1, 0, 2),
}


def preprocess_src_state(src_state: Any) -> Any:
"""Fuses Q/K/V and MLP gate/up projections in the source state."""
if hasattr(src_state, 'flat_state'):
flat_state = list(src_state.flat_state())
new_flat_state = []

layers_q = {}
layers_k = {}
layers_kv = {}
layers_gate = {}
layers_up = {}

for keys, param in flat_state:
src_key = '.'.join(str(k) for k in keys)
if 'attn.q_einsum.w' in src_key:
layer_idx = keys[1]
layers_q[layer_idx] = (keys, param)
elif 'attn.k_einsum.w' in src_key:
layer_idx = keys[1]
layers_k[layer_idx] = (keys, param)
elif 'attn.kv_einsum.w' in src_key:
layer_idx = keys[1]
layers_kv[layer_idx] = (keys, param)
elif 'mlp.gate_proj.kernel' in src_key:
layer_idx = keys[1]
layers_gate[layer_idx] = (keys, param)
elif 'mlp.up_proj.kernel' in src_key:
layer_idx = keys[1]
layers_up[layer_idx] = (keys, param)
else:
new_flat_state.append((keys, param))

sample_kv_val = None
if layers_kv:
sample_kv_val = next(iter(layers_kv.values()))[1]
if hasattr(sample_kv_val, 'value'):
sample_kv_val = sample_kv_val.value

for layer_idx in layers_q:
q_keys, q_param = layers_q[layer_idx]
q_val = q_param.value if hasattr(q_param, 'value') else q_param
hidden_size = q_val.shape[1]
q_val_t = jnp.reshape(jnp.transpose(q_val, (1, 0, 2)), (hidden_size, -1))

if layer_idx in layers_kv:
_, kv_param = layers_kv[layer_idx]
kv_val = kv_param.value if hasattr(kv_param, 'value') else kv_param
k_val = kv_val[0]
v_val = kv_val[1]

k_val_t = jnp.reshape(
jnp.transpose(k_val, (1, 0, 2)), (hidden_size, -1)
)
v_val_t = jnp.reshape(
jnp.transpose(v_val, (1, 0, 2)), (hidden_size, -1)
)

qkv_val = jnp.concatenate([q_val_t, k_val_t, v_val_t], axis=-1)
qkv_keys = q_keys[:-2] + ('qkv_einsum', 'w')
if hasattr(q_param, 'value'):
new_flat_state.append((qkv_keys, nnx.Param(qkv_val)))
else:
new_flat_state.append((qkv_keys, qkv_val))
elif layer_idx in layers_k:
k_keys, k_param = layers_k[layer_idx]
new_flat_state.append((q_keys, q_param))
new_flat_state.append((k_keys, k_param))
elif sample_kv_val is not None:
# KV-shared layer
k_val = jnp.zeros_like(sample_kv_val[0])
v_val = jnp.zeros_like(sample_kv_val[1])
k_val_t = jnp.reshape(
jnp.transpose(k_val, (1, 0, 2)), (hidden_size, -1)
)
v_val_t = jnp.reshape(
jnp.transpose(v_val, (1, 0, 2)), (hidden_size, -1)
)

qkv_val = jnp.concatenate([q_val_t, k_val_t, v_val_t], axis=-1)
qkv_keys = q_keys[:-2] + ('qkv_einsum', 'w')
if hasattr(q_param, 'value'):
new_flat_state.append((qkv_keys, nnx.Param(qkv_val)))
else:
new_flat_state.append((qkv_keys, qkv_val))
else:
new_flat_state.append((q_keys, q_param))

for layer_idx in layers_gate:
gate_keys, gate_param = layers_gate[layer_idx]
_, up_param = layers_up[layer_idx]

gate_val = (
gate_param.value if hasattr(gate_param, 'value') else gate_param
)
up_val = up_param.value if hasattr(up_param, 'value') else up_param

gate_up_val = jnp.concatenate([gate_val, up_val], axis=-1)

gate_up_keys = gate_keys[:-2] + ('gate_up_proj', 'kernel')
if hasattr(gate_param, 'value'):
new_flat_state.append((gate_up_keys, nnx.Param(gate_up_val)))
else:
new_flat_state.append((gate_up_keys, gate_up_val))
src_state = src_state.from_flat_path(new_flat_state)
return src_state


VLLM_JAX_MAPPING: Dict[str, Any] = {
'to_hf_mappings': TO_HF_MAPPINGS,
'lora_to_hf_mappings': LORA_TO_HF_MAPPINGS,
'to_hf_transpose_keys': TO_HF_TRANSPOSE_KEYS,
'preprocess_src_state': preprocess_src_state,
}

__all__ = [
'VLLM_JAX_MAPPING',
]
Loading