
如何把 Llama 的 PyTorch 权重转换为 MLX 格式并加载运行推理【免费下载链接】mlxMLX: An array framework for Apple silicon项目地址: https://gitcode.com/GitHub_Trending/ml/mlx如果你手里有 Llama 的 PyTorch 权重以及随附的 SentencePiece 分词模型想在自己 Apple silicon 的 Mac 上用 MLX 跑 Llama 文本生成需要三步把 PyTorch 的权重键名与存储格式转换成 MLX 可直接加载的 NPZ 文件用mx.load加mlx.utils.tree_unflatten把权重灌入用mlx.nn实现的 Llama 模型最后调用推理脚本逐 token 生成文本。整个过程的前提来自 LLM inference 示例文档你必须已经拿到原始 Llama 权重因为该示例面向官方权重不覆盖从 Hugging Face 下载等其他来源。准备环境从 Build and Install 文档看在 Apple silicon 上从 PyPI 安装 MLX 需要Apple silicon 设备原生nativePython 3.10uname -p应为armmacOS 14.0。安装命令pip install mlx如果系统版本满足要求但 pip 找不到匹配的发行版文档给出的判断方法是检查python -c import platform; print(platform.processor())的输出应为arm若是i386说明你在用非原生RosettaPython需要换回原生 Python。转换脚本本身依赖 PyTorch 和 NumPy脚本里直接import torch与import numpy as np并调用torch.load读取权重文件所以执行转换的机器上这两个库也要可用。注意转换脚本读取的是 PyTorch 的权重文件转换完成后推理端只需要 MLXPyTorch 不再参与推理。编写权重转换脚本LLM inference 示例给出了完整的转换脚本。它的核心是一个map_torch_to_mlx函数负责把 PyTorch 侧的键名改成示例中Llama模型由LlamaAttention、LlamaEncoderLayer组合全部用mlx.nn.Linear、nn.RoPE、nn.RMSNorm实现所期望的键名tok_embedding→embedding.weight含norm的键attention_norm→norm1ffn_norm→norm2wq/wk/wv/wo→query_proj/key_proj/value_proj/out_projFFN 三个矩阵feed_forward.w1→linear1feed_forward.w3→linear2feed_forward.w2→linear3output→out_proj键名含rope的项直接丢弃return None, None因为示例模型用nn.RoPE在线计算位置编码不需要预存的 RoPE 参数。脚本入口用 argparse 接收两个位置参数PyTorch 权重文件路径和输出 NPZ 文件路径然后torch.load读入、np.savez写出import argparse from itertools import starmap import numpy as np import torch def map_torch_to_mlx(key, value): if tok_embedding in key: key embedding.weight elif norm in key: key key.replace(attention_norm, norm1).replace(ffn_norm, norm2) elif wq in key or wk in key or wv in key or wo in key: key key.replace(wq, query_proj) key key.replace(wk, key_proj) key key.replace(wv, value_proj) key key.replace(wo, out_proj) elif w1 in key or w2 in key or w3 in key: # The FFN is a separate submodule in PyTorch key key.replace(feed_forward.w1, linear1) key key.replace(feed_forward.w3, linear2) key key.replace(feed_forward.w2, linear3) elif output in key: key key.replace(output, out_proj) elif rope in key: return None, None return key, value.numpy() if __name__ __main__: parser argparse.ArgumentParser(descriptionConvert Llama weights to MLX) parser.add_argument(torch_weights) parser.add_argument(output_file) args parser.parse_args() state torch.load(args.torch_weights) np.savez( args.output_file, **{k: v for k, v in starmap(map_torch_to_mlx, state.items()) if k is not None} )把它保存为convert.py后按它的 argparse 接口传两个位置参数运行下例中llama-7B/为文档示例使用的权重目录名llama.npz为你指定的输出文件实际按你的权重位置替换python convert.py torch_weights llama.npz输出的.npz文件就是 MLX 可直接加载的权重格式——Saving and Loading 文档说明mx.load按文件扩展名识别格式加载.npz时返回名称到数组的字典这正好是下一步需要的输入。把 NPZ 权重加载进模型推理脚本先用mlx.nn定义出结构相同的模型embedding、若干LlamaEncoderLayer、RMSNorm 和输出投影然后从磁盘读权重并整体更新。文档给出的加载代码是from mlx.utils import tree_unflatten model.update(tree_unflatten(list(mx.load(weight_file).items())))其中mx.load(weight_file)读 NPZ 得到键值字典tree_unflatten把形如layers.2.attention.query_proj.weight的扁平键转回嵌套结构例如{layers: [..., ..., {attention: {query_proj: {weight: ...}}}]}再交给model.update灌入各参数。文档同时提醒这条路径存在从磁盘到 NumPy、再从 NumPy 到 MLX 的几次额外拷贝未来会被直接加载到 MLX 的实现取代。生成侧则是一个 Python 生成器先处理整个 prompt 并保存每层的 key/value 缓存再自回归地逐个yieldtoken采样用mx.random.categorical(y * (1/temp))。由于 MLX 是惰性求值model.generate返回的每个y在真正mx.eval、拼接或打印之前并不会计算你可以选择何时触发实际计算。运行推理并核对输出文档以本地已存在 PyTorch Llama 权重目录llama-7B/为例展示运行方式完整示例代码在官方mlx-examples仓库的llms/llama目录其中convert.py的命令行接口为--torch-path形式与上文内嵌脚本的位置参数接口略有差异python convert.py --torch-path llama-7B/ python llama.py --prompt Call me Ishmael. Some years ago never mind how long precisely文档示例的运行输出M1 Ultra、7B 模型仅作为示例不是每次运行的固定数值[INFO] Loading model from disk: 5.247 s Press enter to start generation ------ , having little or no money in my purse, and nothing of greater consequence in my mind, ... ------ [INFO] Prompt processing: 0.437 s [INFO] Full generation: 4.330 s可以据此核对三个信号权重加载耗时正常打印、prompt 处理耗时明显小于整段生成耗时、生成的文本连贯。文档据此统计4.3 秒生成 100 个 token其中 0.4 秒处理 prompt约合每 token 39 ms换更长的 prompt 再跑每 token 生成时间与 prompt 处理时间几乎保持不变这是该文档给出的扩展验证方式。生成 token 数用--max-tokens控制例如python llama.py --max-tokens 500 --prompt ...限制与注意事项模型必须与键名映射匹配map_torch_to_mlx是针对示例中Llama类结构写的nn.Linear投影 nn.RoPE SwiGLU 的 FFN。如果你自己改过模型结构键名必须同步改否则model.update拿不到对应参数。权重来源该示例的前提是你已有原始 Llama 权重和 SentencePiece 模型文档不覆盖权重的下载与许可问题。拷贝开销目前mx.load后走 NumPy 再转 MLX存在文档明确指出的额外拷贝文档说明未来会改为直接加载到 MLX。调试惰性计算如果生成卡住没有结果通常是没有触发mx.eval或打印等求值操作——MLX 的数组在求值前只是计算图不是结果。更多模型实现细节attention 缓存拼接、RMSNorm 与 SwiGLU 的具体写法可参考 LLM inference 文档的完整LlamaAttention、LlamaEncoderLayer与Llama代码。【免费下载链接】mlxMLX: An array framework for Apple silicon项目地址: https://gitcode.com/GitHub_Trending/ml/mlx创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考