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