自定义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
相关产品推荐
相关产品推荐

