Keras多输入带生成器模型问题:图像与元数据匹配训练
解决Keras多输入回归模型中图像与元数据批量匹配的问题
嘿,我明白你现在的核心困扰——怎么让图像生成器输出的每一批图像,都精准对应到数据框里同一行的表格元数据对吧?这确实是多输入模型拟合时的常见卡点,我来一步步帮你调整代码搞定它:
问题根源拆解
你当前用的flow_images_from_dataframe默认只输出图像和目标值,但我们需要同时拿到对应批次的元数据关联信息。最靠谱的方式是让生成器返回图像对应的行索引,再通过索引从数据框里精准提取元数据。
修正后的完整代码方案
1. 调整图像生成器,保留行索引关联
先修改图像生成器的配置,让它返回每个图像对应的行索引,后续我们用这个索引匹配元数据:
# 图像生成器函数(修改为返回行索引,用于匹配元数据) train_datagen <- image_data_generator(rescale = 1/255) train_generator <- flow_images_from_dataframe( dataframe = joined_path_with_metadata, directory = 'data_dir', x_col = "path", # 临时将y_col设为数据框的行名(后续转成索引用) y_col = rownames(joined_path_with_metadata), generator = train_datagen, target_size = c(150, 150), batch_size = 20, color_mode = 'rgb', class_mode = "sparse" # 因为y_col是索引,用sparse返回整数格式 )
2. 自定义组合生成器
接下来写一个自定义生成器,把图像、对应元数据、目标值打包成模型需要的输入格式:
# 自定义组合生成器:把图像和对应元数据配对 combined_generator <- function(image_gen, metadata_df) { function() { # 获取图像批次和对应的行索引 batch <- image_gen %>% next_element() images_batch <- batch[[1]] row_indices <- batch[[2]] %>% as.integer() # 提取对应行的元数据(只取feature1和feature2,排除path和target) metadata_batch <- metadata_df[row_indices, c("feature1", "feature2")] %>% as.matrix() # 提取对应行的目标值 targets_batch <- metadata_df[row_indices, "target"] %>% as.matrix() # 返回模型需要的输入列表 + 目标值 list(list(images_batch, metadata_batch), targets_batch) } } # 初始化训练用的组合生成器 train_combined_gen <- combined_generator(train_generator, joined_path_with_metadata)
3. 修正表格输入层的形状
你之前的表格输入层shape = ncol(dataframe)不对,因为我们只传入2个特征(feature1和feature2),所以要调整:
# 表格数据ANN模型(修正输入形状) tabular_input <- layer_input(shape = 2, dtype = 'float32') # 对应2个特征 mlp_model <- tabular_input %>% layer_dense(units = 16, kernel_initializer = "uniform", activation = "relu") %>% layer_dropout(rate = 0.1) %>% layer_dense(units = 32, kernel_initializer = "uniform", activation = "relu")
4. 模型拟合
现在就可以用自定义组合生成器来训练模型了,记得指定steps_per_epoch(总样本数除以批量大小):
# 计算每个epoch需要跑多少步 steps_per_epoch <- ceiling(nrow(joined_path_with_metadata) / 20) # 拟合模型(如果是新版本Keras,可用fit替代fit_generator) history <- vqa_model %>% fit_generator( generator = train_combined_gen, steps_per_epoch = steps_per_epoch, epochs = 20 )
关键细节提醒
- 这个方案的核心是用行索引做桥梁,确保每一张图像和它的元数据100%对应,完全避免批量错位问题。
- 如果需要验证集,只要复制上述步骤,创建验证集的图像生成器和组合生成器即可。
- 若你用的是较新的Keras版本,
fit_generator已被fit替代,直接把代码改成vqa_model %>% fit(x = train_combined_gen, ...)就行。
内容的提问来源于stack exchange,提问作者Henryk Borzymowski
相关产品推荐
相关产品推荐

