导入本地MKQA数据集后使用map函数触发TypeError的解决方法
问题:MKQA数据集map处理时触发TypeError错误
场景与问题描述
从本地路径/data/mkqa-Chinese加载MKQA数据集后,尝试用Tokenizer结合map函数处理数据集时触发TypeError,报错提示list indices must be integers or slices, not str。
相关代码与信息
- 数据集加载代码:
from datasets import Dataset, load_dataset raw_dataset = load_dataset(data_path)
- 加载后的数据集结构:
DatasetDict({ train: Dataset({ features: ['query', 'answers', 'queries', 'example_id'], num_rows: 6758 }) })
- Tokenizer处理代码:
model_path = "/data/bigscience/bloomz-3b" tokenizer = AutoTokenizer.from_pretrained(model_path, trust_remote_code=True, use_Fast=True) tokenizer.pad_token = tokenizer.eos_token tokenizer.padding_side = 'right' def tok(sample): prompt_and_chosen = " Human: " + sample['queries']['zh_cn'] + " Assistant: " + sample['answers']['zh_cn'][0]['text'] model_inps = tokenizer(prompt_and_chosen, padding=True, max_length=512, truncation=True) return model_inps tokenized_training_data = raw_dataset['train'].map(tok, batched=True) print(tokenized_training_data) print("pause")
- 报错信息:
processed_inputs = function(*fn_args, *additional_args, **fn_kwargs) File "/home/novo_trl_sft.py", line 548, in tok prompt_and_chosen = " Human: " + sample['queries']['zh_cn'] + " Assistant: " + sample['answers']['zh_cn'][0]['text'] TypeError: list indices must be integers or slices, not str
解决方案
问题根源
你开启了batched=True参数,此时传入tok函数的是批量样本组成的字典——字典中每个字段对应的值都是包含所有样本数据的列表,而不是单条样本的字典。比如sample['queries']['zh_cn']是一个列表,直接用字符串索引自然会报错。
两种修正方式
方式一:保持批量处理(推荐,速度更快)
修改tok函数适配批量数据结构:
def tok(batch): # 遍历批量中的每一组query和answer prompts = [] for query, answer_group in zip(batch['queries']['zh_cn'], batch['answers']['zh_cn']): prompt = f" Human: {query} Assistant: {answer_group[0]['text']}" prompts.append(prompt) # 批量传入tokenizer处理 model_inps = tokenizer(prompts, padding=True, max_length=512, truncation=True) return model_inps
方式二:关闭批量处理
如果不需要批量加速,可以直接将map的batched参数设为False,此时tok函数会接收单条样本,原有的字段访问逻辑可以正常工作:
tokenized_training_data = raw_dataset['train'].map(tok, batched=False)
调试小技巧
如果不确定传入函数的是单条还是批量数据,可以在tok函数开头加一行打印:
print(type(sample), sample['queries']['zh_cn'])
通过输出的类型和内容,就能快速定位数据结构问题。
内容的提问来源于stack exchange,提问作者4daJKong
相关产品推荐
相关产品推荐

