使用Pandas实现列转行适配Keras flow_from_dataframe categorical模式
Pandas透视转换适配tf.keras ImageDataGenerator flow_from_dataframe categorical模式方案
前置说明
默认你的原始DataFrame(命名为df_raw)为长表结构,包含filepath(图片存储路径)、label(图片所属类别)两个核心列。
转换操作代码
你需要的“label类别转为列、文件路径为行”的结构本质是对label做独热编码,直接使用Pandas的get_dummies方法即可完成,无需复杂透视操作:
import pandas as pd # 对label列做独热编码,自动生成对应类别列 df_processed = pd.get_dummies(df_raw, columns=['label'], prefix='', prefix_sep='')
如果是多标签分类场景(单张图片对应多个标签,原始数据中同一路径对应多行不同label),额外增加合并步骤:
# 合并同一路径的多标签结果 df_processed = df_processed.groupby('filepath', as_index=False).max()
转换完成后的DataFrame结构示例:
| filepath | cat | dog | bird |
|---|---|---|---|
| ./train/cat/001.jpg | 1 | 0 | 0 |
| ./train/dog/003.jpg | 0 | 1 | 0 |
| ./train/bird/007.jpg | 0 | 0 | 1 |
flow_from_dataframe调用示例
from tensorflow.keras.preprocessing.image import ImageDataGenerator # 初始化数据生成器 datagen = ImageDataGenerator(rescale=1./255) # 获取所有类别列名 class_list = df_processed.columns.drop('filepath').tolist() train_generator = datagen.flow_from_dataframe( dataframe=df_processed, x_col='filepath', y_col=class_list, target_size=(224, 224), batch_size=32, class_mode='categorical' )
常见问题修复
- 此前用
pd.pivot转换报错是因为pivot要求filepath无重复值,换用上述get_dummies方案即可解决 - 若调用生成器时输出维度异常,检查
y_col是否传入了全类别列的列表,不要仅传入单个列名
内容的提问来源于stack exchange,提问作者plutolaser
相关产品推荐
相关产品推荐

