如何在Keras CNN模型中结合SVM分类器实现癫痫发作预测?
CNN特征提取+SVM分类实现方案
1. 改造CNN为特征提取器
你的现有CNN模型仅包含卷积、池化和批归一化层,没有最终分类层。先完成CNN的训练(保持你原有的fit_generator流程即可),之后构建一个新模型,专门输出CNN最后一层的特征向量,作为后续SVM的输入:
from keras.models import Model # 基于训练好的CNN构建特征提取器 feature_extractor = Model(inputs=model.input, outputs=model.layers[-1].output) # 若后续给CNN添加了分类Dense层,需改为取倒数第二层输出:model.layers[-2].output
2. 从生成器提取特征与标签
因为数据通过生成器加载,没有现成的x/y数组,需要遍历生成器逐个batch提取特征和标签,拼接成SVM可用的数据集:
import numpy as np # 提取训练集特征与标签 train_gen = generate_arrays_for_training(indexPat, train_data, start=0, end=100) train_feats = [] train_labels = [] # 遍历步数与fit_generator的steps_per_epoch一致 for _ in range(int(len(train_data)/2)): x_batch, y_batch = next(train_gen) # 提取当前batch的特征 batch_feats = feature_extractor.predict(x_batch, verbose=0) # 展平特征为一维向量(SVM要求输入为一维) batch_feats_flat = batch_feats.reshape(batch_feats.shape[0], -1) train_feats.append(batch_feats_flat) # 把one-hot标签转为类别索引(二分类下取argmax) train_labels.append(np.argmax(y_batch, axis=1)) # 拼接为完整数组 train_feats = np.concatenate(train_feats, axis=0) train_labels = np.concatenate(train_labels, axis=0) # 同理提取测试集特征与标签 test_gen = generate_arrays_for_training(indexPat, test_data, start=0, end=100) test_feats = [] test_labels = [] for _ in range(int(len(test_data)/2)): x_batch, y_batch = next(test_gen) batch_feats = feature_extractor.predict(x_batch, verbose=0) batch_feats_flat = batch_feats.reshape(batch_feats.shape[0], -1) test_feats.append(batch_feats_flat) test_labels.append(np.argmax(y_batch, axis=1)) test_feats = np.concatenate(test_feats, axis=0) test_labels = np.concatenate(test_labels, axis=0)
3. 标准化特征并训练SVM
SVM对特征尺度敏感,先标准化特征再训练分类器:
from sklearn.svm import SVC from sklearn.preprocessing import StandardScaler # 标准化特征 scaler = StandardScaler() train_feats_scaled = scaler.fit_transform(train_feats) test_feats_scaled = scaler.transform(test_feats) # 初始化并训练SVM(二分类场景,可调整kernel、C等参数优化效果) svm_clf = SVC(kernel='rbf', C=1.0, gamma='scale') svm_clf.fit(train_feats_scaled, train_labels) # 评估模型 train_acc = svm_clf.score(train_feats_scaled, train_labels) test_acc = svm_clf.score(test_feats_scaled, test_labels) print(f"训练集准确率: {train_acc:.4f}, 测试集准确率: {test_acc:.4f}")
注意事项
- 若后续给CNN添加了全连接分类层,特征提取器必须输出分类层之前的隐藏层特征,避免用分类层的输出作为SVM输入。
- 生成器遍历步数必须与
fit_generator中的steps_per_epoch/validation_steps一致,确保覆盖全部数据。 - 可通过网格搜索(
GridSearchCV)调整SVM的kernel、C、gamma等参数,进一步提升分类性能。
内容的提问来源于stack exchange,提问作者Eda
相关产品推荐
相关产品推荐

