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

基于TensorFlow训练ANN:如何通过验证避免模型过拟合?

解决过拟合:基于验证集的训练停止策略实现

核心思路

要通过验证结果避免过拟合,关键是定期在验证集上评估模型性能,当验证集损失不再下降甚至开始上升(与训练集损失曲线交叉)时,立即停止训练,同时保存性能最优的模型。

具体修改步骤

1. 拆分数据集为训练集和验证集

首先将输入的dataset按比例拆分为训练集和验证集,示例采用8:2的划分比例:

from sklearn.model_selection import train_test_split

train_dataset, val_dataset = train_test_split(dataset, test_size=0.2, random_state=42)
num_train_samples = len(train_dataset)
num_val_samples = len(val_dataset)

2. 记录训练集与验证集的损失变化

在训练循环外初始化两个列表,用于存储每轮的平均训练损失和验证损失:

train_loss_history = []
val_loss_history = []

3. 每轮训练后执行验证评估

修改训练循环,完成一轮训练后,遍历验证集计算平均损失(验证阶段不更新模型参数):

for i in range(self.epoch):
    # 训练阶段:遍历训练集计算平均损失
    train_loss = 0.0
    for j in range(num_train_samples):
        loss, _ = sess.run([self.loss, self.train_op], feed_dict={self.x:[train_dataset[j]]})
        train_loss += loss
    avg_train_loss = train_loss / num_train_samples
    train_loss_history.append(avg_train_loss)

    # 验证阶段:遍历验证集计算平均损失
    val_loss = 0.0
    for j in range(num_val_samples):
        loss = sess.run(self.loss, feed_dict={self.x:[val_dataset[j]]})
        val_loss += loss
    avg_val_loss = val_loss / num_val_samples
    val_loss_history.append(avg_val_loss)

    # 按间隔打印损失并保存阶段性模型
    if i % 10 == 0:
        ram_train.append(cpu_usage(1))
        print(f'epoch {i}: 训练损失 = {avg_train_loss:.4f}, 验证损失 = {avg_val_loss:.4f}')
        self.saver.save(sess, f'./model_hidden{self.hidden}_wdw{self.window}_epoch{i}.ckpt')

4. 添加早停逻辑

在每轮验证后,检查验证损失的变化,当验证损失超过训练损失且连续多轮上升时,触发早停:

# 早停参数:允许验证损失上升的最大轮数
patience = 5
best_val_loss = float('inf')
patience_counter = 0

for i in range(self.epoch):
    # 训练和验证步骤同上...

    # 早停判断逻辑
    if avg_val_loss < best_val_loss:
        best_val_loss = avg_val_loss
        patience_counter = 0
        # 保存当前最优模型
        self.saver.save(sess, f'./best_model_hidden{self.hidden}_wdw{self.window}.ckpt')
    else:
        patience_counter += 1
        # 验证损失超过训练损失且耐心耗尽时停止训练
        if avg_val_loss > avg_train_loss and patience_counter >= patience:
            print(f'epoch {i}: 验证损失超过训练损失,触发早停')
            break

5. 优化模型保存逻辑

仅保存最优模型和指定轮次的模型,避免每轮保存浪费存储资源:

# 移除原有的每轮保存代码,仅在指定轮次和最优状态时保存
if i % 10 == 0:
    ram_train.append(cpu_usage(1))
    print(f'epoch {i}: 训练损失 = {avg_train_loss:.4f}, 验证损失 = {avg_val_loss:.4f}')
    self.saver.save(sess, f'./model_hidden{self.hidden}_wdw{self.window}_epoch{i}.ckpt')

完整修改后的train方法

def train(self, dataset):
    # 拆分训练集与验证集
    from sklearn.model_selection import train_test_split
    train_dataset, val_dataset = train_test_split(dataset, test_size=0.2, random_state=42)
    num_train_samples = len(train_dataset)
    num_val_samples = len(val_dataset)
    
    print('Training...')
    tic = time.time()
    # 初始化损失历史记录
    train_loss_history = []
    val_loss_history = []
    # 早停参数配置
    patience = 5
    best_val_loss = float('inf')
    patience_counter = 0

    with tf.compat.v1.Session() as sess:
        sess.run(tf.compat.v1.global_variables_initializer())
        for i in range(self.epoch):
            # 训练阶段计算平均损失
            train_loss = 0.0
            for j in range(num_train_samples):
                loss, _ = sess.run([self.loss, self.train_op], feed_dict={self.x:[train_dataset[j]]})
                train_loss += loss
            avg_train_loss = train_loss / num_train_samples
            train_loss_history.append(avg_train_loss)

            # 验证阶段计算平均损失
            val_loss = 0.0
            for j in range(num_val_samples):
                loss = sess.run(self.loss, feed_dict={self.x:[val_dataset[j]]})
                val_loss += loss
            avg_val_loss = val_loss / num_val_samples
            val_loss_history.append(avg_val_loss)

            # 间隔打印与阶段性模型保存
            if i % 10 == 0:
                ram_train.append(cpu_usage(1))
                print(f'epoch {i}: 训练损失 = {avg_train_loss:.4f}, 验证损失 = {avg_val_loss:.4f}')
                self.saver.save(sess, f'./model_hidden{self.hidden}_wdw{self.window}_epoch{i}.ckpt')

            # 早停触发逻辑
            if avg_val_loss < best_val_loss:
                best_val_loss = avg_val_loss
                patience_counter = 0
                self.saver.save(sess, f'./best_model_hidden{self.hidden}_wdw{self.window}.ckpt')
            else:
                patience_counter += 1
                if avg_val_loss > avg_train_loss and patience_counter >= patience:
                    print(f'epoch {i}: 验证损失超过训练损失,触发早停')
                    break

        # 最终保存最优模型
        self.saver.save(sess, f'./best_model_hidden{self.hidden}_wdw{self.window}.ckpt')
    
    tac = time.time()
    print('Done.')
    return avg_train_loss, avg_val_loss, ram_train, (tac - tic), train_loss_history, val_loss_history

额外建议

  • 除早停策略外,可给模型添加Dropout层或L2正则化,进一步抑制过拟合
  • 替换单样本训练为批量训练(batch training),提升训练效率和模型稳定性

内容的提问来源于stack exchange,提问作者Mariana Flávio

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.14 10:35:22