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

Hugging Face Datasets map(batch=True)报ArrowInvalid列长度不匹配错误

解决ArrowInvalid错误:批量分词拼接后的长度不一致问题

问题根源

当用batched=True调用datasets.map()时,Arrow要求每一列的所有样本长度必须统一(列式存储特性)。你的自定义函数拼接后的input_ids在同一个batch里长度不固定,导致Arrow无法写入,触发长度不匹配的错误。而batch=False是单条处理,每条样本的长度不影响存储,所以没问题。

解决方案:强制统一拼接后的序列长度

核心思路是让每个样本处理后的input_ids长度固定,通过截断+padding实现,确保同一batch内所有样本长度一致。下面是完整的可运行代码示例:

from transformers import AutoTokenizer
from datasets import load_dataset

# 初始化tokenizer和参数
tokenizer = AutoTokenizer.from_pretrained("bert-base-uncased")
MAX_TOTAL_LENGTH = 512  # 最终拼接后的总最大长度
MAX_PER_SEQ = MAX_TOTAL_LENGTH // 2  # 每个单独序列的截断长度

def tokenize_function(examples):
    # 分别对text1、text2批量分词,只保留input_ids,不自动padding
    tokenized_text1 = tokenizer(
        examples["text1"],
        truncation=True,
        max_length=MAX_PER_SEQ,
        padding=False,
        return_attention_mask=False,
        return_token_type_ids=False
    )
    tokenized_text2 = tokenizer(
        examples["text2"],
        truncation=True,
        max_length=MAX_PER_SEQ,
        padding=False,
        return_attention_mask=False,
        return_token_type_ids=False
    )
    
    # 逐个样本拼接并处理长度
    input_ids = []
    for ids1, ids2 in zip(tokenized_text1["input_ids"], tokenized_text2["input_ids"]):
        # 移除第二个序列的[CLS]标记
        ids2_trimmed = ids2[1:] if len(ids2) > 0 else []
        # 拼接两个序列
        combined_ids = ids1 + ids2_trimmed
        # 截断到总最大长度(防止极端情况拼接后超出)
        if len(combined_ids) > MAX_TOTAL_LENGTH:
            combined_ids = combined_ids[:MAX_TOTAL_LENGTH]
        # padding到固定长度
        combined_ids += [tokenizer.pad_token_id] * (MAX_TOTAL_LENGTH - len(combined_ids))
        input_ids.append(combined_ids)
    
    # 生成对应的attention_mask
    attention_mask = []
    for ids in input_ids:
        mask = [1 if token != tokenizer.pad_token_id else 0 for token in ids]
        attention_mask.append(mask)
    
    return {
        "input_ids": input_ids,
        "attention_mask": attention_mask
    }

# 加载数据集并应用分词函数
dataset = load_dataset("your_dataset_name")
processed_dataset = dataset.map(tokenize_function, batched=True, batch_size=32)

关键细节说明

  1. 避免批量级别的列表操作:必须遍历每个样本单独拼接,不能直接对整个batch的input_ids列表进行切片拼接,否则会导致样本对应关系混乱,长度更不可控。
  2. 处理空序列情况:加入了ids2_trimmed = ids2[1:] if len(ids2) > 0 else []的判断,防止text2为空时出现索引错误。
  3. 固定长度保障:通过截断+padding强制每个input_ids长度为MAX_TOTAL_LENGTH,完全符合Arrow的存储要求。

内容的提问来源于stack exchange,提问作者SMMousaviSP

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.17 14:23:08