训练中保存验证误差最优的Estimator模型并恢复训练
Since Hooks aren't feasible here (they execute after every training step and can't access the full validation set's average loss), you have two reliable approaches to save only your top-performing models when validation error decreases:
1. Implement a Custom Exporter (Recommended for train_and_evaluate Workflows)
Estimator's Exporter system is built exactly for this scenario—it runs after each evaluation and can access full validation metrics. You can create a custom exporter that tracks the best validation loss, saves the model only when it improves, and keeps a fixed number of top models.
Here's a concrete implementation:
import os import shutil from tensorflow_estimator.python.estimator.exporter import Exporter class BestModelExporter(Exporter): def __init__(self, name='best_model', eval_metric_name='loss', compare_fn=lambda current, best: current < best, keep_top_n=3): self._name = name self._eval_metric_name = eval_metric_name self._compare_fn = compare_fn # Defines "better" (lower loss = better here) self._keep_top_n = keep_top_n self._best_metric_value = None self._saved_model_paths = [] def export(self, estimator, export_path, checkpoint_path, eval_result, is_the_final_export): current_metric = eval_result[self._eval_metric_name] # Check if current model outperforms the best so far if self._best_metric_value is None or self._compare_fn(current_metric, self._best_metric_value): self._best_metric_value = current_metric # Export to a unique directory with metric value for clarity export_dir = os.path.join(export_path, f"{self._name}_loss_{current_metric:.4f}") estimator.export_saved_model(export_dir, serving_input_receiver_fn=your_serving_input_fn) self._saved_model_paths.append((current_metric, export_dir)) # Sort models by performance and keep only top N self._saved_model_paths.sort(key=lambda x: x[0], reverse=False) # Ascending for lower loss if len(self._saved_model_paths) > self._keep_top_n: _, worst_path = self._saved_model_paths.pop() shutil.rmtree(worst_path) return self._saved_model_paths[-1][1] if self._saved_model_paths else None
To use this, pass it to tf.estimator.train_and_evaluate:
exporters = [BestModelExporter(keep_top_n=5)] tf.estimator.train_and_evaluate(estimator, train_spec, eval_spec)
This exporter runs automatically after each evaluation, checks for validation loss improvement, saves the model, and cleans up older, worse-performing models to keep only your top N.
2. Manual Evaluation & Saving (For Custom Training Loops)
If you're using a custom training loop instead of train_and_evaluate, you can manually run estimator.evaluate() at intervals, track the best validation loss, and save the model when it improves.
Example code:
import os import shutil best_val_loss = float('inf') top_model_paths = [] keep_top_n = 3 # Custom training loop for epoch in range(num_epochs): estimator.train(input_fn=train_input_fn, steps=steps_per_epoch) # Run evaluation every N epochs if epoch % eval_interval == 0: eval_results = estimator.evaluate(input_fn=val_input_fn) current_val_loss = eval_results['loss'] if current_val_loss < best_val_loss: best_val_loss = current_val_loss # Save the model with epoch and loss info save_dir = os.path.join(model_dir, f"best_model_epoch_{epoch}_loss_{current_val_loss:.4f}") estimator.export_saved_model(save_dir, serving_input_receiver_fn=your_serving_input_fn) top_model_paths.append((current_val_loss, save_dir)) # Trim to keep only top N models top_model_paths.sort(key=lambda x: x[0]) if len(top_model_paths) > keep_top_n: _, worst_path = top_model_paths.pop() shutil.rmtree(worst_path)
Key Notes:
- Metric Adjustment: If you're optimizing for a metric where higher is better (like accuracy), tweak the
compare_fnin the exporter or the condition in the manual loop tocurrent_val_acc > best_val_acc. - Storage Management: Using
shutil.rmtreeensures you don't waste space on subpar models—only keep the ones that perform best on validation data. - Serving Input Function: Don't forget to define
your_serving_input_fnto specify how the model accepts input for inference.
内容的提问来源于stack exchange,提问作者tfnopcho

