40 lines
1.2 KiB
Python
40 lines
1.2 KiB
Python
"""Hugging Face <-> MLX parameter-name conversion for purpose-lite BERT."""
|
|||
|
|
|
||
|
|
from __future__ import annotations
|
||
|
|
|
||
|
|
|
||
|
|
_HF_TO_MLX_REPLACEMENTS = (
|
||
|
|
(".layer.", ".layers."),
|
||
|
|
(".self.key.", ".key_proj."),
|
||
|
|
(".self.query.", ".query_proj."),
|
||
|
|
(".self.value.", ".value_proj."),
|
||
|
|
(".attention.output.dense.", ".attention.out_proj."),
|
||
|
|
(".attention.output.LayerNorm.", ".ln1."),
|
||
|
|
(".output.LayerNorm.", ".ln2."),
|
||
|
|
(".intermediate.dense.", ".linear1."),
|
||
|
|
(".output.dense.", ".linear2."),
|
||
|
|
(".embeddings.LayerNorm.", ".embeddings.norm."),
|
||
|
|
(".pooler.dense.", ".pooler."),
|
||
|
|
)
|
||
|
|
|
||
|
|
_MLX_TO_HF_REPLACEMENTS = tuple(
|
||
|
|
(mlx, hugging_face) for hugging_face, mlx in reversed(_HF_TO_MLX_REPLACEMENTS)
|
||
|
|
)
|
||
|
|
|
||
|
|
|
||
|
|
def hugging_face_to_mlx_key(key: str) -> str:
|
||
|
|
"""Return the MLX BERT parameter name corresponding to a Transformers key."""
|
||
|
|
|
||
|
|
for hugging_face, mlx in _HF_TO_MLX_REPLACEMENTS:
|
||
|
|
key = key.replace(hugging_face, mlx)
|
||
|
|
return key
|
||
|
|
|
||
|
|
|
||
|
|
def mlx_to_hugging_face_key(key: str) -> str:
|
||
|
|
"""Return the Transformers parameter name corresponding to an MLX BERT key."""
|
||
|
|
|
||
|
|
for mlx, hugging_face in _MLX_TO_HF_REPLACEMENTS:
|
||
|
|
key = key.replace(mlx, hugging_face)
|
||
|
|
return key
|
||
|
|
|