Skip to content

Commit 53cdfa9

Browse files
kharshith-kdivyashreepathihalligemini-code-assist[bot]
authored
Mistral kerashub to HF Safetensors exporter (#2743)
* Add Mistral export to HuggingFace safetensors format - Implement mistral.py with config, weights mapping, and tokenizer config - Add comprehensive test suite in mistral_test.py - Register Mistral exporters in hf_exporter.py - Support SentencePiece tokenizer with proper file renaming - Verified end-to-end export and text generation with transformers library * Developed standalone export script and xla testing scripts * Remove mistral-inference dependency from export script * Rewrite run_mistral_xla.py: remove mistral-inference, implement Mistral in pure PyTorch * run_mistral_xla: remove xmp.spawn, call generate directly for Colab TPU * run_mistral_xla: use torch_xla.device/sync new API, fallback to xm for older versions * Fix pre-commit: ruff E501/F401, api_gen formatting * Fix mistral_test.py: use_fast=False, in-vocab test text, torch.tensor for HF model * Fix mistral_test.py: read rope_theta from config.json, add json import * Replace tools/mistral/ with tools/checkpoint_export/verify_mistral_export.py * Fix get_mistral_config: add max_position_embeddings, drop attention/mlp_bias * Fix E501: wrap long usage lines in verify_mistral_export.py docstring * Fix mistral_test: drop fragile HF/KerasHub token equality, use .cpu().numpy() * Address Gemini review: filter None special tokens, clarify hardcoded constants * Fix mistral_test: use ops.convert_to_numpy for cross-backend compatibility * Update keras_hub/src/utils/transformers/export/mistral.py Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com> --------- Co-authored-by: Divyashree Sreepathihalli <divyashreepathihalli@gmail.com> Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
1 parent d9c2035 commit 53cdfa9

4 files changed

Lines changed: 791 additions & 2 deletions

File tree

keras_hub/src/utils/transformers/export/hf_exporter.py

Lines changed: 18 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -28,6 +28,15 @@
2828
)
2929
from keras_hub.src.utils.transformers.export.gpt2 import get_gpt2_weights_map
3030

