TensorFlow自定义训练:模型编译、最优保存及Loss日志实现求助
Got it, let's integrate the three features you need into your custom training loop step by step. Here's the modified code with detailed explanations:
1. Model Compilation
While custom training loops don't strictly require model.compile() like Keras' built-in fit, we can still use it to bind your optimizer, loss function, and metrics to the model for consistency. This makes your code more readable and aligns with Keras conventions.
2. Early Stopping & Best Model Saving
We'll track the validation loss across epochs, count how many consecutive epochs it doesn't improve, and stop training once it hits your patience threshold (10 epochs). We'll also save the best model whenever validation loss improves.
3. CSV Logging
We'll use Python's built-in csv module to write each epoch's training and validation loss to a CSV file, just like Keras' CSVLogger.
Modified Full Code
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_acc_metric.update_state(y, logits) return loss_value @tf.function def test_step(x, y): val_logits = model(x, training=False) val_acc_metric.update_state(y, val_logits) # Calculate validation loss for early stopping tracking val_loss = loss_fn(y, val_logits) return val_loss # Initialize optimizer, loss, metrics optimizer = tf.keras.optimizers.SGD(learning_rate=1e-3) loss_fn = tf.keras.losses.MeanSquaredError() train_acc_metric = tf.keras.metrics.MeanSquaredError() val_acc_metric = tf.keras.metrics.MeanSquaredError() batch_size = 16 # Load dataset 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') # Preprocess data 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) # Prepare datasets 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) # Define and compile model (Step 1: Model Compilation) model = test_model(im_width=1, im_height=80, neurons=16, kern_sz=20) model.compile( optimizer=optimizer, loss=loss_fn, metrics=[train_acc_metric] ) model.summary() # Initialize variables for early stopping and logging best_val_loss = float('inf') patience = 10 wait = 0 save_path = 'model.h5' csv_log_path = 'training_log.csv' losses = [] # Initialize your batch loss tracking list # Write CSV header with open(csv_log_path, 'w', newline='') as f: writer = csv.writer(f) writer.writerow(['epoch', 'train_loss', 'val_loss']) ###### Custom training loop ###### epochs = 100 # Increase epoch count since we'll use early stopping for epoch in range(epochs): print("\nStart of epoch %d" % (epoch,)) start_time = time.time() epoch_train_losses = [] # Training loop for step, (x_batch_train, y_batch_train) in enumerate(train_dataset): loss_value = train_step(x_batch_train, y_batch_train) epoch_train_losses.append(float(loss_value)) losses.append(float(loss_value)) # Log every 2 batches if step % 2 == 0: print( "Training loss (for one batch) at step %d: %.4f" % (step, float(loss_value)) ) print("Seen so far: %d samples" % ((step + 1) * batch_size)) # Calculate and display epoch-level training metrics avg_train_loss = np.mean(epoch_train_losses) train_acc = train_acc_metric.result() print("Training loss over epoch: %.4f" % (float(train_acc),)) train_acc_metric.reset_states() # Validation loop epoch_val_losses = [] for x_batch_val, y_batch_val in val_dataset: val_loss = test_step(x_batch_val, y_batch_val) epoch_val_losses.append(float(val_loss)) avg_val_loss = np.mean(epoch_val_losses) val_acc = val_acc_metric.result() val_acc_metric.reset_states() print("Validation loss: %.4f" % (float(val_acc),)) print("Time taken: %.2fs" % (time.time() - start_time)) # Step 3: Log metrics to CSV with open(csv_log_path, 'a', newline='') as f: writer = csv.writer(f) writer.writerow([epoch, float(train_acc), float(val_acc)]) # Step 2: Early stopping logic & save best model if float(val_acc) < best_val_loss: best_val_loss = float(val_acc) wait = 0 model.save(save_path) print(f"✅ Best model saved to {save_path} (val_loss improved to {best_val_loss:.4f})") else: wait += 1 print(f"⚠️ Val loss not improved, wait count: {wait}/{patience}") if wait >= patience: print("🛑 Early stopping triggered - no improvement in 10 epochs") break
Key Changes Explained:
- Model Compilation: Added
model.compile()to bind optimizer, loss, and metrics to the model for Keras-like consistency (optional but makes code more maintainable). - Early Stopping: Added tracking variables to monitor validation loss. If validation loss improves, we save the model and reset the wait counter; if not, we increment the counter until it hits the 10-epoch patience limit, then stop training.
- CSV Logging: Created a CSV file with headers, then appended each epoch's training/validation loss after every epoch completes, matching the behavior of Keras'
CSVLogger. - Validation Loss Tracking: Modified
test_step()to return batch-level validation loss, so we can accurately track epoch-level validation performance for early stopping.
内容的提问来源于stack exchange,提问作者pro

