58 lines
1.9 KiB
Python
58 lines
1.9 KiB
Python
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()
|