如何把 Llama 的 PyTorch 权重转换为 MLX 格式并加载运行推理
2026/9/13 9:06:52 网站建设 项目流程

如何把 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.loadmlx.utils.tree_unflatten把权重灌入用mlx.nn实现的 Llama 模型,最后调用推理脚本逐 token 生成文本。整个过程的前提来自 LLM inference 示例文档:你必须已经拿到原始 Llama 权重,因为该示例面向官方权重,不覆盖从 Hugging Face 下载等其他来源。

准备环境

从 Build and Install 文档看,在 Apple silicon 上从 PyPI 安装 MLX 需要:

  • Apple silicon 设备;
  • 原生(native)Python >= 3.10,uname -p应为arm
  • macOS >= 14.0。

安装命令:

pip install mlx

如果系统版本满足要求但 pip 找不到匹配的发行版,文档给出的判断方法是检查python -c "import platform; print(platform.processor())"的输出:应为arm,若是i386说明你在用非原生(Rosetta)Python,需要换回原生 Python。

转换脚本本身依赖 PyTorch 和 NumPy(脚本里直接import torchimport numpy as np并调用torch.load读取权重文件),所以执行转换的机器上这两个库也要可用。注意:转换脚本读取的是 PyTorch 的权重文件,转换完成后推理端只需要 MLX,PyTorch 不再参与推理。

编写权重转换脚本

LLM inference 示例给出了完整的转换脚本。它的核心是一个map_torch_to_mlx函数,负责把 PyTorch 侧的键名改成示例中Llama模型(由LlamaAttentionLlamaEncoderLayer组合,全部用mlx.nn.Linearnn.RoPEnn.RMSNorm实现)所期望的键名:

  • tok_embeddingembedding.weight
  • norm的键:attention_normnorm1ffn_normnorm2
  • wq/wk/wv/woquery_proj/key_proj/value_proj/out_proj
  • FFN 三个矩阵:feed_forward.w1linear1feed_forward.w3linear2feed_forward.w2linear3
  • outputout_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(description="Convert 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 文档的完整LlamaAttentionLlamaEncoderLayerLlama代码。

【免费下载链接】mlxMLX: An array framework for Apple silicon项目地址: https://gitcode.com/GitHub_Trending/ml/mlx

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询