Pegasus Tokenizer批量处理报错:TypeError类型不匹配问题咨询
解决Pegasus模型批量处理时Tokenizer的TypeError问题
问题背景
我在学习Transformers技术栈,尝试用Pegasus模型对数据集文本生成摘要,以此适配BERT Tokenizer的长度限制。调用map函数时,设置batched=False代码运行正常,但切换为batched=True时,出现报错:
TypeError: TextEncodeInput must be Union[TextInputSequence, Tuple[InputSequence, InputSequence]]
已完成的排查步骤
- 检查并移除了数据集中的空值样本
- 确认输入数据为字符串列表格式
- 尝试从批量样本中直接生成字符串列表作为Tokenizer输入
- 编写输入有效性检查函数,验证单条及批量输入的格式合法性
完整代码示例
from datasets import load_dataset from transformers import PegasusTokenizer, PegasusForConditionalGeneration # 加载数据集(替换为实际本地路径) dataset = load_dataset("csv", data_files="your_dataset.csv") dataset = dataset["train"].filter(lambda x: len(x["text"].strip()) > 0) # 初始化Tokenizer和模型 tokenizer = PegasusTokenizer.from_pretrained("google/pegasus-xsum") model = PegasusForConditionalGeneration.from_pretrained("google/pegasus-xsum") # 原处理函数(单条正常,批量报错) def process_sample(sample): # 错误点:批量时会将每个文本包装成列表,导致输入为[[text1], [text2]] inputs = tokenizer([sample["text"]], truncation=True, max_length=512) return inputs # 报错的调用方式 # tokenized_dataset = dataset.map(process_sample, batched=True) # 输入验证函数 def validate_inputs(batch): for text in batch["text"]: assert isinstance(text, str), f"非字符串输入:{text}" assert len(text.strip()) > 0, "空文本输入" return True # 验证数据集 validate_inputs(dataset[:10]) # 单条验证通过 validate_inputs({"text": dataset[:10]["text"]}) # 批量验证通过
解决方案
1. 修正处理函数适配批量输入
当batched=True时,map传入的是批量字典,每个字段对应的值是该字段的样本列表。处理函数需要直接将列表传入Tokenizer,而非对每个样本单独包装列表:
def process_batch(batch): # 直接传入批量文本列表,无需额外嵌套 inputs = tokenizer(batch["text"], truncation=True, max_length=512, padding="longest") # 生成摘要逻辑(按需添加) outputs = model.generate(**inputs) batch["summary"] = [tokenizer.decode(out, skip_special_tokens=True) for out in outputs] return batch # 正确的批量调用 tokenized_dataset = dataset.map(process_batch, batched=True, batch_size=8)
2. 排查输入格式嵌套问题
报错的核心原因是Tokenizer收到了嵌套列表(如[[text1], [text2]]),而非一维字符串列表。检查处理函数中是否存在将单条文本包装为列表的逻辑,批量时会导致嵌套,需移除多余的列表包装。
3. 批量输入验证强化
在处理函数开头加入批量验证,提前拦截无效输入:
def process_batch(batch): # 批量验证所有文本都是非空字符串 assert all(isinstance(t, str) and len(t.strip()) > 0 for t in batch["text"]), "存在无效文本输入" inputs = tokenizer(batch["text"], truncation=True, max_length=512) return inputs
内容的提问来源于stack exchange,提问作者Arcaici
相关产品推荐
相关产品推荐

