Merge nucleic/sleek-ember-seal-uady into dev

This commit is contained in:
2026-07-30 20:14:37 -07:00
parent 09f98cdd00
commit 90092332db
8 changed files with 1668 additions and 0 deletions
+57
View File
@@ -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()
+72
View File
@@ -0,0 +1,72 @@
import json
import sys
import tempfile
import unittest
from pathlib import Path
import numpy as np
import torch
MODULE_DIR = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(MODULE_DIR))
import train
import train_mlx
class FixedShapeTokenizerTests(unittest.TestCase):
class Tokenizer:
pad_token_id = 0
cls_token_id = 1
sep_token_id = 2
padding_side = "right"
model_input_names = ["input_ids", "attention_mask", "token_type_ids"]
def __call__(self, texts, **_):
return {
"input_ids": [
list(range(10, 10 + int(text.split()[-1]))) for text in texts
]
}
@staticmethod
def num_special_tokens_to_add(pair=False):
return 3 if pair else 2
def test_mlx_encoding_matches_pytorch_encoding(self):
texts = ["tokens 3", "tokens 200"]
pytorch = train.encode_fixed_shape(self.Tokenizer(), texts, torch)
mlx = train_mlx.encode_fixed_shape_numpy(self.Tokenizer(), texts)
self.assertEqual(set(pytorch), set(mlx))
for key in pytorch:
with self.subTest(key=key):
np.testing.assert_array_equal(pytorch[key].numpy(), mlx[key])
class CheckpointConfigTests(unittest.TestCase):
def test_rejects_changed_label_order(self):
config = {
"model_type": "bert",
"hidden_size": 384,
"num_hidden_layers": 6,
"id2label": {
str(index): label
for index, label in enumerate(reversed(train.LABELS))
},
}
with tempfile.TemporaryDirectory() as temp:
model_dir = Path(temp)
(model_dir / "config.json").write_text(
json.dumps(config),
encoding="utf-8",
)
with self.assertRaisesRegex(
train.DataError,
"label order",
):
train_mlx._checkpoint_config(model_dir)
if __name__ == "__main__":
unittest.main()