如何在tf.Dataset上正确适配TextVectorization层
报错原因
tf.data.experimental.make_csv_dataset 加载后返回的数据集结构为 (特征字典, 标签张量) 的二元组:其中特征字典的key是CSV表头的列名,value是对应列的批次张量;标签是你指定label_name="tags"对应的tags列张量。
直接把这个完整数据集传给vectorizer.adapt()时,TextVectorization层接收到的输入是嵌套的元组+字典结构,而非它要求的纯字符串文本张量,因此触发类型错误。
正确适配步骤
- 从原始训练数据集中单独抽取出需要做向量化的文本列,剔除标签和其他无关特征列。先确认自己CSV中存储待处理文本的列名,替换下方代码里的占位符即可:
# 替换为你的CSV里实际的文本列名,比如文本列叫code就写x["code"] text_only_ds = data_provider.train_ds.map(lambda feature_batch, label_batch: feature_batch["your_text_column_name"])
- (可选优化)给文本数据集加预取,加快adapt阶段的处理速度:
text_only_ds = text_only_ds.prefetch(tf.data.AUTOTUNE)
- 传入纯文本数据集执行适配:
vectorizer.adapt(text_only_ds)
注意事项
- 如果你有多列文本需要输入同一个TextVectorization层处理,可以在
map步骤中将多列文本拼接为单个字符串张量后输出,不要保留字典结构 adapt方法不会自动拆分(特征,标签)元组,也不会自动从特征字典中筛选文本列,所有传入adapt的数据集必须保证每一批输出直接是待处理的字符串张量- 适配阶段不要对文本做额外的数值转换,保持原始字符串类型输出即可
内容的提问来源于stack exchange,提问作者1231fdgsg78sdf7
相关产品推荐
相关产品推荐

