如何在Keras中使用Precision而非Accuracy优化CNN模型
二分类Keras模型使用Precision作为训练指标的修复方案
失效原因
直接将字符串'precision'传入model.compile的metrics参数不生效,是因为TensorFlow/Keras的字符串指标别名未明确匹配二分类场景的精确率计算逻辑,需要手动传入实例化的精确率指标对象。
修复后完整代码
from tensorflow.keras.models import Sequential from tensorflow.keras.layers import Dense from tensorflow.keras import optimizers from tensorflow.keras.metrics import Precision model = Sequential() model.add(Dense(64,name = 'Primera', input_dim=8, activation='relu')) model.add(Dense(32 ,name = 'Segunda')) model.add(Dense(1,name = 'Tercera', activation='sigmoid')) optm = optimizers.Adam(learning_rate=0.001, beta_1=0.9, beta_2=0.999, amsgrad=False) # 此处替换为实例化的二分类Precision指标 model.compile(loss='binary_crossentropy', optimizer=optm, metrics=[Precision(name='precision')]) model.summary() history = model.fit(trainX, trainY, epochs=1000, batch_size=16, validation_split=0.1, verbose=1)
补充说明
- 训练过程中会同时输出训练集精确率
precision和验证集精确率val_precision,训练完成后可从history.history字典中提取对应数值绘制训练曲线 - 若业务需要更高的精确率,可在实例化
Precision时指定分类阈值,例如Precision(name='precision', thresholds=0.7),阈值越高,模型输出为正样本的判定标准越严格,精确率通常越高 - 如需同时监控召回率、AUC等其他指标,导入
Recall、AUC等Keras内置指标类,按相同规则添加到metrics列表即可
内容的提问来源于stack exchange,提问作者Agustin Elli
相关产品推荐
相关产品推荐

