- 人工智能
- 大模型
- 预训练
- 微调
- LoRA
- 模型量化
【免费下载链接】lit-llama
Implementation of the LLaMA language model based on nanoGPT. Supports flash attention, Int8 and GPTQ 4bit quantization, LoRA and LLaMA-Adapter fine-tuning, pre-training. Apache 2.0-licensed.
本文基于 howto/tpus.md 编写,完整讲解如何在 Google Cloud 上创建 TPU v4 虚拟机、安装 Lit-LLaMA 依赖、配置 PjRT 运行时,并直接运行 LLaMA 文本生成。读完本文,你将掌握一条从零到一在 TPU 上跑通 Lit-LLaMA 推理的完整命令链路,并理解 XLA 图编译、mark_step与缓存重置等底层机制对推理性能的影响。
TPU 支持的底层机制:lightning.Fabric + PyTorch XLA
Lit-LLaMA 的代码统一构建在lightning.Fabric之上,而 Fabric 本身通过PyTorch XLA提供对 TPU 的支持。也就是说,项目的训练、生成脚本并不直接感知 TPU 设备,而是由 Fabric 负责把张量放到 XLA 设备上,再由 PyTorch XLA 将模型计算编译为可在 TPU 上高效执行的图(graph)。
这一架构有两个直接后果:
- 首次调用慢,后续调用快:XLA 采用惰性执行(lazy execution),第一次运行时需要把算子图编译成 TPU 可执行程序,之后同一形状的计算直接复用编译产物。这正是原文档中"首次生成约需 20 秒、后续约 5 秒"的原因。
- 代码中必须显式处理 XLA 执行边界:源码里可以看到多处针对
xla设备类型的专门分支。
例如在 generate.py 中,生成循环前和每次迭代后都会调用torch_xla.core.xla_model.mark_step(),强制 XLA 执行当前累积的图并同步设备:
if idx.device.type == "xla": import torch_xla.core.xla_model as xm xm.mark_step()又比如 lit_llama/model.py 中的reset_cache,在 XLA 设备上除了清空 KV 缓存外,还会把rope_cache和mask_cache一并置空:
def reset_cache(self) -> None: self.kv_caches.clear() if self.mask_cache.device.type == "xla": # https://github.com/Lightning-AI/lit-parrot/pull/83#issuecomment-1558150179 self.rope_cache = None self.mask_cache = None原因是 XLA 编译的图对张量形状敏感,跨样本复用长度相关的缓存可能导致形状不匹配;在 TPU 上每次生成前重置缓存是必要的。
在 Google Cloud 创建 TPU v4 虚拟机
原文档给出了两条gcloud命令,即可创建一台带 TPU v4 的虚拟机并登录:
gcloud compute tpus tpu-vm create lit-llama --version=tpu-vm-v4-pt-2.0 --accelerator-type=v4-8 --zone=us-central2-b gcloud compute tpus tpu-vm ssh lit-llama --zone=us-central2-b逐参数拆解如下:
| 参数 | 取值 | 含义 |
|---|---|---|
lit-llama | 实例名 | TPU VM 的标识,后续 SSH、删除均要引用它 |
--version=tpu-vm-v4-pt-2.0 | TPU v4 运行时镜像 | 预装了 TPU v4 驱动与 PyTorch 2.0 的官方运行时 |
--accelerator-type=v4-8 | v4-8 | 单 Pod 上挂载 4 个 TPU v4 芯片(共 8 个 core),足以容纳 7B 权重并跑推理 |
--zone=us-central2-b | 区域 | TPU v4 所在可用区,创建与后续 SSH/删除命令必须保持一致 |
执行前请确保已安装并认证gcloud(gcloud auth login),且当前 GCP 项目已启用 Cloud TPU API、具备相应配额。原文档还提示,关于 TPU v4 的完整开通说明与全部可用选项,可参考官方提供的 TPU v4 用户指南。
克隆仓库并安装依赖
进入虚拟机后,克隆仓库并安装依赖:
git clone https://gitcode.com/gh_mirrors/li/lit-llama cd lit-llama pip install -e ".[all]".[all]安装的是 pyproject.toml 中定义的完整依赖集合,核心依赖包括:
torch>=2.0.0:与tpu-vm-v4-pt-2.0镜像内置的 PyTorch 2.0 版本匹配;lightning(master 分支):提供 Fabric,即 TPU 支持得以实现的抽象层;sentencepiece:加载 LLaMA 的tokenizer.model;bitsandbytes:服务于llm.int8量化路径。
[project.optional-dependencies] all额外引入tqdm、numpy <2.0、jsonargparse[signatures]、datasets、zstandard等,分别用于权重转换、数据加载、CLI 参数解析与 RedPajama 数据准备。按.[all]安装即可覆盖推理与后续可能用到的全部脚本。
配置 PjRT 运行时环境变量
PyTorch XLA 自 2.0 起默认使用新的PjRT(Pluginable JAX Runtime)运行时。原文档指出该运行时目前仍标记为"experimental",因此建议显式设置以下两个环境变量:
export PJRT_DEVICE=TPU export ALLOW_MULTIPLE_LIBTPU_LOAD=1PJRT_DEVICE=TPU:告诉 PjRT 使用 TPU 设备插件,确保 XLA 将编译好的图派发到 TPU 硬件执行;ALLOW_MULTIPLE_LIBTPU_LOAD=1:允许libtpu库被多次加载,规避多 worker/多进程场景下的加载冲突。
这两行写入 shell 会话即可立即生效;若希望每次登录自动生效,可追加到~/.bashrc后重新登录。
准备模型权重
由于 TPU VM 是新建的空机器,需要把 LLaMA 权重导入,原文档给出两条途径:
- 使用
gcloud compute tpus tpu-vm scp将本机已有的权重直接拷贝进 VM; - 遵循 权重下载指南:在 VM 内下载 Meta 原始权重或 OpenLLaMA 权重,再用
scripts/convert_checkpoint.py(原始权重)或scripts/convert_hf_checkpoint.py(HuggingFace 格式)转换为 Lit-LLaMA 的lit-llama.pth格式,最终得到类似checkpoints/lit-llama/7B/lit-llama.pth与tokenizer.model的目录结构。
转换后的默认路径(checkpoints/lit-llama/7B/lit-llama.pth与checkpoints/lit-llama/tokenizer.model)与 generate.py 中的默认参数一致,无需额外指定即可运行。若想自定义目录,可参考 路径定制指南:所有脚本均支持-h查看可选项,并通过--checkpoint_path、--tokenizer_path显式传参。
在 TPU 上运行推理
权重就绪后,推理开箱即用:
python3 generate.py --prompt "Hello, my name is" --num_samples 3该命令以Hello, my name is为提示词,连续生成 3 段文本。原文档给出的实测现象为:首次生成约需 20 秒(XLA 编译图的开销),之后每次生成回落到约 5 秒。
generate.py的核心参数如下(与 generate.py 中的函数签名一致):
| 参数 | 默认值 | 说明 |
|---|---|---|
--prompt | "Hello, my name is" | 生成所用的提示词 |
--num_samples | 1 | 生成的样本数量 |
--max_new_tokens | 50 | 每个样本最多生成的新 token 数 |
--top_k | 200 | 采样时仅从概率最高的 k 个 token 中抽取 |
--temperature | 0.8 | 采样随机性控制,值越大随机性越高 |
--checkpoint_path | checkpoints/lit-llama/7B/lit-llama.pth | 模型权重路径 |
--tokenizer_path | checkpoints/lit-llama/tokenizer.model | tokenizer 路径 |
结合源码可以还原 TPU 上的生成流程(见 generate.py):
- 编码提示词后,将输入填充为最终长度
T_new的连续张量,并维护input_pos记录当前位置; - 每次迭代仅对最新位置的 token 做前向,配合
model.reset_cache()清理 KV 缓存; - 使用
top_k裁剪后做 softmax 与torch.multinomial采样,生成下一个 token; - 每轮循环前后调用
xm.mark_step(),让 XLA 推进图的执行; - 命中
eos_id则提前截断返回,否则生成满max_new_tokens个 token。
脚本还会逐样本打印Time for inference与tokens/sec统计,方便在 TPU 上直接验证吞吐表现(见 generate.py)。
微调支持状态
截至原文档编写时,TPU 上的微调标记为 "Coming soon",即仓库尚未提供在 TPU 上运行微调的官方教程与验证。目前 finetune 目录下的lora.py、adapter.py、full.py等脚本主要面向 GPU 环境(如 README 中说明的 LoRA/Adapter 微调需要约 24 GB 显存)。在 TPU 上自行尝试微调前,建议先确认对应脚本在 XLA 设备上的算子兼容性。
使用完毕:删除实例
Cloud TPU 是按使用计费的托管服务,原文档特别提醒:结束后务必删除实例,避免产生持续费用:
gcloud compute tpus tpu-vm delete lit-llama --zone=us-central2-b删除命令与创建命令一样,需要带上正确的--zone。
注意事项与限制小结
- PjRT 仍处于实验阶段:尽管它已成为 PyTorch XLA 2.0 的默认运行时,官方仍建议显式设置
PJRT_DEVICE与ALLOW_MULTIPLE_LIBTPU_LOAD两个环境变量; - 首次编译延迟不可避免:约 20 秒的首个样本延迟来自 XLA 图编译,属于预期行为,后续样本即恢复约 5 秒;
- 区域与镜像保持一致:
create、ssh、delete三条命令的--zone=us-central2-b必须统一,镜像版本以tpu-vm-v4-pt-2.0(PyTorch 2.0)为准,与仓库torch>=2.0.0的依赖要求吻合; - TPU 上不要复用 GPU 的假设:例如精度选择,generate.py 仅在检测到 CUDA 且支持 bf16 时才使用
bf16-true,TPU 环境下会回退到32-true,属于脚本的既定行为,无需干预。
至此,你已经掌握了从创建 TPU v4 虚拟机、安装 Lit-LLaMA、配置 PjRT 环境、导入权重,到最终在 TPU 上完成 LLaMA 文本生成的完整链路。更多推理细节可继续阅读 推理指南。
- 人工智能
- 大模型
- 预训练
- 微调
- LoRA
- 模型量化
【免费下载链接】lit-llama
Implementation of the LLaMA language model based on nanoGPT. Supports flash attention, Int8 and GPTQ 4bit quantization, LoRA and LLaMA-Adapter fine-tuning, pre-training. Apache 2.0-licensed.
相关推荐
20分钟搞定BLIP2-OPT-2.7B环境配置:Windows/Linux/Mac全平台避坑指南
20分钟搞定BLIP2 OPT 2.7B环境配置:Windows/Linux/Mac全平台避坑指南 你是否曾因环境配置失败放弃AI视觉项目?是否在CUDA、Py
人工智能大模型预训练微调LoRA模型量化SkyPilot 上的 Cloud TPU v6e(Trillium)实战:一键创建、Llama 3 8B 训练与 JetStream 推理服务
SkyPilot 上的 Cloud TPU v6e(Trillium)实战:一键创建、Llama 3 8B 训练与 JetStream 推理服务 本指南以仓库
后端任务调度MLOps集群管理
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考