运行Transformer文本分类代码时遇TypeError:batch_text_or_text_pairs需为列表
问题:多类别文本分类代码触发TypeError错误
运行多类别文本分类代码时触发TypeError,错误提示:batch_text_or_text_pairs has to be a list (got )
报错回溯
TypeError Traceback (most recent call last) File :1 Cell In[20], line 5, in regular_encode(texts, tokenizer, maxlen) 1 def regular_encode(texts, tokenizer, maxlen=512): 2 """ 3 encodes text for a model 4 """ ----> 5 enc_di = tokenizer.batch_encode_plus( 6 texts, 7 return_token_type_ids=False, 8 padding='max_length', 9 max_length=None 10 11 ) 13 return np.array(enc_di['input_ids']) File ~/.local/lib/python3.8/site-packages/transformers/tokenization_utils_base.py:2473, in PreTrainedTokenizerBase.batch_encode_plus(self, batch_text_or_text_pairs, add_special_tokens, padding, truncation, max_length, stride, is_split_into_words, pad_to_multiple_of, return_tensors, return_token_type_ids, return_attention_mask, return_overflowing_tokens, return_special_tokens_mask, return_offsets_mapping, return_length, verbose, **kwargs) 2463 # Backward compatibility for 'truncation_strategy', 'pad_to_max_length' 2464 padding_strategy, truncation_strategy, max_length, kwargs = self._get_padding_truncation_strategies( 2465 padding=padding, 2466 truncation=truncation, (...) 2470 **kwargs, ... (...) 382 pad_to_multiple_of=pad_to_multiple_of, 383 ) TypeError: batch_text_or_text_pairs has to be a list (got )
出现问题的函数代码
def regular_encode(texts, tokenizer, maxlen=512): """ encodes text for a model """ enc_di = tokenizer.batch_encode_plus( texts, return_token_type_ids=False, padding='max_length', max_length=None ) return np.array(enc_di['input_ids'])
解决方案
核心原因:传入
regular_encode的texts参数不是列表类型(比如是单个字符串、空值或其他非列表结构),但tokenizer.batch_encode_plus要求输入必须是列表格式,哪怕只有一条文本也要包裹在列表里。具体修复:
- 检查调用
regular_encode的代码,确保传入的texts是列表。例如:- 单条文本:改成
regular_encode([single_text], tokenizer) - DataFrame列:用
regular_encode(df['text_column'].tolist(), tokenizer)转换为列表后传入
- 单条文本:改成
- 在
regular_encode函数内增加参数校验,自动处理非列表输入:def regular_encode(texts, tokenizer, maxlen=512): """ encodes text for a model """ # 统一将输入转为列表格式 if not isinstance(texts, list): texts = [texts] enc_di = tokenizer.batch_encode_plus( texts, return_token_type_ids=False, padding='max_length', max_length=maxlen # 使用传入的maxlen参数,替代原代码的None ) return np.array(enc_di['input_ids']) - 修正原函数的
max_length=None问题:原代码会忽略传入的maxlen参数,使用tokenizer默认最大长度,改成max_length=maxlen才能保证文本被截断/填充到指定长度。
- 检查调用
内容的提问来源于stack exchange,提问作者Bonyl Santhmayor
相关产品推荐
相关产品推荐