31+
# --- Mistral Utils ---
32+
from keras_hub.src.utils.transformers.export.mistral import get_mistral_config
33+
from keras_hub.src.utils.transformers.export.mistral import (
34+
get_mistral_tokenizer_config,
35+
)
36+
from keras_hub.src.utils.transformers.export.mistral import (
37+
get_mistral_weights_map,
38+
)
39+
3140
# --- Qwen Utils ---
3241
from keras_hub.src.utils.transformers.export.qwen import get_qwen_config
3342
from keras_hub.src.utils.transformers.export.qwen import (
@@ -38,20 +47,23 @@
3847
MODEL_CONFIGS = {
3948
"GemmaBackbone": get_gemma_config,
4049
"Gemma3Backbone": get_gemma3_config,
50+
"MistralBackbone": get_mistral_config,
4151
"QwenBackbone": get_qwen_config,
4252
"GPT2Backbone": get_gpt2_config,
4353
}
4454

4555
MODEL_EXPORTERS = {
4656
"GemmaBackbone": get_gemma_weights_map,
4757
"Gemma3Backbone": get_gemma3_weights_map,
58+
"MistralBackbone": get_mistral_weights_map,
4859
"QwenBackbone": get_qwen_weights_map,
4960
"GPT2Backbone": get_gpt2_weights_map,
5061
}
5162

5263
MODEL_TOKENIZER_CONFIGS = {
5364
"GemmaTokenizer": get_gemma_tokenizer_config,
5465
"Gemma3Tokenizer": get_gemma3_tokenizer_config,
66+
"MistralTokenizer": get_mistral_tokenizer_config,
5567
"QwenTokenizer": get_qwen_tokenizer_config,
5668
"GPT2Tokenizer": get_gpt2_tokenizer_config,
5769
}
@@ -211,8 +223,12 @@ def export_tokenizer(tokenizer, path):
211223

212224
# Rename files to match Hugging Face expectations
213225

214-
# 1. SentencePiece Models (Gemma / Gemma 3)
215-
if tokenizer_type in ["GemmaTokenizer", "Gemma3Tokenizer"]:
226+
# 1. SentencePiece Models (Gemma / Gemma 3 / Mistral)
227+
if tokenizer_type in [
228+
"GemmaTokenizer",
229+
"Gemma3Tokenizer",
230+
"MistralTokenizer",
231+
]:
216232
vocab_spm_path = os.path.join(path, "vocabulary.spm")
217233
tokenizer_model_path = os.path.join(path, "tokenizer.model")
218234
if os.path.exists(vocab_spm_path):
Lines changed: 176 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,176 @@
1+
import keras.ops as ops
2+
3+
4+
def get_mistral_config(backbone):
5+
"""Convert Keras Mistral backbone config to Hugging Face dictionary."""
6+
head_dim = backbone.hidden_dim // backbone.num_query_heads
7+
8+
hf_config = {
9+
"architectures": ["MistralForCausalLM"],
10+
"model_type": "mistral",
11+
"vocab_size": backbone.vocabulary_size,
12+
"num_hidden_layers": backbone.num_layers,
13+
"num_attention_heads": backbone.num_query_heads,
14+
"num_key_value_heads": backbone.num_key_value_heads,
15+
"hidden_size": backbone.hidden_dim,
16+
"intermediate_size": backbone.intermediate_dim,
17+
"head_dim": head_dim,
18+
"rms_norm_eps": backbone.layer_norm_epsilon,
19+
"rope_theta": backbone.rope_max_wavelength,
20+
# All Mistral 7B variants use max_position_embeddings=32768.
21+
# The backbone does not expose this field; the HF default (131072)
22+
# differs from the canonical Mistral checkpoints.
23+
"max_position_embeddings": 32768,
24+
"hidden_act": "silu",
25+
"attention_dropout": backbone.dropout,
26+
"tie_word_embeddings": False,
27+
"use_cache": True,
28+
"sliding_window": backbone.sliding_window,
29+
"torch_dtype": backbone.dtype_policy.compute_dtype,
30+
}
31+
32+
return hf_config
33+
34+
35+
def get_mistral_weights_map(backbone, include_lm_head=False):
36+
"""Create a Keras-to-HuggingFace weight name mapping for Mistral.
37+
38+
Args:
39+
backbone: A `keras_hub.models.MistralBackbone` instance.
40+
include_lm_head: If True, includes the ``lm_head.weight`` tensor
41+
(used when exporting a CausalLM task).
42+
43+
Returns:
44+
dict mapping HuggingFace weight keys to Keras tensors.
45+
"""
46+
weights_map = {}
47+
48+
# Token embeddings
49+
weights_map["model.embed_tokens.weight"] = (
50+
backbone.token_embedding.embeddings
51+
)
52+
53+
for i in range(backbone.num_layers):
54+
decoder_layer = backbone.transformer_layers[i]
55+
attn_layer = decoder_layer._self_attention_layer
56+
57+
# --- Normalization ---
58+
# Pre-attention (input) layernorm
59+
weights_map[f"model.layers.{i}.input_layernorm.weight"] = (
60+
decoder_layer._self_attention_layernorm.scale
61+
)
62+
# Pre-MLP (post-attention) layernorm
63+
weights_map[f"model.layers.{i}.post_attention_layernorm.weight"] = (
64+
decoder_layer._feedforward_layernorm.scale
65+
)
66+
67+
# --- Attention projections ---
68+
# Keras Q kernel: (hidden_dim, num_query_heads, head_dim)
69+
# HF q_proj.weight: (num_query_heads * head_dim, hidden_dim)
70+
q_kernel = attn_layer._query_dense.kernel
71+
q_kernel = ops.reshape(q_kernel, (backbone.hidden_dim, -1))
72+
weights_map[f"model.layers.{i}.self_attn.q_proj.weight"] = (
73+
ops.transpose(q_kernel)
74+
)
75+
76+
# Keras K kernel: (hidden_dim, num_key_value_heads, head_dim)
77+
# HF k_proj.weight: (num_key_value_heads * head_dim, hidden_dim)
78+
k_kernel = attn_layer._key_dense.kernel
79+
k_kernel = ops.reshape(k_kernel, (backbone.hidden_dim, -1))
80+
weights_map[f"model.layers.{i}.self_attn.k_proj.weight"] = (
81+
ops.transpose(k_kernel)
82+
)
83+
84+
# Keras V kernel: (hidden_dim, num_key_value_heads, head_dim)
85+
# HF v_proj.weight: (num_key_value_heads * head_dim, hidden_dim)
86+
v_kernel = attn_layer._value_dense.kernel
87+
v_kernel = ops.reshape(v_kernel, (backbone.hidden_dim, -1))
88+
weights_map[f"model.layers.{i}.self_attn.v_proj.weight"] = (
89+
ops.transpose(v_kernel)
90+
)
91+
92+
# Keras O kernel: (num_query_heads, head_dim, hidden_dim)
93+
# HF o_proj.weight: (hidden_dim, num_query_heads * head_dim)
94+
o_kernel = attn_layer._output_dense.kernel
95+
o_kernel = ops.reshape(o_kernel, (-1, backbone.hidden_dim))
96+
weights_map[f"model.layers.{i}.self_attn.o_proj.weight"] = (
97+
ops.transpose(o_kernel)
98+
)
99+
100+
# --- MLP (SwiGLU) ---
101+
# Keras gate/up kernel: (hidden_dim, intermediate_dim)
102+
# HF gate/up_proj.weight: (intermediate_dim, hidden_dim)
103+
gate_kernel = decoder_layer._feedforward_gate_dense.kernel
104+
weights_map[f"model.layers.{i}.mlp.gate_proj.weight"] = ops.transpose(
105+
gate_kernel
106+
)
107+
108+
up_kernel = decoder_layer._feedforward_intermediate_dense.kernel
109+
weights_map[f"model.layers.{i}.mlp.up_proj.weight"] = ops.transpose(
110+
up_kernel
111+
)
112+
113+
# Keras down kernel: (intermediate_dim, hidden_dim)
114+
# HF down_proj.weight: (hidden_dim, intermediate_dim)
115+
down_kernel = decoder_layer._feedforward_output_dense.kernel
116+
weights_map[f"model.layers.{i}.mlp.down_proj.weight"] = ops.transpose(
117+
down_kernel
118+
)
119+
120+
# Final layernorm
121+
weights_map["model.norm.weight"] = backbone.layer_norm.scale
122+
123+
# LM head (only when exporting the full CausalLM task)
124+
if include_lm_head:
125+
# Mistral models typically don't tie embeddings
126+
weights_map["lm_head.weight"] = ops.transpose(
127+
backbone.token_embedding.reverse_embeddings
128+
)
129+
130+
return weights_map
131+
132+
133+
def get_mistral_tokenizer_config(tokenizer):
134+
"""Build a HuggingFace-compatible tokenizer_config.json for Mistral."""
135+
tokenizer_config = {
136+
"add_bos_token": True,
137+
"add_eos_token": False,
138+
"added_tokens_decoder": {},
139+
"bos_token": tokenizer.start_token,
140+
"clean_up_tokenization_spaces": False,
141+
"eos_token": tokenizer.end_token,
142+
"legacy": False,
143+
# 32768 matches the canonical Mistral context window (same as
144+
# max_position_embeddings set in get_mistral_config).
145+
"model_max_length": 32768,
146+
"pad_token": None,
147+
"sp_model_kwargs": {},
148+
"spaces_between_special_tokens": False,
149+
"tokenizer_class": "LlamaTokenizer",
150+
# Mistral's SentencePiece model always uses "<unk>" as the unknown
151+
# token piece; MistralTokenizer does not expose an unk_token property.
152+
"unk_token": "<unk>",
153+
"use_default_system_prompt": False,
154+
}
155+
156+
# Add added_tokens_decoder; filter out any tokens that are None (e.g. if
157+
# a subclass does not define start_token / end_token).
158+
added_tokens_decoder = {}
159+
special_tokens = [
160+
t
161+
for t in [tokenizer.start_token, tokenizer.end_token, "<unk>"]
162+
if t is not None
163+
]
164+
for token in special_tokens:
165+
token_id = tokenizer.token_to_id(token)
166+
if token_id is not None:
167+
added_tokens_decoder[str(token_id)] = {
168+
"content": token,
169+
"special": True,
170+
"single_word": False,
171+
"lstrip": False,
172+
"rstrip": False,
173+
"normalized": False,
174+
}
175+
tokenizer_config["added_tokens_decoder"] = added_tokens_decoder
176+
return tokenizer_config
Lines changed: 160 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,160 @@
1+
import json
2+
import os
3+
4+
import keras.ops as ops
5+
import numpy as np
6+
import torch
7+
from transformers import AutoModelForCausalLM
8+
from transformers import AutoTokenizer
9+
10+
from keras_hub.src.models.mistral.mistral_backbone import MistralBackbone
11+
from keras_hub.src.models.mistral.mistral_causal_lm import MistralCausalLM
12+
from keras_hub.src.models.mistral.mistral_causal_lm_preprocessor import (
13+
MistralCausalLMPreprocessor,
14+
)
15+
from keras_hub.src.models.mistral.mistral_tokenizer import MistralTokenizer
16+
from keras_hub.src.tests.test_case import TestCase
17+
from keras_hub.src.utils.transformers.export.hf_exporter import (
18+
export_to_safetensors,
19+
)
20+
21+
22+
class TestMistralExport(TestCase):
23+
def test_export_to_hf(self):
24+
# 1. Create tokenizer from test vocab
25+
proto = os.path.join(self.get_test_data_dir(), "mistral_test_vocab.spm")
26+
tokenizer = MistralTokenizer(proto=proto)
27+
28+
# 2. Create a small backbone
29+
backbone = MistralBackbone(
30+
vocabulary_size=tokenizer.vocabulary_size(),
31+
num_layers=2,
32+
num_query_heads=4,
33+
num_key_value_heads=2,
34+
hidden_dim=64,
35+
intermediate_dim=128,
36+
rope_max_wavelength=10000,
37+
layer_norm_epsilon=1e-6,
38+
sliding_window=512,
39+
)
40+
41+
# 3. Create preprocessor & model
42+
preprocessor = MistralCausalLMPreprocessor(
43+
tokenizer=tokenizer, sequence_length=16
44+
)
45+
keras_model = MistralCausalLM(
46+
backbone=backbone, preprocessor=preprocessor
47+
)
48+
49+
# 4. Set all weights to deterministic random values
50+
rng = np.random.default_rng(42)
51+
weights = keras_model.get_weights()
52+
for i in range(len(weights)):
53+
weights[i] = rng.random(weights[i].shape).astype(weights[i].dtype)
54+
keras_model.set_weights(weights)
55+
56+
# 5. Export to Hugging Face format
57+
export_path = os.path.join(self.get_temp_dir(), "export_task")
58+
export_to_safetensors(keras_model, export_path)
59+
60+
# 6. Verify exported files
61+
exported = os.listdir(export_path)
62+
self.assertIn("config.json", exported)
63+
self.assertIn("model.safetensors", exported)
64+
self.assertIn("tokenizer.model", exported)
65+
self.assertIn("tokenizer_config.json", exported)
66+
67+
# 7. Load with Hugging Face Transformers
68+
# use_fast=False: the tiny test vocab is not a standard BPE/Unigram
69+
# SentencePiece model, so the fast tokenizer conversion fails.
70+
hf_tokenizer = AutoTokenizer.from_pretrained(
71+
export_path, use_fast=False
72+
)
73+
hf_model = AutoModelForCausalLM.from_pretrained(export_path)
74+
75+
# 8. Verify configuration
76+
hf_config = hf_model.config
77+
self.assertEqual(
78+
hf_config.vocab_size,
79+
backbone.vocabulary_size,
80+
"Vocabulary sizes do not match",
81+
)
82+
self.assertEqual(
83+
hf_config.num_hidden_layers,
84+
backbone.num_layers,
85+
"Number of layers do not match",
86+
)
87+
self.assertEqual(
88+
hf_config.num_attention_heads,
89+
backbone.num_query_heads,
90+
"Number of query heads do not match",
91+
)
92+
self.assertEqual(
93+
hf_config.num_key_value_heads,
94+
backbone.num_key_value_heads,
95+
"Number of key-value heads do not match",
96+
)
97+
self.assertEqual(
98+
hf_config.hidden_size,
99+
backbone.hidden_dim,
100+
"Hidden dimensions do not match",
101+
)
102+
self.assertEqual(
103+
hf_config.intermediate_size,
104+
backbone.intermediate_dim,
105+
"Intermediate dimensions do not match",
106+
)
107+
self.assertEqual(
108+
hf_config.rms_norm_eps,
109+
backbone.layer_norm_epsilon,
110+
"Layer norm epsilons do not match",
111+
)
112+
# rope_theta was added as an explicit MistralConfig attribute in a
113+
# later transformers release; read from the saved config.json directly
114+
# to stay version-agnostic.
115+
with open(os.path.join(export_path, "config.json")) as f:
116+
saved_config = json.load(f)
117+
self.assertEqual(
118+
saved_config.get("rope_theta"),
119+
backbone.rope_max_wavelength,
120+
"RoPE theta values do not match",
121+
)
122+
self.assertEqual(
123+
hf_config.sliding_window,
124+
backbone.sliding_window,
125+
"Sliding window values do not match",
126+
)
127+
128+
# 9. Test tokenizer functionality
129+
# Verify the HF tokenizer loads and produces output. A direct token-ID
130+
# comparison between KerasHub and HF is fragile: newer versions of
131+
# LlamaTokenizer prepend a leading-space prefix (▁) before each word,
132+
# causing small test-vocab words to become UNK (id=0) on one side but
133+
# not the other. We skip the strict equality check and instead drive
134+
# the model forward pass with KerasHub token IDs, which are guaranteed
135+
# to be valid indices for the exported weight matrix.
136+
test_text = "the quick brown fox"
137+
keras_tokens = tokenizer(test_text)
138+
# Smoke-check: the HF tokenizer must at least produce some output.
139+
hf_tokens_check = hf_tokenizer.encode(
140+
test_text, add_special_tokens=False
141+
)
142+
self.assertGreater(
143+
len(hf_tokens_check),
144+
0,
145+
"HF tokenizer produced empty output",
146+
)
147+
148+
# 10. Test model inference - verify shapes match
149+
# ops.convert_to_numpy works across all Keras backends (torch/MPS,
150+
# JAX, TensorFlow) unlike backend-specific calls like .cpu().numpy().
151+
keras_tokens_list = ops.convert_to_numpy(keras_tokens).tolist()
152+
input_ids = torch.tensor([keras_tokens_list], dtype=torch.long)
153+
hf_outputs = hf_model(input_ids=input_ids, return_dict=True)
154+
155+
# Check output shape
156+
self.assertEqual(
157+
hf_outputs.logits.shape[-1],
158+
backbone.vocabulary_size,
159+
"Output vocabulary size does not match",
160+
)

0 commit comments

Comments
 (0)