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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.06 05:00:55