You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

请求生成将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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.23 04:54:16