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

TensorFlow自定义训练:模型编译、最优保存及Loss日志实现求助

Solution for Your Custom TensorFlow Training Loop

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.04 09:55:14