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

如何将基于基因表达谱的疾病预测ANN模型准确率从69%提升至90%?

如何提升基因表达谱ANN模型的准确率(从69%到90%)

问题背景

基于基因表达谱构建的二分类ANN模型,用于预测个体是否患病,当前测试准确率卡在约69%,训练日志及代码如下:

训练输出日志

Epoch 1/100
45049/45049 [==============================] - 106s 2ms/step - loss: 0.6041 - accuracy: 0.6888 - val_loss: 0.6004 - val_accuracy: 0.6928
Epoch 2/100
45049/45049 [==============================] - 106s 2ms/step - loss: 0.6016 - accuracy: 0.6905 - val_loss: 0.5996 - val_accuracy: 0.6881
Epoch 3/100
45049/45049 [==============================] - 108s 2ms/step - loss: 0.6013 - accuracy: 0.6912 - val_loss: 0.5994 - val_accuracy: 0.6934
Epoch 4/100
45049/45049 [==============================] - 105s 2ms/step - loss: 0.6013 - accuracy: 0.6913 - val_loss: 0.5996 - val_accuracy: 0.6881
Epoch 5/100
45049/45049 [==============================] - 109s 2ms/step - loss: 0.6010 - accuracy: 0.6919 - val_loss: 0.5999 - val_accuracy: 0.6949
Epoch 6/100
45049/45049 [==============================] - 111s 2ms/step - loss: 0.6009 - accuracy: 0.6917 - val_loss: 0.5998 - val_accuracy: 0.6937
Epoch 7/100
45049/45049 [==============================] - 133s 3ms/step - loss: 0.6019 - accuracy: 0.6913 - val_loss: 0.6000 - val_accuracy: 0.6894
Epoch 8/100
45049/45049 [==============================] - 132s 3ms/step - loss: 0.6014 - accuracy: 0.6918 - val_loss: 0.5987 - val_accuracy: 0.6959
Epoch 9/100
45049/45049 [==============================] - 121s 3ms/step - loss: 0.6007 - accuracy: 0.6925 - val_loss: 0.5994 - val_accuracy: 0.6946
Epoch 10/100
45049/45049 [==============================] - 126s 3ms/step - loss: 0.6007 - accuracy: 0.6929 - val_loss: 0.6000 - val_accuracy: 0.6941
Epoch 11/100
45049/45049 [==============================] - 137s 3ms/step - loss: 0.6019 - accuracy: 0.6918 - val_loss: 0.5999 - val_accuracy: 0.6883
Epoch 12/100
45049/45049 [==============================] - 136s 3ms/step - loss: 0.6009 - accuracy: 0.6925 - val_loss: 0.5985 - val_accuracy: 0.6957
Epoch 13/100
45049/45049 [==============================] - 137s 3ms/step - loss: 0.6013 - accuracy: 0.6922 - val_loss: 0.5987 - val_accuracy: 0.6958
Epoch 14/100
45049/45049 [==============================] - 138s 3ms/step - loss: 0.6006 - accuracy: 0.6931 - val_loss: 0.5996 - val_accuracy: 0.6939
Epoch 15/100
45049/45049 [==============================] - 137s 3ms/step - loss: 0.6006 - accuracy: 0.6928 - val_loss: 0.6001 - val_accuracy: 0.6868
Epoch 16/100
45049/45049 [==============================] - 136s 3ms/step - loss: 0.6007 - accuracy: 0.6927 - val_loss: 0.5990 - val_accuracy: 0.6956
Epoch 17/100
45049/45049 [==============================] - 138s 3ms/step - loss: 0.6008 - accuracy: 0.6926 - val_loss: 0.6003 - val_accuracy: 0.6921
Epoch 18/100
45049/45049 [==============================] - 138s 3ms/step - loss: 0.6011 - accuracy: 0.6918 - val_loss: 0.5992 - val_accuracy: 0.6892
Epoch 19/100
45049/45049 [==============================] - 138s 3ms/step - loss: 0.6010 - accuracy: 0.6924 - val_loss: 0.6000 - val_accuracy: 0.6886
Epoch 20/100
45049/45049 [==============================] - 137s 3ms/step - loss: 0.6007 - accuracy: 0.6925 - val_loss: 0.6001 - val_accuracy: 0.6885
Epoch 21/100
45049/45049 [==============================] - 141s 3ms/step - loss: 0.6012 - accuracy: 0.6912 - val_loss: 0.5990 - val_accuracy: 0.6896
Epoch 22/100
45049/45049 [==============================] - 138s 3ms/step - loss: 0.6010 - accuracy: 0.6917 - val_loss: 0.5994 - val_accuracy: 0.6889
12514/12514 [==============================] - 21s 2ms/step - loss: 0.5988 - accuracy: 0.6957
ANN Test accuracy: 0.6957491040229797

