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