请求生成将Train/Test图片文件夹与CSV标签转为CNN输入数据集的示例代码
基于CSV标签的图片数据集转为CNN输入格式示例代码
以下提供TensorFlow/Keras和PyTorch两种主流框架的实现方案,适配你的Train/Test独立文件夹+CSV标签的数据集结构(6个类别)。
TensorFlow/Keras 实现示例
核心逻辑
- 读取CSV标签文件,构建图片名称→标签的映射字典
- 使用
tf.data.Dataset加载图片路径,结合映射字典获取对应标签 - 对图片进行预处理(缩放、归一化、数据增强可选)
- 生成模型可直接接收的批量数据集
import pandas as pd import tensorflow as tf import os # 配置参数 IMAGE_SIZE = (224, 224) # 根据你的CNN模型输入尺寸调整 BATCH_SIZE = 32 NUM_CLASSES = 6 TRAIN_DIR = "./Train" # 替换为你的Train文件夹路径 TEST_DIR = "./Test" # 替换为你的Test文件夹路径 TRAIN_CSV = "./train_labels.csv" # 替换为你的训练集CSV路径 TEST_CSV = "./test_labels.csv" # 替换为你的测试集CSV路径 # 1. 读取CSV标签,构建映射 def load_label_mapping(csv_path): df = pd.read_csv(csv_path, header=None) # 假设CSV无表头,第一列是图片名,第二列是标签 label_map = dict(zip(df[0], df[1])) # 若标签是字符串,转为整数编码(标签已是整数可跳过) unique_labels = sorted(df[1].unique()) str_to_int = {label: idx for idx, label in enumerate(unique_labels)} label_map = {k: str_to_int[v] for k, v in label_map.items()} return label_map train_label_map = load_label_mapping(TRAIN_CSV) test_label_map = load_label_mapping(TEST_CSV) # 2. 定义图片加载与预处理函数 def load_and_preprocess_image(image_path, label_map): # 获取图片文件名 img_name = tf.strings.split(image_path, os.sep)[-1] # 获取对应标签 label = label_map[img_name.numpy().decode()] # 加载图片 img = tf.io.read_file(image_path) img = tf.image.decode_jpeg(img, channels=3) # PNG格式用decode_png img = tf.image.resize(img, IMAGE_SIZE) img = tf.cast(img, tf.float32) / 255.0 # 归一化到[0,1] return img, label # 3. 构建数据集 def create_dataset(image_dir, label_map): # 获取所有图片路径 image_paths = tf.data.Dataset.list_files(f"{image_dir}/*") # 加载图片和标签 dataset = image_paths.map(lambda x: tf.py_function( func=load_and_preprocess_image, inp=[x, label_map], Tout=[tf.float32, tf.int32] )) # 打乱、分批、预取 dataset = dataset.shuffle(buffer_size=1000).batch(BATCH_SIZE).prefetch(tf.data.AUTOTUNE) # 转换标签为one-hot编码(模型用categorical_crossentropy时需要) dataset = dataset.map(lambda x, y: (x, tf.one_hot(y, depth=NUM_CLASSES))) return dataset train_dataset = create_dataset(TRAIN_DIR, train_label_map) test_dataset = create_dataset(TEST_DIR, test_label_map) # 验证数据集格式 for img_batch, label_batch in train_dataset.take(1): print(f"图片批量形状: {img_batch.shape}") print(f"标签批量形状: {label_batch.shape}")
PyTorch 实现示例
核心逻辑
- 自定义
Dataset子类,实现从CSV读取标签、加载图片的逻辑 - 使用
DataLoader生成批量数据 - 结合
transforms完成图片预处理(缩放、归一化等)
import pandas as pd import os from PIL import Image import torch from torch.utils.data import Dataset, DataLoader from torchvision import transforms # 配置参数 IMAGE_SIZE = (224, 224) BATCH_SIZE = 32 NUM_CLASSES = 6 TRAIN_DIR = "./Train" TEST_DIR = "./Test" TRAIN_CSV = "./train_labels.csv" TEST_CSV = "./test_labels.csv" # 1. 自定义数据集类 class ImageDataset(Dataset): def __init__(self, image_dir, csv_path, transform=None): self.image_dir = image_dir self.transform = transform # 读取CSV标签 df = pd.read_csv(csv_path, header=None) self.img_names = df[0].tolist() self.labels = df[1].tolist() # 标签字符串转整数编码(标签已是整数可跳过) unique_labels = sorted(set(self.labels)) self.str_to_int = {label: idx for idx, label in enumerate(unique_labels)} self.labels = [self.str_to_int[label] for label in self.labels] def __len__(self): return len(self.img_names) def __getitem__(self, idx): img_name = self.img_names[idx] img_path = os.path.join(self.image_dir, img_name) image = Image.open(img_path).convert("RGB") label = self.labels[idx] if self.transform: image = self.transform(image) return image, torch.tensor(label, dtype=torch.long) # 2. 定义预处理变换 transform = transforms.Compose([ transforms.Resize(IMAGE_SIZE), transforms.ToTensor(), # 转换为Tensor,自动归一化到[0,1] transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) # 可自定义均值方差 ]) # 3. 创建数据集与数据加载器 train_dataset = ImageDataset(TRAIN_DIR, TRAIN_CSV, transform=transform) test_dataset = ImageDataset(TEST_DIR, TEST_CSV, transform=transform) train_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True, num_workers=4) test_loader = DataLoader(test_dataset, batch_size=BATCH_SIZE, shuffle=False, num_workers=4) # 验证数据集格式 for img_batch, label_batch in train_loader: print(f"图片批量形状: {img_batch.shape}") print(f"标签批量形状: {label_batch.shape}") break
注意事项
- 若你的标签已经是0-5的整数编码,可去掉代码中标签字符串转整数的部分
- 图片尺寸
IMAGE_SIZE需与你的CNN模型输入尺寸匹配 - 可根据需求添加数据增强(比如TensorFlow的
tf.image.random_flip_left_right,PyTorch的transforms.RandomHorizontalFlip)
内容的提问来源于stack exchange,提问作者Md Towfiqur Rahman
相关产品推荐
相关产品推荐