模型代码

import pandas as pd
import numpy as np
from sklearn.model_selection import train_test_split
from sklearn.preprocessing import StandardScaler
from sklearn.impute import SimpleImputer
from tensorflow import keras
from tensorflow.keras import layers

# Load data from the dataset file
dataset_file = 'concatenated_dataset.csv'
df = pd.read_csv(dataset_file)

# Check if there are any missing values in the 'VALUE' column
if df['VALUE'].isnull().any():
    # Handling missing values with SimpleImputer
    imputer = SimpleImputer(strategy='mean')
    df['VALUE'] = imputer.fit_transform(df['VALUE'].values.reshape(-1, 1))

# Split the data into features (X) and target variable (y)
X = df['VALUE'].values.reshape(-1, 1)
y = df['Target'].values

# Step 2: Split the data into training and testing sets
X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2, random_state=42)

# Step 3: Feature Scaling (optional, but recommended for neural networks)
scaler = StandardScaler()
X_train_scaled = scaler.fit_transform(X_train)
X_test_scaled = scaler.transform(X_test)

# Step 4: Build the ANN model
input_dim = X_train_scaled.shape[1]

model = keras.Sequential([
    layers.Dense(units=256, activation='relu', input_shape=(input_dim,)),
    layers.Dropout(0.3),
    layers.Dense(units=128, activation='relu'),
    layers.Dropout(0.2),
    layers.Dense(units=64, activation='relu'),
    layers.Dropout(0.1),
    layers.Dense(units=1, activation='sigmoid')  # For binary classification
])

# Step 5: Compile the model
optimizer = keras.optimizers.Adam(learning_rate=0.001)
model.compile(optimizer=optimizer, loss='binary_crossentropy', metrics=['accuracy'])

# Step 6: Train the ANN model with early stopping
early_stopping = keras.callbacks.EarlyStopping(patience=10, restore_best_weights=True)
history = model.fit(X_train_scaled, y_train, epochs=100, batch_size=32,
                    validation_split=0.1, callbacks=[early_stopping])

# Step 7: Evaluate the ANN model on the test set
ann_loss, ann_accuracy = model.evaluate(X_test_scaled, y_test)
print("ANN Test accuracy:", ann_accuracy)

核心优化方案

一、数据层面(最关键的问题)

  • 重构特征矩阵:当前代码仅使用VALUE单个特征,完全浪费了基因表达谱的多维信息。需将数据集转换为样本-基因矩阵:以样本ID为行索引,基因为列名,表达量为值,可通过df.pivot(index='样本ID', columns='基因名', values='VALUE')实现,这是提升准确率的核心前提。
  • 优化数据预处理:
    • 缺失值处理:替换均值填充为KNN填充(sklearn.impute.KNNImputer),更适配基因表达这类高维生物数据的分布特性;
    • 特征选择:用方差过滤移除变异度极低的基因(sklearn.feature_selection.VarianceThreshold),或用互信息法筛选与目标变量相关性高的基因(sklearn.feature_selection.mutual_info_classif),减少噪声干扰;
    • 数据转换:对基因表达值做log2(VALUE + 1)转换,修正偏态分布,让数据更符合模型假设。
  • 检查类别平衡:统计y_train中患病/健康样本的比例,若差异超过2:1,属于不平衡数据。可通过SMOTE过采样少数类、欠采样多数类,或在model.fit()中设置class_weight='balanced'解决。

二、模型与训练优化

  • 调整模型结构:特征矩阵修正为高维后,先从简单模型开始尝试:比如2层全连接层(如128+64神经元),再根据验证集表现逐步调整层数和神经元数量,避免过度复杂的模型拟合噪声;
  • 优化训练策略:
    • 降低学习率至1e-4,或加入学习率调度器:
      lr_scheduler = keras.callbacks.ReduceLROnPlateau(factor=0.5, patience=5, min_lr=1e-6)
      
      并在fit()中加入该回调;
    • 替换优化器为RMSprop,部分场景下对生物数据的收敛效果更好;
    • 增加L2正则化:在Dense层中加入kernel_regularizer=keras.regularizers.l2(0.001),抑制过拟合。
  • 扩充训练数据:如果数据集规模较小,可对基因表达值加入轻微的高斯噪声做数据增强,或引入同类型公开数据集(如TCGA)做迁移学习。

三、评估指标补充

不要仅依赖准确率,尤其在类别不平衡场景下,需同时查看精确率、召回率、F1-score、ROC-AUC等指标,更全面判断模型性能,避免被表面的准确率误导。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.15 08:37:02