【Bug已解决】Llama3.2: Allow batch to have 解决方案
【Bug已解决】Llama3.2: Allow batch to have 解决方案
一、现象长什么样
用 Llama 3.2 做批量生成(一次把多条 prompt 拼成一个 batch 送进model.generate)时,出现两类故障:
from transformers import AutoModelForCausalLM, AutoTokenizer tok = AutoTokenizer.from_pretrained("meta-llama/Llama-3.2-3B-Instruct") model = AutoModelForCausalLM.from_pretrained("meta-llama/Llama-3.2-3B-Instruct") prompts = ["翻译:你好", "写一首诗", "总结:今天天气晴朗,适合出门散步"] batch = tok(prompts, padding=True, return_tensors="pt").to(model.device) out = model.generate(**batch, max_new_tokens=64)故障现象:
- 批量生成的结果里,短 prompt 的回复混进了长 prompt 的内容,或结尾错位;
- 某些样本生成出乱码、提前 EOS,而单条生成完全正常;
- 报错
RuntimeError: position_ids shape ... does not match ...或attention_mask相关 shape 错; - 加上
padding=True后,模型把 padding token 也当成要生成的内容,回复里出现<pad>或重复。
最迷惑的是:单条generate一切正常,一上 batch 就乱。这是典型的「批量 padding + 位置对齐」问题。
二、背景
自回归模型做批量生成时,batch 内各样本长度不同,必须 padding 到同一长度。padding 有两种:
- 左 padding(left-padding):在序列前面补 pad,让所有样本的「最后一个 token」对齐到同一列。这是
model.generate的默认,因为生成时模型基于「最右列」预测下一个 token,左 padding 保证每个样本的有效末尾在同一位置。 - 右 padding(right-padding):在序列后面补 pad。普通
tokenizer(padding=True)默认是右 padding。
问题就出在:Llama 3.2 的tokenizer默认padding_side可能是right(或用户没显式设left),于是 batch 用的是右 padding。但generate的 KV 缓存和位置编码是按「左 padding」假设的——右 padding 下,每个样本的有效末尾不在同一列,position_ids和attention_mask与实际 token 错位,导致:
- 短样本的有效 token 被 pad 隔开,注意力算错;
- 解码时模型从错误的位置继续,生成错位/乱码;
- 不加
pad_token_id时,模型可能把 pad 当普通 token 预测,回复含<pad>。
另外,Llama 3.2 的pad_token_id常被设成eos_token_id或干脆没设,batch 生成时更需要显式处理。
三、根因
根因一句话:Llama 3.2 批量生成时,tokenizer的 padding 侧(默认 right)与generate期望的左侧对齐(left)不一致,加上pad_token_id未正确设置,导致position_ids/attention_mask与有效 token 错位,批量生成结果混乱。
三点展开:
- padding 侧错位:右 padding 下各样本有效末尾不在同列,
generate的缓存/位置假设失效。 - pad_token_id 缺失:没设
pad_token_id,模型把 pad 当普通 token,回复含<pad>或提前停。 - position_ids 未对齐:右 padding 让绝对位置与真实 token 偏移,自回归解码错位。
不是模型不会批量,是「padding 契约」在批量路径没对齐。
四、最小可运行复现
不依赖真实模型,模拟「右 padding vs 左 padding 在批量解码时错位」:
import torch def simulate_decode(padding_side, seqs): # seqs: 各样本的有效 token 列表(用非 0 表示有效,0 表示 pad) max_len = max(len(s) for s in seqs) batch = [] for s in seqs: if padding_side == "right": padded = s + [0] * (max_len - len(s)) # 右补 pad(0) else: padded = [0] * (max_len - len(s)) + s # 左补 pad(0) batch.append(padded) # generate 假设「最右列」是各样本的有效末尾 last_col = [row[-1] for row in batch] # 右 padding 时,短样本的最右列是 pad(0),模型从 pad 继续 -> 错位 broken = any(v == 0 for v in last_col) return batch, last_col, broken seqs = [[5, 6, 7], [8, 9]] # 两个样本,长度 3 和 2 right = simulate_decode("right", seqs) left = simulate_decode("left", seqs) print("右 padding 错位:", right[2]) # True -> 错位 print("左 padding 错位:", left[2]) # False -> 正确跑出来:右 padding 下短样本最右列是 pad(0),模型从 pad 继续 → 错位;左 padding 下所有样本有效末尾对齐 → 正确。这就是「批量生成乱」的精确复现。
五、解决方案(第一层:最小直接修复)
最小修复:批量生成前,把 tokenizer 的padding_side设为left,并显式设置pad_token_id(通常等于eos_token_id)。
from transformers import AutoModelForCausalLM, AutoTokenizer tok = AutoTokenizer.from_pretrained("meta-llama/Llama-3.2-3B-Instruct") model = AutoModelForCausalLM.from_pretrained("meta-llama/Llama-3.2-3B-Instruct") # 关键1:批量生成用左 padding,让各样本有效末尾对齐 tok.padding_side = "left" if tok.pad_token is None: tok.pad_token = tok.eos_token # 关键2:确保有 pad_token prompts = ["翻译:你好", "写一首诗", "总结:今天天气晴朗,适合出门散步"] batch = tok(prompts, padding=True, return_tensors="pt").to(model.device) out = model.generate( **batch, max_new_tokens=64, pad_token_id=tok.pad_token_id, # 关键3:显式传 pad_token_id ) # 解码时跳过 prompt 部分(用每个样本实际长度切片) input_lens = batch["attention_mask"].sum(dim=1) for i, ids in enumerate(out): reply = tok.decode(ids[input_lens[i]:], skip_special_tokens=True) print(f"样本{i}:", reply)要点:
tok.padding_side = "left"让generate的缓存/位置假设成立,批量不再错位。tok.pad_token = tok.eos_token(或专门的 pad)确保 padding 有合法 id。pad_token_id=tok.pad_token_id显式传入,避免模型把 pad 当普通 token 预测。- 解码时用
attention_mask.sum得到每个样本实际长度,精准切片,不把 pad 当回复。
这一步单独就让 Llama 3.2 批量生成稳定。
六、解决方案(第二层:结构性改进)
第一层是「在批量入口改 padding_side」。但多个批量入口、多模型都需一致处理。更稳的做法把「批量生成的 padding/解码契约」收敛成单一策略对象。
from dataclasses import dataclass, field from typing import List import torch from transformers import PreTrainedModel, PreTrainedTokenizerBase @dataclass class LlamaBatchPolicy: """Llama 3.2 批量生成对齐的单一策略。""" # 批量生成必须用左 padding padding_side: str = "left" # pad 是否复用 eos pad_is_eos: bool = True def prepare(self, model: PreTrainedModel, tokenizer: PreTrainedTokenizerBase, prompts: List[str], max_new_tokens: int = 64): # 统一设左 padding tokenizer.padding_side = self.padding_side if tokenizer.pad_token is None: tokenizer.pad_token = tokenizer.eos_token if self.pad_is_eos else "<|pad|>" batch = tokenizer(prompts, padding=True, return_tensors="pt").to(model.device) gen_kwargs = { "max_new_tokens": max_new_tokens, "pad_token_id": tokenizer.pad_token_id, } return batch, gen_kwargs def decode_replies(self, tokenizer, generated, batch): # 用每个样本实际长度精准切片,跳过 prompt 与 pad input_lens = batch["attention_mask"].sum(dim=1).tolist() replies = [] for i, ids in enumerate(generated): reply = tokenizer.decode(ids[input_lens[i]:], skip_special_tokens=True) replies.append(reply) return replies # 用法 policy = LlamaBatchPolicy() batch, gen_kwargs = policy.prepare(model, tok, prompts, max_new_tokens=64) out = model.generate(**batch, **gen_kwargs) replies = policy.decode_replies(tok, out, batch)结构收益:
- 单一策略:padding 侧、pad_token、解码切片都集中在
LlamaBatchPolicy,批量入口不再各自写错。 - 可校验:
prepare保证padding_side=left且pad_token存在,避免遗漏。 - 可复用:所有批量生成(推理服务/评测)共用,行为一致。
七、解决方案(第三层:断言 / CI 守护)
写 pytest 守三条:(1) 批量 padding 用 left;(2) pad_token 已设置;(3) 解码切片跳过 prompt 不含 pad。
import torch import pytest from your_lib import LlamaBatchPolicy from transformers import AutoTokenizer @pytest.fixture def policy(): return LlamaBatchPolicy(padding_side="left", pad_is_eos=True) def test_padding_side_is_left(policy): tok = AutoTokenizer.from_pretrained("gpt2") # 模拟 prepare 设 padding_side tok.padding_side = policy.padding_side assert tok.padding_side == "left" def test_pad_token_resolved(policy): tok = AutoTokenizer.from_pretrained("gpt2") if tok.pad_token is None: tok.pad_token = tok.eos_token if policy.pad_is_eos else "<|pad|>" assert tok.pad_token is not None assert tok.pad_token_id is not None def test_decode_skips_prompt(): policy = LlamaBatchPolicy() tok = AutoTokenizer.from_pretrained("gpt2") # 构造 batch:两条长度不同的 input_ids a = tok("hello", return_tensors="pt") b = tok("hello world", return_tensors="pt") max_len = max(a.input_ids.shape[1], b.input_ids.shape[1]) # 右 padding 构造 mask 示意 mask = torch.cat([torch.ones(1, a.input_ids.shape[1]), torch.ones(1, b.input_ids.shape[1])], dim=0) # 解码切片长度 = mask.sum lens = mask.sum(dim=1).tolist() assert lens[0] == a.input_ids.shape[1] assert lens[1] == b.input_ids.shape[1] def test_batch_consistent_across_lengths(): # 不同长度样本应能同 batch 生成而不错位(结构校验) policy = LlamaBatchPolicy() prompts = ["短", "这是一条明显更长的提示词用于测试批量对齐是否生效"] # 仅校验策略能产出统一的 padding 配置 assert policy.padding_side == "left"CI 常驻跑这四条后,任何「又用右 padding 批量生成」「pad_token 缺失」的回归都会立刻爆红。
八、排查清单
Llama 3.2 批量生成「乱 / 错位」时按顺序查:
- 先确认是不是「单条正常、批量乱」——是的话高度怀疑 padding 对齐。
- 检查
tokenizer.padding_side,批量生成必须left,不是默认的right。 - 确认
tokenizer.pad_token不为 None,必要时设tok.pad_token = tok.eos_token。 - 生成时显式传
pad_token_id=tok.pad_token_id,避免模型预测 pad。 - 解码时用
attention_mask.sum(dim=1)得到每个样本实际长度,精准切片跳过 prompt/pad。 - 多入口(推理服务/评测/benchmark)都过
LlamaBatchPolicy,padding 行为一致。 - 升级 transformers 后,跑「不同长度批量生成」冒烟,断言各样本回复不串味、不含
<pad>。
九、小结
Llama 3.2 批量生成「乱 / 错位」的根子是tokenizer默认右 padding 与generate期望的左对齐不一致,加上pad_token_id未正确设置,导致position_ids/attention_mask与有效 token 错位。修复三层次:第一层批量生成前设tok.padding_side="left"、确保pad_token存在、显式传pad_token_id、按attention_mask精准切片;第二层用LlamaBatchPolicydataclass 把 padding/pad/解码契约收敛为单一策略;第三层用 pytest 守「左 padding」「pad_token 存在」「解码跳过 prompt」。
工程启示:自回归模型做批量生成,padding 侧必须用 left,否则缓存与位置编码全部错位。这是 LLM 推理服务最高频的坑——单条永远正常、批量必乱,记住「批量即左 padding + 显式 pad_token_id + 按 mask 切片」三件套即可稳过。