You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.13 20:12:42