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

训练中保存验证误差最优的Estimator模型并恢复训练

Solution for Saving Best Estimator Models on Validation Loss Drop

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:

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_fn in the exporter or the condition in the manual loop to current_val_acc > best_val_acc.
  • Storage Management: Using shutil.rmtree ensures 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_fn to specify how the model accepts input for inference.

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.26 10:29:27