Merge nucleic/sleek-ember-seal-uady into dev
This commit is contained in:
@@ -0,0 +1,57 @@
|
||||
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()
|
||||
Reference in New Issue
Block a user