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

如何使用tf-slim将验证损失与准确率记录至TensorBoard?

Got it, I totally get where you're coming from—Keras's fit_generator is such a smooth, high-level way to handle training with automatic validation and TensorBoard logging without getting stuck in low-level tf.Session loops. Let's break down how to replicate that exact workflow using tf-slim, keeping things clean and avoiding manual session management:

Replicating Keras-style fit_generator with tf-slim

The core idea is to use tf-slim's built-in training utilities plus TensorFlow's SessionRunHook system to add epoch-end validation and logging. Here's a step-by-step implementation:

1. Prepare your data generators (reuse your Keras-style ones!)

First, wrap your existing train_gen() and valid_gen() into TensorFlow Dataset objects—this lets tf-slim seamlessly consume them without rewriting your generator logic:

import tensorflow as tf
from tensorflow.contrib import slim

def get_train_dataset(batch_size=32):
    return tf.data.Dataset.from_generator(
        train_gen,  # Your existing Keras-style training generator
        output_types=(tf.float32, tf.int32),  # Match your generator's output types
        output_shapes=((None, 224, 224, 3), (None,))  # Match your data's shape
    ).repeat().batch(batch_size)

def get_valid_dataset(batch_size=32):
    return tf.data.Dataset.from_generator(
        valid_gen,  # Your existing Keras-style validation generator
        output_types=(tf.float32, tf.int32),
        output_shapes=((None, 224, 224, 3), (None,))
    ).batch(batch_size)

2. Define your model with tf-slim

Use tf-slim's layer APIs to build your model (this replaces Keras's Model definition):

def my_model(inputs, num_classes=10, is_training=True):
    with slim.arg_scope([slim.conv2d, slim.fully_connected],
                        activation_fn=tf.nn.relu,
                        weights_regularizer=slim.l2_regularizer(1e-4)):
        net = slim.conv2d(inputs, 32, [3,3], scope='conv1')
        net = slim.max_pool2d(net, [2,2], scope='pool1')
        net = slim.conv2d(net, 64, [3,3], scope='conv2')
        net = slim.max_pool2d(net, [2,2], scope='pool2')
        net = slim.flatten(net, scope='flatten')
        net = slim.fully_connected(net, 128, scope='fc1')
        net = slim.dropout(net, 0.5, is_training=is_training, scope='dropout')
        logits = slim.fully_connected(net, num_classes, activation_fn=None, scope='fc2')
    return logits

3. Define loss and metrics

