Gemma模型在TPU环境模仿CUDA但显存不足提示的异质内存配置方法

文章导读
TPU 的显存不足提示通常来自 XLA 的 Allocator,而不是 CUDA。绝大多数情况下,异质内存配置不是单一开关,而是组合使用半精度、分片加载和手动权重搬运。如果 Gemma 模型在 TPU 上直接加载失败,建议先确认 TPU 类型和主机内存,再按下面步骤逐步调整。
📋 目录
  1. A 前置条件与问题确认
  2. B 配置与安装:三种方式组合
  3. C 验证配置是否生效
  4. D 失败时回退方案
A A

TPU 的显存不足提示通常来自 XLA 的 Allocator,而不是 CUDA。绝大多数情况下,异质内存配置不是单一开关,而是组合使用半精度、分片加载和手动权重搬运。如果 Gemma 模型在 TPU 上直接加载失败,建议先确认 TPU 类型和主机内存,再按下面步骤逐步调整。

前置条件与问题确认

开始配置前,先确认三个信息:TPU 类型(v3、v4 或 v5)、JAX/PyTorch-XLA 版本、主机内存是否足够。在 JAX 环境中执行以下代码确认设备数量与名称:

import jax, jax.numpy as jnp
print(jax.devices())
print(jax.local_device_count())

如果提示信息中同时出现"Memory allocation failed"和"HBM"字样,说明 TPU 内存耗尽;如果只出现"System RAM",则是主机内存不足。异质内存配置针对的是第一种情况,目标是尽量把部分权重和中间结果放到主机内存。

配置与安装:三种方式组合

方式一:强制 bfloat16 权重压缩

Gemma 模型参数为 bfloat16 时,内存占用约为 float32 的一半。在加载前设置环境变量通常最直接:

Gemma模型在TPU环境模仿CUDA但显存不足提示的异质内存配置方法
# 在运行脚本前设置,适用于 JAX 与 PyTorch/XLA
export XLA_USE_BF16=1

注意:此方式会改变模型精度,可能影响输出质量。建议在推理场景先用一个短样本对比输出。

方式二:使用 pjit 手动分片权重

在 JAX 中,通过 pjit 将模型参数分成多个 PartitionSpec,可以让部分数组驻留在 CPU 上。以下是一个最小骨架,需根据模型结构微调:

from jax.experimental.pjit import pjit
from jax.sharding import Mesh, PartitionSpec
import jax.numpy as jnp

mesh = Mesh(jax.devices(), 'd')

# params 是模型参数,inputs 是输入
# 这里假设把第1层参数放在 host CPU,其他层正常分片
def step(params, inputs):
    # 先执行需要 CPU 的分区计算
    # 返回结果仍需放到 TPU devices
    return ...

executor = pjit(step,
                in_shardings=(PartitionSpec('d', None), PartitionSpec('d')),
                out_shardings=PartitionSpec('d'))

此方式要求理解 XLA 的分片规则,不能照搬到所有模型。如果你使用的是 Hugging Face Transformers 加载 Gemma,通常 device_map 在 TPU 上不可用,需要手动把权重拆分成两批:先计算前半层,再释放 TPU 内存,从 CPU 加载后半层。

方式三:环境变量降低预分配空间

TPU 默认行为是尽可能预先占用大部分 HBM。可尝试设置:

Gemma模型在TPU环境模仿CUDA但显存不足提示的异质内存配置方法
# 动态分配,避免一次性吃满 TPU HBM
export XLA_ALLOCATOR=Platform

该变量不是所有 TPU 运行时都支持,需要结合当前 XLA 版本确认。如果启动后运行报错,回退到不设置。

验证配置是否生效

执行模型推理或训练,观察是否还出现 OOM。同时用 vmstat 1 监控主机内存变化:当 TPU 内存被降低分配后,进程的 RES 内存应上升。更准确的办法是通过 jax.device_info() 查询每个 TPU 设备的峰值内存,但各版本接口有差异,建议用输出日志中的 XLA 内存分配信息判断。

具体验证流程:先用 4 句话的小样本跑一次,记录最大 TPU 内存占用;再按上述方式二把模型中间层放到 CPU,再次运行看是否仍然 OOM。如果内存占用下降但不明显,说明配置有效;如果完全无变化,说明分片代码没有被实际执行。

Gemma模型在TPU环境模仿CUDA但显存不足提示的异质内存配置方法

失败时回退方案

异质内存配置在 TPU 上并不是银弹。如果尝试上述方法后仍然 OOM,回退到以下两条更简单的路径:

  • 降低工作负载:把 batch size 调为 1、缩短序列长度,或者改用 Gemma-2B 而不是 7B。
  • 使用 TPU v3 的 64GB 版本(如果资源允许),或者改用多卡 TPU 通过数据并行分摊单卡内存。

对于推理场景,还可以考虑把模型拆成多个片段,在 CPU 上逐层计算,每计算完一层再搬运到 TPU。虽然速度较慢,但至少能跑通。如果只是为了验证输出质量,优先用半精度和更小 batch,避免过早进入分片调试。

最后,如果你发现只是想把 TPU 当作 GPU 使用,建议先确认自己的安装环境是否正确包含 PJRT TPU 运行时,因为部分"显存不足"提示其实是框架找不到 TPU 时的假错误。检查 PJRT_DEVICE 是否设置为 TPU,并用官方测试脚本确认基础计算能跑通。