如何为基于RNN的图像序列二分类任务导入数据集?
PyTorch 实现方案
1. 自定义数据集类
继承torch.utils.data.Dataset实现专属数据集类,负责读取序列文件夹、加载图像并返回序列与对应标签:
import os import torch from torch.utils.data import Dataset, DataLoader from PIL import Image from torchvision import transforms class ImageSequenceDataset(Dataset): def __init__(self, root_dir, label_map, transform=None): self.root_dir = root_dir self.label_map = label_map # 示例:{'film_1':0, 'film_2':1},需按实际标签规则定义 self.transform = transform self.sequence_folders = [f for f in os.listdir(root_dir) if os.path.isdir(os.path.join(root_dir, f))] def __len__(self): return len(self.sequence_folders) def __getitem__(self, idx): seq_folder = self.sequence_folders[idx] # 按文件名排序保证序列时序正确 img_paths = sorted([os.path.join(self.root_dir, seq_folder, f) for f in os.listdir(os.path.join(self.root_dir, seq_folder)) if f.endswith(('.png', '.jpg', '.jpeg'))]) # 加载并预处理图像 sequence = [] for img_path in img_paths: img = Image.open(img_path).convert('L') # 转灰度图,RGB则去掉.convert('L') if self.transform: img = self.transform(img) sequence.append(img) # 堆叠成 (time, channels, width, height) 张量 sequence = torch.stack(sequence) label = self.label_map[seq_folder] return sequence, label
2. 数据加载与划分
定义预处理规则,划分训练/验证集并通过DataLoader批量加载:
# 图像预处理:统一尺寸、转张量 transform = transforms.Compose([ transforms.Resize((64, 64)), transforms.ToTensor() ]) # 构建标签映射,需根据你的二分类规则调整 label_map = {} # 示例:film_1至film_100为类别0,film_101至film_200为类别1 for i in range(1, 101): label_map[f'film_{i}'] = 0 for i in range(101, 201): label_map[f'film_{i}'] = 1 # 创建完整数据集 full_dataset = ImageSequenceDataset(root_dir='你的根文件夹路径', label_map=label_map, transform=transform) # 8:2划分训练/验证集 train_size = int(0.8 * len(full_dataset)) val_size = len(full_dataset) - train_size train_dataset, val_dataset = torch.utils.data.random_split(full_dataset, [train_size, val_size]) # DataLoader:自动批量加载、打乱训练集 train_loader = DataLoader(train_dataset, batch_size=8, shuffle=True) val_loader = DataLoader(val_dataset, batch_size=8, shuffle=False)
DataLoader核心作用:自动将数据集按指定批次拆分,shuffle=True打乱训练集避免过拟合,返回的每个批次形状为(batch_size, time_steps, channels, width, height)
3. RNN模型与训练
模型定义
import torch.nn as nn class SequenceClassifier(nn.Module): def __init__(self, input_size, hidden_size, num_layers, num_classes): super().__init__() self.rnn = nn.RNN(input_size, hidden_size, num_layers, batch_first=True) self.fc = nn.Linear(hidden_size, num_classes) def forward(self, x): # 将图像展平为向量:(batch, time, 64*64) x = x.view(x.size(0), x.size(1), -1) # 初始化隐藏层 h0 = torch.zeros(self.num_layers, x.size(0), self.hidden_size).to(x.device) # RNN前向传播,取最后一个时间步输出 out, _ = self.rnn(x, h0) out = self.fc(out[:, -1, :]) return out # 初始化模型 input_size = 64*64 # 对应64x64灰度图的展平尺寸 model = SequenceClassifier(input_size, hidden_size=128, num_layers=2, num_classes=2)
训练循环
criterion = nn.CrossEntropyLoss() optimizer = torch.optim.Adam(model.parameters(), lr=0.001) device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') model.to(device) num_epochs = 10 for epoch in range(num_epochs): # 训练阶段 model.train() train_loss = 0.0 for sequences, labels in train_loader: sequences, labels = sequences.to(device), labels.to(device) outputs = model(sequences) loss = criterion(outputs, labels) optimizer.zero_grad() loss.backward() optimizer.step() train_loss += loss.item() * sequences.size(0) train_loss /= len(train_loader.dataset) # 验证阶段 model.eval() val_loss = 0.0 correct = 0 with torch.no_grad(): for sequences, labels in val_loader: sequences, labels = sequences.to(device), labels.to(device) outputs = model(sequences) loss = criterion(outputs, labels) val_loss += loss.item() * sequences.size(0) _, predicted = torch.max(outputs.data, 1) correct += (predicted == labels).sum().item() val_loss /= len(val_loader.dataset) val_acc = correct / len(val_loader.dataset) print(f'Epoch [{epoch+1}/{num_epochs}], Train Loss: {train_loss:.4f}, Val Loss: {val_loss:.4f}, Val Acc: {val_acc:.4f}')
Keras 实现方案
1. 数据加载与划分
用tf.data.Dataset构建序列数据集:
import tensorflow as tf from tensorflow.keras import layers, models import os def load_sequence(seq_folder, label_map, img_size=(64,64)): # 按文件名排序加载图像 img_paths = sorted([os.path.join(seq_folder, f) for f in os.listdir(seq_folder) if f.endswith(('.png', '.jpg', '.jpeg'))]) sequence = [] for path in img_paths: img = tf.io.read_file(path) img = tf.image.decode_jpeg(img, channels=1) # 灰度图用1,RGB用3 img = tf.image.resize(img, img_size) img = img / 255.0 # 归一化到0-1 sequence.append(img) sequence = tf.stack(sequence) label = label_map[os.path.basename(seq_folder)] return sequence, label # 标签映射,同PyTorch逻辑 label_map = {} for i in range(1, 101): label_map[f'film_{i}'] = 0 for i in range(101, 201): label_map[f'film_{i}'] = 1 # 划分训练/验证集文件夹 root_dir = '你的根文件夹路径' seq_folders = [os.path.join(root_dir, f) for f in os.listdir(root_dir) if os.path.isdir(os.path.join(root_dir, f))] train_size = int(0.8 * len(seq_folders)) train_folders, val_folders = seq_folders[:train_size], seq_folders[train_size:] # 构建训练/验证数据集 train_dataset = tf.data.Dataset.from_tensor_slices(train_folders) train_dataset = train_dataset.map(lambda x: load_sequence(x, label_map), num_parallel_calls=tf.data.AUTOTUNE) train_dataset = train_dataset.shuffle(100).batch(8).prefetch(tf.data.AUTOTUNE) val_dataset = tf.data.Dataset.from_tensor_slices(val_folders) val_dataset = val_dataset.map(lambda x: load_sequence(x, label_map), num_parallel_calls=tf.data.AUTOTUNE) val_dataset = val_dataset.batch(8).prefetch(tf.data.AUTOTUNE)
2. RNN模型与训练
模型定义
input_shape = (None, 64, 64, 1) # None支持可变长度序列 model = models.Sequential([ layers.TimeDistributed(layers.Flatten(), input_shape=input_shape), layers.SimpleRNN(128, return_sequences=False), layers.Dense(64, activation='relu'), layers.Dense(2, activation='softmax') ]) model.compile(optimizer='adam', loss='sparse_categorical_crossentropy', metrics=['accuracy'])
训练模型
history = model.fit(train_dataset, epochs=10, validation_data=val_dataset)
关键注意事项
- 序列时序:必须按文件名排序加载图像,否则序列顺序混乱会直接影响RNN效果
- 图像尺寸统一:所有图像需resize到相同尺寸,否则无法堆叠成张量
- 可变序列长度:若不同序列的图像数量不一致,需用padding/truncate统一长度,PyTorch可使用
torch.nn.utils.rnn.pad_sequence,Keras可使用tf.keras.preprocessing.sequence.pad_sequences - 标签映射:需根据你的二分类规则自定义,若有标签csv文件,可先读取文件生成标签字典
内容的提问来源于stack exchange,提问作者user25720613
相关产品推荐
相关产品推荐