Create a helper function to compute training/validation loss and accuracy (matches Keras's built-in metrics):

def compute_loss_and_metrics(logits, labels):
    # Cross-entropy loss + regularization losses
    slim.losses.softmax_cross_entropy(logits, labels)
    total_loss = slim.losses.get_total_loss()
    
    # Accuracy metric
    predictions = tf.argmax(logits, axis=1)
    accuracy = tf.reduce_mean(tf.cast(tf.equal(predictions, labels), tf.float32))
    
    return total_loss, accuracy

4. Custom Hook for Epoch-End Validation

This is the key part: a custom SessionRunHook that runs validation after every epoch and logs results to TensorBoard. No manual session loops needed!

class ValidationHook(tf.train.SessionRunHook):
    def __init__(self, valid_dataset, num_valid_batches, log_dir, steps_per_epoch):
        self.valid_dataset = valid_dataset
        self.num_valid_batches = num_valid_batches  # Total batches in validation set
        self.log_dir = log_dir
        self.steps_per_epoch = steps_per_epoch  # Training steps per epoch
        self.summary_writer = None
        self.valid_loss = None
        self.valid_acc = None
        self.iterator = None

    def begin(self):
        # Set up validation data iterator
        self.iterator = self.valid_dataset.make_initializable_iterator()
        val_inputs, val_labels = self.iterator.get_next()
        
        # Build validation model (disable dropout/batch norm training mode)
        val_logits = my_model(val_inputs, is_training=False)
        self.valid_loss, self.valid_acc = compute_loss_and_metrics(val_logits, val_labels)
        
        # Create TensorBoard summaries for validation metrics
        loss_summary = tf.summary.scalar('validation_loss', self.valid_loss)
        acc_summary = tf.summary.scalar('validation_accuracy', self.valid_acc)
        self.summary_op = tf.summary.merge([loss_summary, acc_summary])
        
        # Initialize summary writer
        self.summary_writer = tf.summary.FileWriter(self.log_dir, tf.get_default_graph())

    def after_run(self, run_context, run_values):
        # Check if we've finished an epoch
        global_step = run_context.session.run(tf.train.get_global_step())
        if global_step % self.steps_per_epoch == 0 and global_step != 0:
            # Reset validation iterator and run full validation pass
            run_context.session.run(self.iterator.initializer)
            total_loss = 0.0
            total_acc = 0.0

            for _ in range(self.num_valid_batches):
                val_loss, val_acc, summary = run_context.session.run(
                    [self.valid_loss, self.valid_acc, self.summary_op]
                )
                total_loss += val_loss
                total_acc += val_acc

            # Calculate averages and log
            avg_loss = total_loss / self.num_valid_batches
            avg_acc = total_acc / self.num_valid_batches
            self.summary_writer.add_summary(summary, global_step)
            print(f"Epoch {global_step // self.steps_per_epoch}: "
                  f"Val Loss = {avg_loss:.4f}, Val Acc = {avg_acc:.4f}")

    def end(self, session):
        self.summary_writer.close()

5. Put it all together with tf-slim's training loop

Use slim.learning.train() to handle the entire training workflow—this replaces Keras's fit_generator and eliminates manual session code:

def main():
    # Configuration
    log_dir = "./tf_slim_logs"
    num_epochs = 10
    batch_size = 32
    steps_per_epoch = 1000  # Number of training steps per epoch
    num_valid_batches = 100  # Number of batches in your validation set
    num_classes = 10

    # Set up training data
    train_dataset = get_train_dataset(batch_size)
    train_iterator = train_dataset.make_initializable_iterator()
    train_inputs, train_labels = train_iterator.get_next()

    # Build training model and compute loss/metrics
    train_logits = my_model(train_inputs, is_training=True)
    total_loss, train_acc = compute_loss_and_metrics(train_logits, train_labels)

    # Training TensorBoard summaries
    train_loss_summary = tf.summary.scalar('training_loss', total_loss)
    train_acc_summary = tf.summary.scalar('training_accuracy', train_acc)
    train_summary_op = tf.summary.merge([train_loss_summary, train_acc_summary])

    # Optimizer
    optimizer = tf.train.AdamOptimizer(learning_rate=1e-3)
    train_op = slim.learning.create_train_op(total_loss, optimizer)

    # Initialize global step counter
    global_step = tf.train.get_or_create_global_step()

    # Set up hooks
    valid_dataset = get_valid_dataset(batch_size)
    validation_hook = ValidationHook(
        valid_dataset, num_valid_batches, log_dir, steps_per_epoch
    )
    train_summary_hook = tf.train.SummarySaverHook(
        save_steps=100,  # Log training metrics every 100 steps
        output_dir=log_dir,
        summary_op=train_summary_op
    )

    # Start training (no manual session loops!)
    slim.learning.train(
        train_op,
        log_dir=log_dir,
        global_step=global_step,
        number_of_steps=num_epochs * steps_per_epoch,
        hooks=[validation_hook, train_summary_hook],
        local_init_op=tf.group(train_iterator.initializer, tf.local_variables_initializer())
    )

if __name__ == "__main__":
    main()

Key Notes:

  • No manual tf.Session code: slim.learning.train() handles session creation, initialization, and training loop logic for you.
  • Reuse your generators: The tf.data.Dataset.from_generator wrapper lets you keep your existing Keras-style data generators.
  • Automatic logging: Both training and validation metrics are written to TensorBoard, just like Keras's fit_generator.
  • Epoch-end validation: The custom ValidationHook triggers validation after every full epoch, matching Keras's behavior.

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.25 06:50:57