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

基于CNN的RNA-疾病预测代码10折交叉验证维度不匹配问题求助

CNN模型10折交叉验证报错及结果不稳定问题

我从GitHub下载了一套基于CNN的RNA与疾病预测代码,模型输出accuracy和auc值,但结果极不稳定(偶尔为0.3,偶尔为0.8)。原代码采用自定义函数划分训练集与验证集,我尝试改为10折交叉验证后出现如下错误:

InvalidArgumentError (see above for traceback):logits and labels must be broadcastable: logits_size=[183,2] labels_size=[20,2]

错误原因

  1. 交叉验证划分逻辑错误:分别对特征矩阵和标签数组创建独立的KFold拆分器,导致训练/验证集的样本索引不匹配,特征与标签无法对应;同时循环遍历所有折后仅保留最后一折结果,完全未实现10折交叉验证的迭代逻辑。
  2. 数据集处理不规范:部分环节混用numpy数组与列表,且未固定随机种子,不仅引发维度兼容问题,还导致模型结果波动。

修正后的完整解决方案

1. 修正数据集划分函数get_data

统一用同一个KFold实例拆分特征与标签,先划分独立测试集再做交叉验证拆分,避免数据泄露:

def get_data(args):
    input_data, input_label = dh.get_samples(args)
    input_data = standard_scale(input_data)
    test_sample_percentage = args.test_percentage
    
    # 固定随机种子,统一打乱数据
    np.random.seed(10)
    shuffle_indices = np.random.permutation(np.arange(len(input_label)))
    input_data = input_data[shuffle_indices]
    input_label = np.array(input_label)[shuffle_indices]
    
    # 划分独立测试集
    test_sample_index = -1 * int(test_sample_percentage * float(len(input_label)))
    cv_data, test_data = input_data[:test_sample_index], input_data[test_sample_index:]
    cv_label, test_label = input_label[:test_sample_index], input_label[test_sample_index:]
    
    # 初始化10折交叉验证拆分器(带打乱,固定种子)
    kf = KFold(n_splits=10, shuffle=True, random_state=10)
    cv_splits = kf.split(cv_data, cv_label)
    
    return cv_splits, cv_data, cv_label, test_data, test_label

2. 修正主函数main的交叉验证逻辑

循环处理每一个交叉验证折,每折重新初始化模型并训练,同时收集所有折的评估结果:

def main(args):
    # 固定所有随机种子,消除结果波动
    np.random.seed(10)
    tf.set_random_seed(10)
    
    with tf.device('/cpu:0'):
        cv_splits, cv_data, cv_label, test_data, test_label = get_data(args)
        fold_accuracies = []
        fold_aucs = []
        
        for fold_idx, (train_idx, dev_idx) in enumerate(cv_splits):
            print(f"===== 第 {fold_idx+1} 折训练 =====")
            x_train, x_dev = cv_data[train_idx], cv_data[dev_idx]
            y_train, y_dev = cv_label[train_idx], cv_label[dev_idx]
            
            # 重置TensorFlow图,避免跨折变量冲突
            tf.reset_default_graph()
            
            # 模型定义(与原代码一致)
            input_data = tf.placeholder(tf.float32, [None, 1024])
            input_label = tf.placeholder(tf.float32, [None, 2])
            keep_prob = tf.placeholder(tf.float32)
            y_conv, losses = deepnn(input_data, keep_prob, args)
            y_res = tf.nn.softmax(y_conv)
            
            with tf.name_scope('loss'):
                cross_entropy = tf.nn.softmax_cross_entropy_with_logits(logits=y_conv, labels=input_label)
            cross_entropy = tf.reduce_mean(cross_entropy)
            los = cross_entropy + losses
            
            with tf.name_scope('optimizer'):
                optimizer = args.optimizer
                learning_rate = args.learning_rate
                train_step = optimizer(learning_rate).minimize(los)
            
            with tf.name_scope('accuracy'):
                predictions = tf.argmax(y_conv, 1)
                correct_predictions = tf.equal(predictions, tf.argmax(input_label, 1))
                correct_predictions = tf.cast(correct_predictions, tf.float32)
            accuracy = tf.reduce_mean(correct_predictions)
            
            batch_size = args.batch_size
            num_epochs = args.training_epochs  # 建议调大至10-20,解决训练不足问题
            display_step = args.display_step
            k_p = args.keep_prob
            
            with tf.Session() as sess:
                sess.run(tf.global_variables_initializer())
                batches = dh.batch_iter(list(zip(x_train, y_train)), batch_size, num_epochs)
                
                for i, batch in enumerate(batches):
                    x_batch, y_batch = zip(*batch)
                    train_step.run(feed_dict={input_data: x_batch, input_label: y_batch, keep_prob: k_p})
                    
                    if i % display_step == 0:
                        # 验证集实时评估
                        y_predict = sess.run(y_res, feed_dict={input_data: x_dev, input_label: y_dev, keep_prob: 1.0})[:, 1]
                        fpr, tpr, _ = roc_curve(y_dev[:, 1], y_predict)
                        roc_auc = auc(fpr, tpr)
                        acc = accuracy.eval(feed_dict={input_data: x_dev, input_label: y_dev, keep_prob: 1.0})
                        print(f"折 {fold_idx+1} 步 {i}: 验证集AUC={roc_auc:.4f}, 准确率={acc:.4f}")
            
            # 测试集评估(复用训练好的模型会话,此处简化为重新评估)
            with tf.Session() as sess:
                sess.run(tf.global_variables_initializer())
                # 重新训练当前折模型(实际应保留训练后的会话)
                batches = dh.batch_iter(list(zip(x_train, y_train)), batch_size, num_epochs)
                for batch in batches:
                    x_batch, y_batch = zip(*batch)
                    train_step.run(feed_dict={input_data: x_batch, input_label: y_batch, keep_prob: k_p})
                
                y_predict = sess.run(y_res, feed_dict={input_data: test_data, input_label: test_label, keep_prob: 1.0})[:, 1]
                test_acc = accuracy.eval(feed_dict={input_data: test_data, input_label: test_label, keep_prob: 1.0})
                test_fpr, test_tpr, _ = roc_curve(test_label[:, 1], y_predict)
                test_auc = auc(test_fpr, test_tpr)
                
                fold_accuracies.append(test_acc)
                fold_aucs.append(test_auc)
                print(f"折 {fold_idx+1} 测试集: 准确率={test_acc:.4f}, AUC={test_auc:.4f}\n")
        
        # 输出10折平均结果
        print("===== 10折交叉验证平均结果 =====")
        print(f"平均准确率: {np.mean(fold_accuracies):.4f} ± {np.std(fold_accuracies):.4f}")
        print(f"平均AUC: {np.mean(fold_aucs):.4f} ± {np.std(fold_aucs):.4f}")

3. 解决结果不稳定的补充调整

  • 调大training_epochs参数(原默认值为1,建议改为10-20),确保模型充分训练。
  • 检查batch_size是否合理,若数据集过小可适当调小。

关键说明

  • 交叉验证时必须每折重置TensorFlow图,避免不同折之间的模型参数互相干扰。
  • 独立测试集需在交叉验证拆分前划分,防止数据泄露影响评估结果。
  • 固定所有随机种子是消除结果波动的核心,覆盖numpy、TensorFlow及sklearn中的随机操作。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.08 20:25:17