如何使用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:
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.Sessioncode:slim.learning.train()handles session creation, initialization, and training loop logic for you. - Reuse your generators: The
tf.data.Dataset.from_generatorwrapper 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
ValidationHooktriggers validation after every full epoch, matching Keras's behavior.
内容的提问来源于stack exchange,提问作者Abraham Ben

