如何修改TensorFlow-Keras代码实现单epoch内多角度图像旋转训练
实现单轮次内遍历所有旋转角度的训练需求
我来帮你调整代码实现这个需求!核心思路是把角度切换的时机从epoch结束改成每个batch训练时,这样一个epoch内就能依次遍历所有指定的旋转角度,同时保持同一轮次持续训练。
修改要点
- 移除原有的
CustomCallback:不再需要在epoch结束时触发角度切换 - 调整
CIFAR10Sequence的逻辑:- 删掉专门的
change_angle方法 - 在
__getitem__中,根据当前的batch索引计算要使用的旋转角度,通过idx % len(self.angles)实现四个角度的循环遍历
- 删掉专门的
修改后的完整代码
from skimage.io import imread from skimage.transform import resize, rotate import numpy as np import tensorflow as tf from tensorflow import keras from tensorflow.keras import layers from keras.utils import Sequence from keras.models import Sequential from keras.layers import Conv2D, Activation, Flatten, Dense # Model architecture (dummy) model = Sequential() model.add(Conv2D(32, (3, 3), input_shape=(15, 15, 4))) model.add(Activation('relu')) model.add(Flatten()) model.add(Dense(1)) model.add(Activation('sigmoid')) model.compile(loss='binary_crossentropy', optimizer='rmsprop', metrics=['accuracy']) # Data iterator - 修改后版本 class CIFAR10Sequence(Sequence): def __init__(self, filenames, labels, batch_size): self.filenames, self.labels = filenames, labels self.batch_size = batch_size self.angles = [0,90,180,270] def __len__(self): return int(np.ceil(len(self.filenames) / float(self.batch_size))) # 每个batch切换一次旋转角度,同一epoch内依次遍历四个角度 def __getitem__(self, idx): # 根据batch索引计算当前要用的角度,循环遍历四个选项 angle_idx = idx % len(self.angles) angle = self.angles[angle_idx] print(f"Rotating Angle: {angle}") batch_x = self.filenames[idx * self.batch_size:(idx + 1) * self.batch_size] batch_y = self.labels[idx * self.batch_size:(idx + 1) * self.batch_size] return (np.array([rotate(resize(imread(filename), (15, 15)), angle) for filename in batch_x]), np.array(batch_y)) # Create data reader sequence = CIFAR10Sequence(["f1.PNG"]*10, [0, 1]*5, 8) # 不再需要自定义回调,直接训练即可 model.fit(sequence, epochs=10)
运行效果说明
每个epoch训练时,会按顺序输出角度日志(以你的测试数据为例,一个epoch有2个batch):
Rotating Angle: 0 Rotating Angle: 90
如果你的数据集更大、一个epoch包含更多batch,日志会依次输出0→90→180→270→0→90...,完美实现单轮次内依次遍历所有指定旋转角度的需求。
内容的提问来源于stack exchange,提问作者maubere
相关产品推荐
相关产品推荐

