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

如何在基于生成器的Keras CNN癫痫预测模型中集成SVM分类器?

实现CNN+SVM的癫痫发作预测方案

步骤1:重构CNN作为特征提取器

先修改你的CNN模型,移除最后的Softmax分类层,保留到Dense(256, activation='sigmoid')层作为特征输出(也可根据需求选择Flatten层等其他层作为特征源)。同时要补充输入层的形状,确保和生成器输出的数据维度匹配。

from keras.models import Model
from keras.layers import Input, Conv2D, MaxPooling2D, BatchNormalization, Dense, Dropout, Flatten
from keras import regularizers, optimizers

# 定义输入层,替换成你的数据实际形状,比如(64,64,1)
input_layer = Input(shape=(你的输入高度, 你的输入宽度, 通道数))
x = Conv2D(32, (3, 3), strides=(1,1), padding='same', activation='relu')(input_layer)
x = MaxPooling2D(pool_size=(2, 2), padding='same')(x)
x = BatchNormalization()(x)
x = Dense(32, kernel_regularizer=regularizers.l2(0.1))(x)  # 移除input_dim,由Keras自动推断输入维度
x = Dropout(0.6)(x)
x = Flatten()(x)
x = Dropout(0.6)(x)
feature_layer = Dense(256, activation='sigmoid')(x)  # 该层作为特征输出
x = Dropout(0.6)(feature_layer)
output_layer = Dense(2, activation='softmax')(x)

# 构建完整训练模型,先做监督预训练
full_model = Model(inputs=input_layer, outputs=output_layer)
opt_adam = optimizers.Adam(lr=0.0001, beta_1=0.9, beta_2=0.999, epsilon=1e-08, decay=0.0)
full_model.compile(loss=categorical_focal_loss(), optimizer=opt_adam, metrics=['accuracy'])

# 执行原训练流程
history = full_model.fit_generator(
    generate_arrays_for_training(indexPat, train_data, start=0, end=100),
    validation_data=generate_arrays_for_training(indexPat, test_data, start=0, end=100),
    steps_per_epoch=int((len(train_data)/2)),
    validation_steps=int((len(test_data)/2)),
    verbose=2, epochs=65, max_queue_size=2, shuffle=True
)

# 构建特征提取专用模型,去掉最后的分类层
feature_extractor = Model(inputs=input_layer, outputs=feature_layer)

步骤2:从生成器提取特征与标签

通过遍历生成器,批量提取特征和对应的标签,解决无明确X、y的问题:

import numpy as np
from sklearn.svm import SVC
from sklearn.metrics import accuracy_score

def extract_features_labels(generator, total_steps):
    features = []
    labels = []
    for _ in range(total_steps):
        x_batch, y_batch = next(generator)
        # 提取当前批次的特征
        feat_batch = feature_extractor.predict(x_batch, verbose=0)
        features.append(feat_batch)
        # 将one-hot标签转为单分类标签
        labels.append(np.argmax(y_batch, axis=1))
    # 拼接成完整的特征矩阵和标签数组
    return np.concatenate(features, axis=0), np.concatenate(labels, axis=0)

# 提取训练集特征与标签
train_gen = generate_arrays_for_training(indexPat, train_data, start=0, end=100)
X_train_feat, y_train = extract_features_labels(train_gen, steps=int(len(train_data)/2))

# 提取测试集特征与标签
test_gen = generate_arrays_for_training(indexPat, test_data, start=0, end=100)
X_test_feat, y_test = extract_features_labels(test_gen, steps=int(len(test_data)/2))

步骤3:训练并评估SVM分类器

用提取到的特征训练SVM,完成最终分类:

# 初始化SVM,可根据数据集调整核函数、正则参数等
svm_clf = SVC(kernel='rbf', C=1.0, gamma='scale', random_state=42)

# 训练SVM
svm_clf.fit(X_train_feat, y_train)

# 预测并评估
y_pred = svm_clf.predict(X_test_feat)
test_acc = accuracy_score(y_test, y_pred)
print(f"SVM测试集准确率: {test_acc:.4f}")

关键注意事项

  • 必须替换Input(shape=...)中的参数,确保和生成器输出的x_batch维度完全匹配
  • 若不需要预训练CNN,可直接训练特征提取层,但监督预训练通常能获得更有效的特征
  • SVM的超参数(如kernel、C、gamma)需根据数据集调优,可使用GridSearchCV进行网格搜索
  • 提取特征时,total_steps要和fit_generator中的steps_per_epoch/validation_steps一致,避免数据遗漏

内容的提问来源于stack exchange,提问作者Eda

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.15 03:05:19