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

使用Hugging Face Trainer训练CodeT5-small(eth_py150_open数据集)遇TypeError求助

问题:训练CodeT5-small时出现TypeError: can only join an iterable错误

错误日志

***** Running training *****
  Num examples = 74749
  Num Epochs = 12
  Instantaneous batch size per device = 8
  Total train batch size (w. parallel, distributed & accumulation) = 8
  Gradient Accumulation steps = 1
  Total optimization steps = 112128
---------------------------------------------------------------------------
TypeError                                 Traceback (most recent call last)
<ipython-input-28-3435b262f1ae> in <module>
----> 1 trainer.train()

3 frames
/usr/local/lib/python3.7/dist-packages/transformers/trainer.py in _prepare_inputs(self, inputs)
   2414         if len(inputs) == 0:
   2415             raise ValueError(
-> 2416                 "The batch received was empty, your model won't be able to train on it. Double-check that your "
   2417                 f"training dataset contains keys expected by the model: {','.join(self._signature_columns)}."
   2418             )

TypeError: can only join an iterable

相关代码

import torch
import transformers
from datasets import load_dataset_builder
from datasets import load_dataset

corpus=load_dataset("eth_py150_open", split='train')

training_args = transformers.TrainingArguments( #general training arguments
    per_device_train_batch_size = 8,
    warmup_steps = 0,
    weight_decay = 0.01,
    learning_rate = 1e-4,
    num_train_epochs = 12,
    output_dir = './runs/run2/output/',
    logging_dir = './runs/run2/logging/',
    logging_steps = 50,
    save_steps= 10000,
    remove_unused_columns=False,
)

model = transformers.T5ForConditionalGeneration.from_pretrained('Salesforce/codet5-small').cuda()

trainer = transformers.Trainer(
    model = model,
   args = training_args,
    train_dataset = corpus,
)

解决方案

这个错误的核心原因是:你的原始数据集没有经过预处理,缺少CodeT5模型要求的输入列(input_ids、attention_mask、labels),加上remove_unused_columns=False的设置,导致Trainer无法识别有效输入列,最终触发空迭代器的拼接错误。

以下是修正后的完整代码,关键步骤已标注:

import torch
import transformers
from datasets import load_dataset

# 1. 加载原始数据集
corpus = load_dataset("eth_py150_open", split='train')

# 2. 加载对应tokenizer和模型(必须配套使用)
tokenizer = transformers.AutoTokenizer.from_pretrained('Salesforce/codet5-small')
model = transformers.T5ForConditionalGeneration.from_pretrained('Salesforce/codet5-small').cuda()

# 3. 定义预处理函数:将原始代码转换为模型可识别的格式
def preprocess_function(examples):
    # 对代码文本进行tokenize,生成input_ids和attention_mask
    inputs = tokenizer(
        examples["code"],
        padding="max_length",
        truncation=True,
        max_length=512
    )
    # T5模型需要labels,且需将padding token替换为-100(损失计算时会忽略这些位置)
    labels = tokenizer(
        examples["code"],
        padding="max_length",
        truncation=True,
        max_length=512
    )
    labels["input_ids"] = [
        [(l if l != tokenizer.pad_token_id else -100) for l in label] 
        for label in labels["input_ids"]
    ]
    
    # 将labels加入输入字典
    inputs["labels"] = labels["input_ids"]
    return inputs

# 4. 批量预处理数据集
processed_corpus = corpus.map(preprocess_function, batched=True)

# 5. 调整训练参数:移除remove_unused_columns=False(默认True会自动保留模型需要的列)
training_args = transformers.TrainingArguments(
    per_device_train_batch_size=8,
    warmup_steps=0,
    weight_decay=0.01,
    learning_rate=1e-4,
    num_train_epochs=12,
    output_dir='./runs/run2/output/',
    logging_dir='./runs/run2/logging/',
    logging_steps=50,
    save_steps=10000,
)

# 6. 初始化Trainer并开始训练
trainer = transformers.Trainer(
    model=model,
    args=training_args,
    train_dataset=processed_corpus,
)

trainer.train()

关键说明

  • 无需手动转换为torch Dataset:Hugging Face的Dataset对象经过map处理后可直接给Trainer使用
  • T5模型的labels必须特殊处理:将padding token替换为-100,否则会影响损失计算
  • 保留默认的remove_unused_columns=True:Trainer会自动过滤掉模型不需要的列(比如原始数据集中的repo列),避免空输入问题

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.15 16:25:38