CNN+LSTM困倦分类模型验证精度停滞问题求助
困倦分类CNN+LSTM模型训练异常问题
数据集与划分
- 采用SUSTDDD数据集,19名受试者按**13:3:3(70%:15%:15%)**划分为训练/验证/测试集
模型架构
基于冻结权重的预训练VGG16搭建CNN+LSTM视频分类模型,核心代码如下:
参数定义
image_height = 224 image_width = 224 batch_size = 6 n_of_frames = 50 NUM_OF_CLASSES = 2 NUM_OF_CHANNELS = 3
模型构建
input_shape = (n_of_frames, image_height, image_width, NUM_OF_CHANNELS) video = Input(shape=input_shape) cnn_base = VGG16(input_shape=(image_height, image_width, NUM_OF_CHANNELS), weights="imagenet", include_top=False) cnn_base.trainable = False cnn_out = GlobalAveragePooling2D()(cnn_base.output) cnn = Model(cnn_base.input, cnn_out) encoded_frames = TimeDistributed(cnn)(video) encoded_sequence1 = LSTM(1024)(encoded_frames) hidden_layer1 = Dense(512, activation="relu")(encoded_sequence1) hidden_layer2 = Dropout(0.3)(hidden_layer1) outputs = Dense(NUM_OF_CLASSES, activation="softmax")(hidden_layer2) model = Model(video, outputs) optimizer = Adam(learning_rate=0.0001) model.compile(loss='categorical_crossentropy', optimizer=optimizer, metrics=['accuracy']) model.summary()
模型参数统计
Model: "model_1" _________________________________________________________________ Layer (type) Output Shape Param # ================================================================= input_1 (InputLayer) [(None, 50, 224, 224, 3 0 )] time_distributed (TimeDist (None, 50, 512) 14714688 ributed) lstm (LSTM) (None, 1024) 6295552 dense (Dense) (None, 512) 524800 dropout (Dropout) (None, 512) 0 dense_1 (Dense) (None, 2) 1026 ================================================================= Total params: 21536066 (82.15 MB) Trainable params: 6821378 (26.02 MB) Non-trainable params: 14714688 (56.13 MB)
训练异常表现
- 训练精度波动且无法提升
- 验证精度卡在1-2个固定值
- 网络始终只预测单一类别
- 训练曲线显示异常趋势,混淆矩阵呈现单一类别预测特征
自定义数据生成器代码
class CustomSequenceGenerator(tf.keras.utils.Sequence): def __init__(self, directory, batch_size, n_of_frames, target_size, shuffle=True): self.directory = directory self.batch_size = batch_size self.n_of_frames = n_of_frames self.target_size = target_size self.shuffle = shuffle self.skip_step = 1 self.sequences = self.load_filepaths() np.random.shuffle(self.sequences) if self.directory =='/workspace/SUSTDDD/val' or self.directory =='/workspace/SUSTDDD/test': self.sequences = self.balance_dataset(self.sequences, 161) if self.directory =='/workspace/SUSTDDD/train': self.sequences = self.balance_dataset(self.sequences, 2000) print(self.list_sequences(self.sequences)) def load_filepaths(self): filepaths = [] count = 0 total_jpg_files = 0 for clas in os.listdir(self.directory): for vid in os.listdir(os.path.join(self.directory, clas)): single_sequence = [] if vid.endswith(".zip"): continue sub_vids = os.listdir(os.path.join(self.directory, clas, vid)) for sub_vid in sub_vids: files = os.listdir(os.path.join(self.directory, clas, vid, sub_vid)) files.sort() for idx, file in enumerate(files): if file.lower().endswith('.png') or file.lower().endswith('.jpg') or file.lower().endswith('.jpeg'): if idx % self.skip_step == 0: filepath = os.path.join(self.directory, clas, vid, sub_vid, file) single_sequence.append(filepath) if len(single_sequence) == self.n_of_frames: filepaths.append(single_sequence) single_sequence = [] if self.shuffle: np.random.shuffle(filepaths) return filepaths def list_sequences(self, filepaths): classes = {"D": 0, "ND": 0} for idx, filepath in enumerate(filepaths): for k, v in classes.items(): if k == filepath[0].split('/')[4]: classes[k] += 1 return classes def balance_dataset(self, filepaths, max_cap): classes = {"D": 0, "ND": 0} to_remove = [] for idx, filepath in enumerate(filepaths): for k, v in classes.items(): if k == filepath[0].split('/')[4]: classes[k] += 1 if v >= max_cap: to_remove.append(filepath) for j in to_remove: filepaths.remove(j) return filepaths def __len__(self): return len(self.sequences) // self.batch_size def __getitem__(self, idx): batch_sequences = self.sequences[idx * self.batch_size:(idx + 1) * self.batch_size] batch_images = [] batch_labels = [] for j, sequence in enumerate(batch_sequences): batch_single_sequence = [] current_class = -1 for image_path in sequence: tf_image = tf.io.read_file(image_path) decoded_image = tf.image.decode_image(tf_image, dtype=tf.float32) image_resized = tf.image.resize(decoded_image, self.target_size) image_batch = tf.keras.applications.vgg16.preprocess_input(image_resized) batch_single_sequence.append(image_batch) current_class = image_path.split('/')[4] if current_class == "D": current_class = 1 elif current_class == "ND": current_class = 0 batch_images.append(batch_single_sequence) batch_labels.append(current_class) batch_images = np.array(batch_images) batch_labels = np.array(batch_labels) batch_labels = to_categorical(batch_labels, num_classes=NUM_OF_CLASSES) #print(np.shape(batch_images), np.shape(batch_labels)) return batch_images, batch_labels def on_epoch_end(self): if self.shuffle: np.random.shuffle(self.sequences)
已尝试的优化手段
- 调整模型复杂度:替换为3个Conv层+1个LSTM层
- 更换预训练模型:VGG16、VGG19、ResNet50、InceptionV3(匹配对应输入尺寸)
- 调整学习率:范围从0.01到1e-6
- 更换优化器:Adam、SGD、RMSprop
- 更换数据集:RLDD
- 添加正则化:Dropout或Recurrent Dropout
- 损失函数与激活调整:使用Binary Cross Entropy搭配Sigmoid激活函数
- 数据预处理:使用TensorFlow内置归一化工具
- 数据集平衡:手动平衡训练/验证/测试集样本分布
求助需求
目前已尝试上述所有方案,模型仍无法正常收敛,请求进一步的排查方向与解决方案。
内容的提问来源于stack exchange,提问作者user24778149
相关产品推荐
相关产品推荐

