Files

58 lines
1.9 KiB
Python
Raw Permalink Normal View History

import sys
import unittest
from pathlib import Path
MODULE_DIR = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(MODULE_DIR))
from mlx_checkpoint import hugging_face_to_mlx_key, mlx_to_hugging_face_key
class MLXCheckpointTests(unittest.TestCase):
def test_representative_bert_keys_round_trip(self):
keys = (
"bert.embeddings.LayerNorm.weight",
"bert.embeddings.word_embeddings.weight",
"bert.encoder.layer.0.attention.self.query.weight",
"bert.encoder.layer.2.attention.self.key.bias",
"bert.encoder.layer.4.attention.self.value.weight",
"bert.encoder.layer.5.attention.output.dense.bias",
"bert.encoder.layer.1.attention.output.LayerNorm.weight",
"bert.encoder.layer.3.intermediate.dense.weight",
"bert.encoder.layer.3.output.dense.bias",
"bert.encoder.layer.3.output.LayerNorm.weight",
"bert.pooler.dense.weight",
"classifier.weight",
)
for key in keys:
with self.subTest(key=key):
self.assertEqual(
key,
mlx_to_hugging_face_key(hugging_face_to_mlx_key(key)),
)
def test_expected_mlx_names(self):
self.assertEqual(
"bert.encoder.layers.0.attention.query_proj.weight",
hugging_face_to_mlx_key(
"bert.encoder.layer.0.attention.self.query.weight"
),
)
self.assertEqual(
"bert.encoder.layers.0.ln1.bias",
hugging_face_to_mlx_key(
"bert.encoder.layer.0.attention.output.LayerNorm.bias"
),
)
self.assertEqual(
"bert.encoder.layers.0.ln2.weight",
hugging_face_to_mlx_key(
"bert.encoder.layer.0.output.LayerNorm.weight"
),
)
if __name__ == "__main__":
unittest.main()