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)
关键细节说明
- 避免批量级别的列表操作:必须遍历每个样本单独拼接,不能直接对整个batch的
input_ids列表进行切片拼接,否则会导致样本对应关系混乱,长度更不可控。 - 处理空序列情况:加入了
ids2_trimmed = ids2[1:] if len(ids2) > 0 else []的判断,防止text2为空时出现索引错误。 - 固定长度保障:通过截断+padding强制每个
input_ids长度为MAX_TOTAL_LENGTH,完全符合Arrow的存储要求。
内容的提问来源于stack exchange,提问作者SMMousaviSP
相关产品推荐
相关产品推荐

