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

使用tf.data.Dataset.from_tensor_slices创建DialoGPT微调数据集遇类型错误

解决DialoGPT微调中tf.data.Dataset创建的错误

问题场景

作为机器学习新手,在微调DialoGPT模型时,完成Tokenizer加载、CSV数据读取、数据集划分及文本编码后,使用tf.data.Dataset.from_tensor_slices创建训练数据集时遇到两个错误:

已完成的代码:

tokenizerDialoGPT = AutoTokenizer.from_pretrained("microsoft/DialoGPT-medium")
modelDialoGPT = AutoModelForCausalLM.from_pretrained("microsoft/DialoGPT-medium")

df = pd.read_csv('E:/smart_reply/Test_Dataset_DialoGPT.csv', sep=',', names=["Comment", "Reply"])
Comment_DialoGPT = list(df['Comment'])
Reply_DialoGPT = list(df['Reply'])
Comment_DialoGPT_train, Comment_DialoGPT_test, Reply_DialoGPT_train, Reply_DialoGPT_test = train_test_split(Comment_DialoGPT, Reply_DialoGPT, test_size = 0.20, random_state = 0)
train_encodings_Comment_DialoGPT = tokenizerDialoGPT(Comment_DialoGPT_train, truncation=True)
test_encodings_Comment_DialoGPT = tokenizerDialoGPT(Comment_DialoGPT_test, truncation=True)
train_encodings_Reply_DialoGPT = tokenizerDialoGPT(Reply_DialoGPT_train, truncation=True)
test_encodings_Reply_DialoGPT = tokenizerDialoGPT(Reply_DialoGPT_test, truncation=True)

数据集样例:

>>> df.head()
                         Comment                               Reply
0                             Hi                 Hello! Good Morning
1  Could you please modify this?  Sure! I will do that. But, why so?
2          Are you sure, or not?                            Yes. No.
3          What will be the MAU?         Hard to predict. Can't say.
4             Looking good to me                      Great. Thanks.
>>>

尝试的代码及错误

  1. 代码:
train_dataset_dialoGPT = tf.data.Dataset.from_tensor_slices( (dict(train_encodings_Comment_DialoGPT), dict(train_encodings_Reply_DialoGPT)))

错误:无法将非矩形Python序列转换为Tensor

  1. 代码:
train_dataset_dialoGPT = tf.data.Dataset.from_tensor_slices( [dict(train_encodings_Comment_DialoGPT), dict(train_encodings_Reply_DialoGPT)])

错误:ValueError: 不支持将dict类型转换为Tensor

错误原因

  • 第一个错误:Tokenizer仅设置了truncation=True,未做padding处理,导致生成的input_ids是变长序列(每个样本长度不同),而tf.data.Dataset.from_tensor_slices要求输入是矩形张量(所有样本维度一致)。
  • 第二个错误:from_tensor_slices不接受包含字典的列表作为输入,它仅支持张量、数组,或值为张量/数组的字典。

解决方法

步骤1:修正Tokenization,添加Padding

在调用Tokenizer时添加padding=True,确保所有序列长度统一:

train_encodings_Comment_DialoGPT = tokenizerDialoGPT(Comment_DialoGPT_train, truncation=True, padding=True)
test_encodings_Comment_DialoGPT = tokenizerDialoGPT(Comment_DialoGPT_test, truncation=True, padding=True)
train_encodings_Reply_DialoGPT = tokenizerDialoGPT(Reply_DialoGPT_train, truncation=True, padding=True)
test_encodings_Reply_DialoGPT = tokenizerDialoGPT(Reply_DialoGPT_test, truncation=True, padding=True)

步骤2:构造符合模型要求的数据集

DialoGPT作为因果语言模型,训练时需要input_ids、attention_mask以及labels(目标输出)。这里提供两种常见构造方式:

方式1:将Comment作为输入,Reply作为标签

直接把用户评论作为模型输入,回复作为训练目标:

import tensorflow as tf

# 构造模型训练所需的字典
train_dataset_dict = {
    'input_ids': train_encodings_Comment_DialoGPT['input_ids'],
    'attention_mask': train_encodings_Comment_DialoGPT['attention_mask'],
    'labels': train_encodings_Reply_DialoGPT['input_ids']
}

# 创建数据集并做预处理
train_dataset_dialoGPT = tf.data.Dataset.from_tensor_slices(train_dataset_dict)
# 打乱+分批,batch大小可根据显存调整
train_dataset_dialoGPT = train_dataset_dialoGPT.shuffle(1000).batch(8)

方式2:拼接Comment与Reply(更贴合对话模型训练逻辑)

将用户评论和回复用DialoGPT的结束符<|endoftext|>拼接,作为完整输入,同时将整个序列作为标签(因果语言模型会自动学习从上下文生成回复):

# 拼接评论与回复,添加模型专用分隔符
train_pairs = [f"{comment}{tokenizerDialoGPT.eos_token}{reply}" for comment, reply in zip(Comment_DialoGPT_train, Reply_DialoGPT_train)]

# 对拼接后的文本编码,设置统一长度
train_encodings = tokenizerDialoGPT(train_pairs, truncation=True, padding=True, max_length=128)

# 构造数据集
train_dataset_dict = {
    'input_ids': train_encodings['input_ids'],
    'attention_mask': train_encodings['attention_mask'],
    'labels': train_encodings['input_ids']
}

train_dataset_dialoGPT = tf.data.Dataset.from_tensor_slices(train_dataset_dict).shuffle(1000).batch(8)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.20 14:48:24