衣物颜色识别项目:ImageDataGenerator预处理参数异常排查
问题原因
ImageDataGenerator配合flow_from_dataframe使用时,默认会先将图像读取为numpy数组,再传入预处理函数。而你的预处理逻辑基于文件路径实现(需要读取文件做YOLO检测),因此类型不匹配触发Expected file_path to be str...错误。
解决方案
方案1:自定义数据生成器(推荐)
完全掌控数据加载流程,直接从DataFrame获取文件路径,完成YOLO检测、裁剪等预处理:
import numpy as np import cv2 from tensorflow.keras.utils import Sequence # 假设你已实现YOLOv3检测类 class YOLOv3Detector: def __init__(self): # 初始化YOLOv3模型权重、参数等 pass def detect_human(self, img_rgb): # 输入RGB格式的numpy数组,返回单个人体的边界框[x1, y1, x2, y2] pass class ColorRecogGenerator(Sequence): def __init__(self, df, img_dir, batch_size, target_size, label_map): self.df = df self.img_dir = img_dir self.batch_size = batch_size self.target_size = target_size self.label_map = label_map # 颜色到整数标签的映射,如{'red':0, 'blue':1} self.yolo = YOLOv3Detector() self.indexes = np.arange(len(self.df)) def __len__(self): # 返回每个epoch的步数 return int(np.ceil(len(self.df) / self.batch_size)) def __getitem__(self, idx): # 生成单个batch的数据 batch_idx = self.indexes[idx*self.batch_size : (idx+1)*self.batch_size] batch_data = self.df.iloc[batch_idx] X_batch = [] y_batch = [] for _, row in batch_data.iterrows(): # 读取原始图像 img_path = f"{self.img_dir}/{row['img_name']}" img_bgr = cv2.imread(img_path) img_rgb = cv2.cvtColor(img_bgr, cv2.COLOR_BGR2RGB) # YOLO检测并裁剪人体 x1, y1, x2, y2 = self.yolo.detect_human(img_rgb) cropped_img = img_rgb[y1:y2, x1:x2] # 调整尺寸并归一化 cropped_img = cv2.resize(cropped_img, self.target_size) cropped_img = cropped_img / 255.0 X_batch.append(cropped_img) # 处理标签 color_label = self.label_map[row['color']] y_batch.append(color_label) return np.array(X_batch), np.array(y_batch)
使用方式:
# 假设已加载训练集DataFrame和标签映射 train_generator = ColorRecogGenerator( df=train_df, img_dir="./train_images", batch_size=32, target_size=(224, 224), label_map={'black':0, 'white':1, 'red':2} ) # 模型训练 model.fit(train_generator, epochs=15)
方案2:修改预处理函数适配numpy数组
如果你的YOLO检测逻辑支持直接处理numpy数组,可调整预处理函数,接收ImageDataGenerator传入的数组而非文件路径:
# 全局初始化YOLO检测器(避免每次预处理重复初始化) yolo_detector = YOLOv3Detector() def preprocess_image(img_array): # img_array是ImageDataGenerator加载的RGB格式数组 # 转换为YOLO需要的格式(如果你的YOLO用BGR) img_bgr = cv2.cvtColor(img_array, cv2.COLOR_RGB2BGR) # 检测人体并裁剪 x1, y1, x2, y2 = yolo_detector.detect_human(img_array) cropped_img = img_array[y1:y2, x1:x2] # 调整尺寸和归一化 cropped_img = cv2.resize(cropped_img, (224, 224)) return cropped_img / 255.0
使用方式:
from tensorflow.keras.preprocessing.image import ImageDataGenerator datagen = ImageDataGenerator(preprocessing_function=preprocess_image) train_generator = datagen.flow_from_dataframe( dataframe=train_df, directory="./train_images", x_col="img_name", y_col="color", target_size=(224, 224), batch_size=32, class_mode="categorical" # 根据你的任务类型调整,如"sparse" )
内容的提问来源于stack exchange,提问作者Manas Bisht
相关产品推荐
相关产品推荐

