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

自定义TensorFlow训练循环:早停、模型保存与CSV日志实现

问题解决方案

1. 自定义训练循环中的模型编译

在自定义训练循环中,不需要强制调用model.compile()。因为你已经手动实现了损失计算、梯度求解和权重更新的完整流程,compile()方法主要是为Keras内置的model.fit() API服务的,用于绑定优化器、损失函数和指标。

如果你希望模型保存/加载时能保留优化器、损失等配置信息,或者后续可能切换到model.fit(),可以补充调用:

model.compile(optimizer=optimizer, loss=loss_fn, metrics=['mse'])

但这对自定义训练循环的执行逻辑没有影响,训练过程依然依赖你手动定义的train_step和test_step。

2. 实现早停机制

需要手动跟踪验证集损失的变化,当连续指定epoch数损失未下降时停止训练并保存最优模型:

  • 初始化最优损失、等待计数器和耐心值
  • 每个epoch结束后对比当前验证损失与最优损失
  • 若损失下降则更新最优损失、保存模型并重置计数器;否则计数器加1
  • 当计数器达到耐心值时终止训练

3. 保存训练/验证损失到CSV文件

使用Python内置csv模块实现逐行写入,每个epoch记录训练损失和验证损失,实现类似Keras CSVLogger的功能。


修改后的完整代码

import tensorflow as tf
import numpy as np
import time
import csv

@tf.function
def train_step(x, y):
    with tf.GradientTape() as tape:
        logits = model(x, training=True)
        loss_value = loss_fn(y, logits)
    grads = tape.gradient(loss_value, model.trainable_weights)
    optimizer.apply_gradients(zip(grads, model.trainable_weights))
    train_loss_metric.update_state(y, logits)
    return loss_value

@tf.function
def test_step(x, y):
    val_logits = model(x, training=False)
    val_loss_metric.update_state(y, val_logits)

# 初始化优化器、损失函数
optimizer = tf.keras.optimizers.SGD(learning_rate=1e-3)
loss_fn = tf.keras.losses.MeanSquaredError()
batch_size = 16

# 加载数据集
x_train = np.load('x_train_data.npy') 
x_valid = np.load('x_valid_data.npy') 
y_train = np.load('y_train_data.npy') 
y_valid = np.load('y_valid_data.npy') 

# 数据预处理
x_train = np.expand_dims(x_train, axis=2)
x_valid = np.expand_dims(x_valid, axis=2)
y_train = np.expand_dims(y_train, axis=2)
y_valid = np.expand_dims(y_valid, axis=2)

# 构建TF数据集
train_dataset = tf.data.Dataset.from_tensor_slices((x_train, y_train))
train_dataset = train_dataset.shuffle(buffer_size=1024).batch(batch_size)

val_dataset = tf.data.Dataset.from_tensor_slices((x_valid, y_valid))
val_dataset = val_dataset.batch(batch_size)

# 初始化指标(统一命名为loss避免混淆)
train_loss_metric = tf.keras.metrics.MeanSquaredError()
val_loss_metric = tf.keras.metrics.MeanSquaredError()

# 加载模型
model = test_model(im_width=1, im_height=80, neurons=16, kern_sz=20) 
model.summary()

# -------------------------- 新增功能初始化 --------------------------
# 早停参数
patience = 10
best_val_loss = float('inf')
wait = 0
save_path = 'model.h5'

# CSV日志初始化
csv_path = 'training_log.csv'
# 写入表头
with open(csv_path, 'w', newline='') as f:
    writer = csv.writer(f)
    writer.writerow(['epoch', 'train_loss', 'val_loss'])

# -------------------------- 自定义训练循环 --------------------------
epochs = 100  # 设足够大的epoch数,由早停控制终止
losses = []  # 记录批次损失

for epoch in range(epochs):
    print(f"\nStart of epoch {epoch}")
    start_time = time.time()

    # 训练批次循环
    for step, (x_batch_train, y_batch_train) in enumerate(train_dataset):
        loss_value = train_step(x_batch_train, y_batch_train)
        losses.append(float(loss_value))

        # 每2步打印批次损失
        if step % 2 == 0:
            print(f"Training loss (for one batch) at step {step}: {float(loss_value):.4f}")
            print(f"Seen so far: {(step + 1) * batch_size} samples")

    # 获取 epoch 级训练损失
    train_loss = train_loss_metric.result()
    print(f"Training loss over epoch: {float(train_loss):.4f}")
    train_loss_metric.reset_states()

    # 验证循环
    for x_batch_val, y_batch_val in val_dataset:
        test_step(x_batch_val, y_batch_val)
    
    val_loss = val_loss_metric.result()
    val_loss_metric.reset_states()
    print(f"Validation loss: {float(val_loss):.4f}")
    print(f"Time taken: {time.time() - start_time:.2f}s")

    # -------------------------- 早停逻辑 --------------------------
    if val_loss < best_val_loss:
        best_val_loss = val_loss
        model.save(save_path)
        print(f"New best model saved to {save_path}")
        wait = 0  # 重置等待计数器
    else:
        wait += 1
        if wait >= patience:
            print(f"Validation loss hasn't improved in {patience} epochs. Stopping training.")
            break

    # -------------------------- 写入CSV日志 --------------------------
    with open(csv_path, 'a', newline='') as f:
        writer = csv.writer(f)
        writer.writerow([epoch, float(train_loss), float(val_loss)])

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.04 05:20:48