news 2026/10/10 5:10:23

Lit-LLaMA TPU 支持实战指南:在 Google Cloud TPU v4 上运行 LLaMA 推理

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
Lit-LLaMA TPU 支持实战指南:在 Google Cloud TPU v4 上运行 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.

项目地址:https://gitcode.com/gh_mirrors/li/lit-llama
点击查看免费下载

本文基于 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)。

这一架构有两个直接后果:

  1. 首次调用慢,后续调用快:XLA 采用惰性执行(lazy execution),第一次运行时需要把算子图编译成 TPU 可执行程序,之后同一形状的计算直接复用编译产物。这正是原文档中"首次生成约需 20 秒、后续约 5 秒"的原因。
  2. 代码中必须显式处理 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.0TPU v4 运行时镜像预装了 TPU v4 驱动与 PyTorch 2.0 的官方运行时
--accelerator-type=v4-8v4-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=1
  • PJRT_DEVICE=TPU:告诉 PjRT 使用 TPU 设备插件,确保 XLA 将编译好的图派发到 TPU 硬件执行;
  • ALLOW_MULTIPLE_LIBTPU_LOAD=1:允许libtpu库被多次加载,规避多 worker/多进程场景下的加载冲突。

这两行写入 shell 会话即可立即生效;若希望每次登录自动生效,可追加到~/.bashrc后重新登录。

准备模型权重

由于 TPU VM 是新建的空机器,需要把 LLaMA 权重导入,原文档给出两条途径:

  1. 使用gcloud compute tpus tpu-vm scp将本机已有的权重直接拷贝进 VM;
  2. 遵循 权重下载指南:在 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_samples1生成的样本数量
--max_new_tokens50每个样本最多生成的新 token 数
--top_k200采样时仅从概率最高的 k 个 token 中抽取
--temperature0.8采样随机性控制,值越大随机性越高
--checkpoint_pathcheckpoints/lit-llama/7B/lit-llama.pth模型权重路径
--tokenizer_pathcheckpoints/lit-llama/tokenizer.modeltokenizer 路径

结合源码可以还原 TPU 上的生成流程(见 generate.py):

  1. 编码提示词后,将输入填充为最终长度T_new的连续张量,并维护input_pos记录当前位置;
  2. 每次迭代仅对最新位置的 token 做前向,配合model.reset_cache()清理 KV 缓存;
  3. 使用top_k裁剪后做 softmax 与torch.multinomial采样,生成下一个 token;
  4. 每轮循环前后调用xm.mark_step(),让 XLA 推进图的执行;
  5. 命中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.

项目地址:https://gitcode.com/gh_mirrors/li/lit-llama
点击查看免费下载
上一篇:终极Office文档安全分析:oletools完整指南与高效应用
下一篇:django-simple-history用户跟踪终极教程:自动记录操作者信息

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

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/10/10 5:09:27

Spring Boot中JSONPath实战:优雅解析嵌套JSON与第三方接口数据

接手一个跨境商城项目时&#xff0c;最让我头疼的不是业务逻辑&#xff0c;而是第三方接口返回的那一大坨嵌套 JSON —— 订单信息、商品快照、支付流水、物流轨迹全揉在一起&#xff0c;层级深得离谱。为了从里面抠出一个状态码或者金额&#xff0c;我写过一堆JSONObject.getJ…

作者头像 李华
网站建设 2026/10/10 5:08:30

stm32cubemx 固件 FW-F4 V1.28.0离线安装教程

由于现在stm32cubemx 下载需要myST登录&#xff0c;但是注册myST又经常无反应&#xff0c;所以我就找了STM32Cube FW-F4 V1.28.0固件版本&#xff0c;进行本地安装&#xff0c;以下是本地安装教程及固件下载路径安装流程1&#xff0c;以管理员身份打开CUBE,单击INSTALL/REMOVE2…

作者头像 李华
网站建设 2026/10/10 5:06:29

【通信原理笔记】【一】确定信号分析——1.6 频带信号的复包络

文章目录前言一、频带信号的复包络二、频带信号的三种表示三、等效基带分析总结前言 上一篇我们学习了解析信号&#xff0c;它将信号的负频率部分镜像叠加到正频率部分便于分析。然而&#xff0c;频带信号有着不同的载频&#xff0c;分析起来还是不够方便&#xff0c;这篇我们…

作者头像 李华