批量推理长度分桶怎么配置或排查
来源:17golang原创
时间:2026-09-13 09:16:44 221浏览 收藏
我第一次给一批长短差异明显的文本做推理时,batch size 明明没有变,显存却随着某几个长样本突然冲高。原因通常不是“模型随机变大”,而是一个 batch 会按最长输入补齐:短文本也占用了最长文本的 token 空间。长度分桶的正确做法,是先在预处理阶段取得每条样本的 input_ids 长度,把相近长度放进同一批,再用动态 padding 只补到该批最长值。
group_by_length是 Transformers Trainer 的采样策略,离线批量推理要在 DataLoader 上配置batch_sampler。- 分桶边界应由真实 token 长度分布决定;
padding="longest"与pad_to_multiple_of=8解决的是补齐方式,不是分桶。 - 排查先看每批最长长度、有效 token 比例和峰值显存,再决定缩小 batch、收紧上限还是调整边界。
先分清训练参数和推理分桶
Transformers 的 LengthGroupedSampler 会把长度相近的样本放在相邻位置,主要用于 Trainer 的数据加载;新版本文档也把 group_by_length 定义为减少 padding 的训练采样策略。它不会自动接管一个自定义推理循环。推理服务或离线脚本若直接写 for batch in DataLoader(dataset, batch_size=...),默认仍可能把完全不同长度的样本混在一起。
因此配置要拆成两层:第一层是“哪些样本属于一批”,由长度桶和 batch_sampler 决定;第二层是“批内补多少”,由 tokenizer 或 DataCollatorWithPadding 决定。两层混在一起,是最常见的排查误区。

按 token 长度建立可控的批次
边界不要照搬模型的最大上下文长度。先统计业务样本的 token 长度,再用少量有意义的断点,例如 128、256、512、1024;超出最后边界的样本进入单独的长样本桶。下面的实现只负责产生批次索引,模型仍由你的推理代码调用。
import torch
from torch.utils.data import DataLoader
from transformers import DataCollatorWithPadding
def make_batches(lengths, batch_size, boundaries=(128, 256, 512, 1024)):
# 先按 token 长度入桶,避免短文本被最长文本大量补齐
buckets = [[] for _ in range(len(boundaries) + 1)]
for index, length in enumerate(lengths):
bucket_id = next((i for i, limit in enumerate(boundaries) if length
这里的边界是“软约束”:最后一批可能不足 batch_size,长样本也可能把桶撑高。若模型是生成式模型,还要把输入长度和 max_new_tokens 分开预算;后者变大时,即使输入分桶合理,KV cache 仍会增加。

用可观测指标排查配置
不要只盯着平均耗时。给每个 batch 记录 batch_size、max_input_tokens、sum_input_tokens 和 padding 比例:
padding_ratio = 1 - sum_input_tokens / (batch_size * max_input_tokens)。如果比例长期很高,先调整桶边界或减小单批样本数;如果比例不高但显存仍冲高,重点检查 max_new_tokens、模型 dtype、KV cache 和是否混入超长样本。
| 现象 | 优先看什么 | 处理方向 |
|---|---|---|
| 短文本批次也接近长文本显存 | 每批最大输入长度 | 缩窄桶边界,长样本单独处理 |
| 吞吐忽高忽低 | 有效 token 比例与批次 token 数 | 按 token 预算,而非只固定样本数 |
| 结果顺序错乱 | 自定义批次返回的原始 index | 输出携带 index,写回原数组 |
| 对齐后反而变慢 | padding 到的实际长度 | 比较 8/16 对齐与不对齐的实测结果 |
常见问题
分桶是不是越细越快?
不一定。桶太细会让尾批变多、批次变小,调度和设备利用率可能变差。先保证批次 token 数稳定,再用指标调整。
可以直接把 max_length 当分桶边界吗?
可以作为最后一道上限,但它不等于业务分布。超过上限的输入要明确截断或拒绝,不能静默丢掉关键信息。
为什么开启 padding=True 仍然浪费很多 token?
padding=True 通常只表示补到当前批次最长序列;若同批长度差异本身很大,仍需先改批次组成,而不是继续改 padding 开关。
我的经验是先把分桶当成数据装载问题处理,再谈模型参数优化:先确认每批 token 形状稳定,再逐步调 batch size、对齐倍数和生成上限。这样出现显存峰值时,通常能很快判断是输入分布、padding,还是生成阶段的缓存。
-
501 收藏
-
501 收藏
-
501 收藏
-
501 收藏
-
501 收藏
-
- 前端进阶之JavaScript设计模式
- 设计模式是开发人员在软件开发过程中面临一般问题时的解决方案,代表了最佳的实践。本课程的主打内容包括JS常见设计模式以及具体应用场景,打造一站式知识长龙服务,适合有JS基础的同学学习。
- 立即学习 543次学习
-
- GO语言核心编程课程
- 本课程采用真实案例,全面具体可落地,从理论到实践,一步一步将GO核心编程技术、编程思想、底层实现融会贯通,使学习者贴近时代脉搏,做IT互联网时代的弄潮儿。
- 立即学习 516次学习
-
- 简单聊聊mysql8与网络通信
- 如有问题加微信:Le-studyg;在课程中,我们将首先介绍MySQL8的新特性,包括性能优化、安全增强、新数据类型等,帮助学生快速熟悉MySQL8的最新功能。接着,我们将深入解析MySQL的网络通信机制,包括协议、连接管理、数据传输等,让
- 立即学习 500次学习
-
- JavaScript正则表达式基础与实战
- 在任何一门编程语言中,正则表达式,都是一项重要的知识,它提供了高效的字符串匹配与捕获机制,可以极大的简化程序设计。
- 立即学习 487次学习
-
- 从零制作响应式网站—Grid布局
- 本系列教程将展示从零制作一个假想的网络科技公司官网,分为导航,轮播,关于我们,成功案例,服务流程,团队介绍,数据部分,公司动态,底部信息等内容区块。网站整体采用CSSGrid布局,支持响应式,有流畅过渡和展现动画。
- 立即学习 485次学习