TensorFlow XLA:TPU 编译报 "Ran out of memory in memory space hbm" 怎么排查?
【免费下载链接】tensorflowAn Open Source Machine Learning Framework for Everyone项目地址: https://gitcode.com/GitHub_Trending/te/tensorflow
当 JAX/XLA:TPU 程序在编译阶段抛出类似下面的报错时,说明程序需要的静态分配总量超过了 TPU 芯片物理 HBM(High Bandwidth Memory)容量,这是 XLA 错误分类中的 **E1000: Compile Time: HBM OOM**(后端:TPU):
RESOURCE_EXHAUSTED: XLA:TPU compile permanent error. Ran out of memory in memory space hbm. Used 49.34G of 32.00G hbm. Exceeded hbm capacity by 17.34G.(以上是 E1000 文档 中的示例错误消息。)
XLA 在编译时会检查所有必要静态分配的总和是否能放进设备 HBM。编译器管理的 TPU HBM 分配包括六类:程序输入输出(训练批次、优化器状态等)、TPU temporaries(激活、梯度等中间计算)、编译产物(TensorCore 和 SparseCore 的机器码)、系统开销(XLA Runtime 预留空间)、常量(内嵌在 HLO IR 中的常量)、编译器内部分配(如 mesh 中节点的路由信息)。当这六类的总和放不进 HBM 时,就会报出这个错。
排查路径在 XLA 仓库文档中是明确的:E1000 文档 负责"按错误消息分诊",Debug OOM errors with XProf 负责"用 XProf Memory Viewer 定位峰值内存",XLA flags guidance 提供最后的内存 flag 调优项。下面按这条路径展开。
第一步:先看错误消息属于哪种形态
error_1000.md 要求"仔细分析错误消息和日志",然后进入对应分支:
- 错误明确给出了 TC/SC 用量分解,形如
TC Hbm usage: X, SC Hbm usage Y(示例:TPU TensorCore Hbm usage: 34.82G, SparseCore Hbm usage 174.10G, exceeding available bytes: 95.74G)→ 进入"TC/SC 失衡"分支; - 错误是
Ran out of memory in memory space hbm,且日志中列出了异常大的分配(单个张量超过 HBM 上限的 50%)→ 进入"大分配"分支; - 错误是
Ran out of memory in memory space hbm,但日志中没有异常大的张量 → 进入"累积压力"分支,需要用 XProf 可视化峰值内存。
分支一:错误显示了 TC/SC 用量分解
此时是 TensorCore(TC)+ SparseCore(SC)的总用量超过了 HBM 上限,对比两个数值找出瓶颈:
- SparseCore 用量高时,文档给出的检查项:
- HBM stack 用量随
feature_width、max_unique_nz_per_row和logical_replica_count增长。可以用--xla_sc_num_serialized_tables_to_optimize_hbmflag 把 table 的处理串行化以降低峰值 stack 用量,代价是并行度下降; - 检查 padding 开销:SparseCore 会把 embedding table 对齐到 32B(8 个 float)。feature width 较小的表(例如 < 8 个 float)会产生显著 padding 浪费;
maximum_parallel_iterations取值过大会把更多输入数据预取进 HBM heap,调低该值可以释放内存;- 确认 embedding table 是否在所有 chip 之间正确做了 mod sharding。
- HBM stack 用量随
- TensorCore 用量高:转到大分配分支(分支二)继续排查。
- 两者都不高但总和超限:说明已经到了芯片容量上限,需要同时降低两个组件的用量,按分支二、三的建议综合处理。
分支二:日志里有异常大的分配(> 50% HBM 上限)
E1000 文档明确指出:出现这种大分配时"几乎从来不是硬件容量问题,通常是配置错误"。具体检查:
- 查看大分配的 XLA label(如果存在),label 里通常有指向 JAX 源码位置的提示;
- 移除调试残留:在大规模运行里使用
jax.debug.print()会强制编译器把完整张量实体化到 HBM 再传回 CPU,破坏融合并抬高峰值内存。删掉遗留的jax.debug.print(); - 修正 mesh shape 或 sharding 标注:错误的 mesh shape 或缺失的 sharding 标注会让编译器退化为replication——把非常大的张量整个塞进单块芯片。检查大分配的 shape,确认 sharding 被正确指定并被 XLA 传播。
分支三:没有单一大分配——用 XProf 定位峰值内存
当总分配超限但没有明显大张量时,需要先看"峰值时刻到底是谁占着 HBM"。oom_debugging.md 给出的完整流程:
给程序加 profiling trace。文档示例(文档示例)中,一个触发 OOM 的 JAX 程序这样写:
import jax from jax import random import jax.numpy as jnp @jax.profiler.trace("/tmp/xprof") @jax.jit def oom(): a = random.normal(random.PRNGKey(1), (327680, 327680), dtype=jnp.bfloat16) return a @ a if __name__ == "__main__": oom()在你的程序中,把
jax.profiler.trace装饰到需要捕获的入口函数上,第一个参数是 profile 存储目录。文档特别建议用jax.profiler.trace而不是jax.profiler.start_trace/stop_trace,因为前者是上下文管理器,在异常情况下也能安全结束 profiling。安装并启动 XProf,指定 profile 目录和端口:
pip install xprof xprof --logdir=/tmp/xprof/ --port=6006打开 Memory Viewer。本地机器上访问
http://localhost:6006,在Tools下拉框选择Memory Viewer,在Memory Types下拉框选择HBM(通常默认已选中):看 "HLO Ops at Peak Memory Allocation" 区块。该区块展示峰值内存使用点的 buffer 图,buffer 包括:
- Program Inputs and Outputs:训练批次、优化器状态等;
- TensorCore and SparseCore Temporaries:中间计算(激活、梯度等)所需的动态内存。
鼠标悬停在 buffer 图上可以看到该 Op 的 size、shape、allocation type 等细节,用来识别占用高或生命周期长的 temporaries,以及 padding 低效的大输入/中间/输出张量。
下面按 Memory Viewer 里看到的主力归因,选择文档给出的优化项。
配置层面的调整(往往最先有效)
- 减小 batch size:中间激活和梯度的内存与 batch size 成正比。注意减小 batch size 可能需要同步重调学习率、动量或优化器超参以维持训练稳定性;
- 捐赠输入 buffer:如果某个输入在计算后不再使用,且其 shape 和元素类型与某个输出匹配,可以通过
jax.jit的donate_argnums参数把该输入 buffer 捐给输出,内存减少量约为被捐 buffer 的大小; - 对最大张量启用 bfloat16 或量化(如模型架构和质量要求允许)。这会改变数值行为,需要谨慎评估;
- Micro-batching(可选):当无法减小全局 batch size 或增加芯片数、且单芯片 batch size 已接近下限时,把每个 batch 拆成
n个 micro-batch,逐个跑前向和反向,最后累积梯度并整体更新权重——激活内存从M降到约M/n。文档标注的代价:step 时间变长(多次前向反向),且模型与 micro-batch 尺寸差距过大会带来收敛问题。
架构与 sharding 层面
当配置调整不够时,可能是模型拓扑对当前硬件过大:
- 换用更新一代的 TPU(单芯片 HBM 更大);
- 在更大的芯片拓扑上运行,把权重 shard 到更多芯片;
- 使用更高级的 data/tensor/pipeline 并行,并为中间值和输出指定 sharding hint。注意把张量切到多芯片会带来网络通信开销;
- Host offloading:把大张量(激活、优化器状态)卸载到 host CPU 内存。数值上安全,但文档明确警告会严重影响性能——系统要不断在 TPU HBM 和 CPU RAM 之间搬大张量,属于最后手段。
检查 tensor padding 与对齐
TPU 上的低效形状是 OOM 的常见且隐蔽的成因。为了达到峰值性能,XLA 会把 tensor 维度做 padding——minor-most 维度通常对齐到 128 的倍数,第二小的维度对齐到 8 的倍数。padding 影响输入数组和中间 tensor(HLO temporaries),在小维度上可能显著放大内存用量。
- 在 XProf Memory Viewer 中悬停 buffer 查看详情卡里的 padding 信息(文档以 TPU v5 默认 layout 为例):shape
(129, 1024)可能被 pad 到(256, 1024),产生近 50% 的内存浪费(文档示例);改成(128, 1024)则不需要 padding。 - 把大 tensor 的维度(batch size、embedding 维度、hidden size)调整为 128 的倍数。这会改变模型行为,需要谨慎评估。
Rematerialization 与手动 checkpointing
模型接近能放进内存时,可以用jax.checkpoint装饰器配合jax.grad手动控制哪些中间值在前向保留、哪些在反向重算——用算力换 HBM。
也可以让XLA::Rematerializationpass 优先省内存(代价是编译变慢)。E1000 文档列出的 flag 及取舍:
| Flag | 作用 | 影响 / 取舍 |
|---|---|---|
--xla_tpu_max_hbm_size_mib | 手动设置 Rematerialization pass 使用的 HBM 上限 | 强迫编译器把程序塞进比物理 HBM 更小的限制 |
--xla_tpu_rematerialization_algo=PEAK_PRIORITY | 把优化集中到内存峰值点 | 相比默认算法可能更高效地削减内存 |
--xla_tpu_rematerialization_max_block_size_limit=32 | 控制一次可 rematerialize 的 block 内最大指令数 | 调大可以省更多内存,但显著增加编译时间 |
--xla_tpu_rematerialization_block_effort_factor=10.0 | 定义搜索可 rematerialize block 的编译努力程度 | 值越大搜索越彻底,编译时间越长 |
--xla_tpu_pre_fusion_remat=true | 在 fusion pass 之前增加一次 Rematerialization pass | 能发现更多内存节省,但编译时间增加,且可能影响数值稳定性 |
文档明确提示:修改 XLA flag 应作为最后手段,可能损害性能。
最后一档:XLA Memory Flags
flags_guidance.md 的 Memory Flags 一节 说明:这些 flag 就是为了解决编译期的 HBM OOM 提供的,只在遇到 HBM out of memory 时调整,其余场景保持默认值,改动可能损害性能。其中的默认值与建议值:
| Flag | 默认值 | 文档建议值 |
|---|---|---|
xla_latency_hiding_scheduler_rerun | 1 | 5(每次 rerun 会逐步下调调度内存上限,文档标注超过 10 次意义不大) |
xla_tpu_rwb_fusion | true | false(关闭 reduce+broadcast 融合可降内存) |
xla_memory_scheduler | kDefault | kBrkga(更省内存的调度算法,代价是编译更慢) |
xla_tpu_enable_latency_hiding_scheduler | true | false(以放弃异步 collective 的性能收益换内存) |
xla_jf_spmd_threshold_for_windowed_einsum_mib | -1 | 10Mb~1Gb(提高阈值可省内存,代价是失去 collective matmul 机会) |
如何对照症状选手段
E1000 文档末尾给了一张干预手段速查表(下表为按原文整理,"典型症状"列是文档给出的 telltale signs,帮助你确认当前瓶颈是否对得上):
| 手段 | 是否改变程序行为 | 典型症状(对上再动手) |
|---|---|---|
| 高级 sharding 技术 | 基本不改变数值正确性,但增加网络通信开销 | Memory Viewer 中单个张量远大于其他(如被复制到所有 TPU);TensorBoard hooks 中数组显示未分片 |
| 减小 batch size | 改变训练动态,通常要重调学习率(micro-batching 是不改变行为的替代) | 梯度计算时 "Temporaries" 分配失败;Op 名含 "JVP";内存 profile 中大量 batch 形状张量 |
| 启用混合精度(bfloat16) | 有风险,改变数值精度,可能影响结果或导致不收敛 | Memory Viewer 确认最大张量目前是float32 |
手动 checkpointing(jax.checkpoint) | 不改变行为,用计算换内存 | 反向传播时大量完全同尺寸的张量占满内存,常伴随 "JVP" Op 名 |
捐赠输入 buffer(donate_argnums) | 不改变实验完整性;用错会直接报清晰的错误 | 无特定信号,属于"白捡的赢面",值得先试 |
| 修改模型维度 | 改变模型行为,可能直接破坏与数据集的兼容性 | Memory Viewer 显示大量 padding 浪费(维度不是 128 的倍数等) |
| Host offloading | 数值安全但性能上是大坑 | 通常只作为超大优化器状态或重预处理步骤的最后手段 |
修复后如何确认
这个错误发生在编译期检查阶段,所以验证方式就是重新运行程序:修改生效后,XLA:TPU 编译应当通过,不再出现RESOURCE_EXHAUSTED: ... Ran out of memory in memory space hbm。另外注意,上文中会改变程序行为的几个手段(减小 batch size、bfloat16/量化、修改模型维度、xla_tpu_pre_fusion_remat)在文档中都有"可能影响模型行为/收敛/数值稳定性"的标注——对这类改动,编译通过只是第一步,还需要按文档提示重新核对训练指标(如重调学习率、观察收敛),而纯 sharding、donate_argnums、jax.checkpoint这类不改变数值正确性的手段,编译通过即可视为问题解除。
参考文档:E1000 - Compile Time: HBM OOM、Debug OOM errors with XProf、XLA flags guidance、XLA 错误总览。
【免费下载链接】tensorflowAn Open Source Machine Learning Framework for Everyone项目地址: https://gitcode.com/GitHub_Trending/te/tensorflow
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考