基于滑动窗口微调PubMedBERT实现NER任务报错求助
长文本NER场景下滑动窗口的适用性及错误解决
滑动窗口的适用性
滑动窗口完全适用于你这种长文本命名实体识别场景。PubMedBERT的最大输入长度为512,通过滑动窗口可以在不破坏完整语义的前提下,将长文本拆分为多个重叠子窗口,每个子窗口长度控制在512以内,最后合并各窗口的预测结果得到完整文本的NER标注,完美解决长文本处理问题。
错误分析与解决
1. ValueError: expected sequence of length 4079 at dim 1 (got 5846)
这个错误的核心原因是输入序列与标签序列长度不匹配,或者批量处理时窗口长度未统一。常见触发场景:
- 滑动窗口切分后,部分窗口未对齐到模型要求的512长度,导致批量加载时张量维度冲突;
- 标签序列填充逻辑错误,未对短窗口的标签用
-100(Transformer默认忽略损失计算的标签值)填充,导致输入与标签长度不一致。
解决方法:
- 强制所有窗口的
input_ids和labels长度为512,不足部分用padding补全:输入用模型的pad token id(通常为0)填充,标签用-100填充; - 使用Hugging Face提供的
DataCollatorForTokenClassification自动处理批量数据的padding和标签对齐,避免手动处理的误差。
2. IndexError: Invalid Key: 49 is Out of bounds for size 0
这个错误是因为访问了空列表/张量的索引,常见触发场景:
- 短文本(长度小于512)未单独处理,滑动窗口循环生成的窗口列表为空,后续代码尝试访问索引导致报错;
- 滑动窗口的循环条件错误,比如
range的起始/结束值计算失误,导致生成的窗口数量为0。
解决方法:
- 在切分窗口前增加判断:如果文本长度小于等于512,直接生成一个补全后的窗口,避免空列表;
- 检查滑动窗口的循环逻辑,确保循环能覆盖所有文本片段,比如最后一段不足窗口长度时,从文本末尾向前取512长度并补全。
修正后的滑动窗口切分示例代码
from transformers import DataCollatorForTokenClassification def sliding_window_process(tokenized_text, window_size=512, stride=256): input_ids = tokenized_text["input_ids"] labels = tokenized_text["labels"] total_len = len(input_ids) windows = [] # 处理短文本 if total_len <= window_size: pad_len = window_size - total_len padded_input = input_ids + [0] * pad_len padded_labels = labels + [-100] * pad_len windows.append({ "input_ids": padded_input, "attention_mask": [1]*total_len + [0]*pad_len, "labels": padded_labels }) return windows # 滑动切分长文本 for start in range(0, total_len - window_size + 1, stride): end = start + window_size windows.append({ "input_ids": input_ids[start:end], "attention_mask": [1]*window_size, "labels": labels[start:end] }) # 处理末尾剩余片段 if end < total_len: start = total_len - window_size remaining_input = input_ids[start:] remaining_labels = labels[start:] pad_len = window_size - len(remaining_input) windows.append({ "input_ids": remaining_input + [0]*pad_len, "attention_mask": [1]*len(remaining_input) + [0]*pad_len, "labels": remaining_labels + [-100]*pad_len }) return windows # 数据加载时使用专用的collator data_collator = DataCollatorForTokenClassification(tokenizer)
内容的提问来源于stack exchange,提问作者Danial
相关产品推荐
相关产品推荐

