使用Huggingface DistilBERT做问答时input_ids超512如何正确截断
DistilBERT问答场景input_ids长度超限解决方案
错误根源
你之前的截断逻辑存在问题:先单独对context做截断得到长度为512的token id,再将question和该id传入encode_plus时,tokenizer会额外添加question对应的token、<[BOS_never_used_51bce0c785ca2f68081bfa7d91973934]>起始符、两个[SEP]分隔符,最终总token长度必然超过512,所以触发报错。
方案1:快速截断修复(适合短上下文场景)
直接将question和原始context一起传给encode_plus,指定截断规则即可,无需提前处理context:
把askQuestion函数里的encoding行替换为以下代码:
encoding = tokenizer.encode_plus( question, context, truncation="only_second", # 仅截断第二个输入(即context),保留完整question max_length=512, padding="max_length", return_attention_mask=True, return_tensors='pt' # 直接返回pytorch张量,省去后续转tensor的步骤 )
同时后面调用模型的代码可以简化为:
start_scores, end_scores = model(**encoding)
这个方案改动最小,能快速解决长度超限问题,但如果答案在context的后半部分,直接截断会导致答案丢失。
方案2:滑动窗口分块处理(长上下文最优方案)
针对你使用超长大文本作为context的场景,滑动窗口是更合理的解决方案:
- 先计算question的token长度,预留出特殊符的位置,得到每个context块的最大长度:
chunk_size = 512 - len(tokenizer(question).input_ids) - 3(3对应1个<[BOS_never_used_51bce0c785ca2f68081bfa7d91973934]>+2个[SEP]的长度) - 把整个context切成多个重叠的chunk,重叠长度设为32-128之间,避免答案刚好落在块的边界被截断
- 每个chunk分别和question拼接后输入模型,计算所有chunk输出的start_scores和end_scores的最大值,最终取置信度最高的答案
这个方案不会遗漏context里的内容,问答准确率远高于直接截断,同时也适配树莓派的低内存环境,每次只加载一小块文本跑模型,内存占用可控。
树莓派移植注意事项
- 可以直接使用
transformers的pipeline封装好的问答接口,指定truncation=True和max_length=512,自动处理截断逻辑,代码更简洁 - 模型加载时开启量化:
model = DistilBertForQuestionAnswering.from_pretrained(..., torch_dtype=torch.float16),减少一半内存占用,适配树莓派的有限内存。
内容的提问来源于stack exchange,提问作者Scott Bing
相关产品推荐
相关产品推荐

