Files

40 lines
1.2 KiB
Python
Raw Permalink Normal View History

"""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