使用带stride的预训练模型处理长文本时预测中断的问题求助
带stride的预训练模型处理长文本时预测中断的问题求助
嘿,我之前也踩过Hugging Face pipeline处理长文本NER时的类似坑,你的问题大概率是pipeline没有自动利用tokenizer的stride和overflow参数做滑动窗口预测导致的——它可能只处理了文本的第一个512-token窗口,后面的内容直接被忽略了,所以才会出现标注到一半就停的情况。我给你梳理下解决思路和具体操作:
问题根源
你在tokenizer里设置了stride=128和return_overflowing_tokens=True,这会把长文本切割成多个重叠的token窗口(每个窗口512token,重叠128token),但pipeline("token-classification")的默认逻辑并不会自动遍历这些窗口做完整预测,它只会处理第一个窗口的结果,自然就断在中间了。
解决办法:手动实现滑动窗口预测与结果合并
放弃依赖pipeline的默认处理,手动完成文本切割、模型预测、结果映射这三步,才能让stride真正发挥作用。给你一个可参考的代码示例:
import torch from transformers import AutoTokenizer, AutoModelForTokenClassification # 加载训练好的模型和tokenizer model = AutoModelForTokenClassification.from_pretrained(model_path) tokenizer = AutoTokenizer.from_pretrained(model_path, model_max_length=512, is_split_into_words=True) # 假设你的长文本是按词分割好的列表(如果是原始字符串就去掉is_split_into_words=True) long_text = ["你的", "长文本", "内容", "..."] # 1. 用tokenizer切割长文本为带重叠的窗口 encoding = tokenizer( long_text, stride=128, return_overflowing_tokens=True, truncation=True, is_split_into_words=True, return_tensors="pt" ) # 保存窗口到原文本的映射关系,后续用来合并结果 overflow_mapping = encoding.pop("overflow_to_sample_mapping") # 获取所有窗口的word_ids,用来把token级预测映射回原文本的词 all_word_ids = [encoding.word_ids(batch_index=i) for i in range(len(encoding["input_ids"]))] # 2. 模型批量预测 model.eval() with torch.no_grad(): outputs = model(**encoding) predictions = outputs.logits.argmax(dim=-1).tolist() # 3. 合并多个窗口的预测结果,处理重叠部分 final_predictions = [] for window_pred, word_ids in zip(predictions, all_word_ids): prev_word_idx = None for idx, pred in enumerate(window_pred): word_idx = word_ids[idx] # 跳过特殊token(如[CLS]/[SEP])和同一个词的子token if word_idx is None or word_idx == prev_word_idx: continue prev_word_idx = word_idx # 处理重叠:如果当前词已经有预测,用后续窗口的结果覆盖(也可以改成取多数投票) if word_idx >= len(final_predictions): final_predictions.append(pred) else: final_predictions[word_idx] = pred # 最终final_predictions就是原文本每个词对应的标注结果
额外注意事项
- 训练阶段你设置的tokenizer参数是正确的,要确保训练时每个窗口样本都被用来训练了(比如数据加载时没有过滤overflow样本),这样模型才具备处理长文本的能力。
- 如果坚持想用pipeline,可以试试在初始化pipeline时传入
tokenizer_kwargs={"stride":128, "return_overflowing_tokens":True},但实测下来还是手动处理更可靠,因为pipeline对overflow的支持并不完善。
备注:内容来源于stack exchange,提问作者Despe1990
相关产品推荐
相关产品推荐

