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()