如何保存Keras实验运行产生的结果(而非模型)?
Hey there! Let's tackle your two questions step by step—saving the raw experiment data and the accuracy/loss plots you're generating.
1. Saving Raw Experiment Data (Accuracy, Loss Values)
When you train a Keras model, the history object returned by model.fit() stores all epoch-wise metrics (like training/validation accuracy, loss) as a dictionary (history.history). You can save this data for later analysis using either pickle or numpy:
Using Pickle
Pickle lets you save the entire metrics dictionary directly:
import pickle # After training your model with open('training_history.pkl', 'wb') as file: pickle.dump(cnn.history.history, file) # To load the data later with open('training_history.pkl', 'rb') as file: loaded_history = pickle.load(file) # Access metrics like loaded_history['acc'] or loaded_history['val_loss']
Using Numpy
If you prefer a lightweight format, numpy can save the dictionary as a .npy file:
import numpy as np # Save the metrics np.save('training_history.npy', cnn.history.history) # Load the data later loaded_history = np.load('training_history.npy', allow_pickle=True).item()
2. Saving Training/Validation Plots (Accuracy & Loss)
Your existing plotting code is already set up—you just need to add a line to save the figure before calling plt.show() (saving first ensures you capture the full plot without any blank areas). Here's how to modify your example code to save both accuracy and loss plots:
import matplotlib.pyplot as plt import numpy as np # Assume nb_epochs is defined, and cnn is your trained model plot = True if plot: # --- Training vs Validation Accuracy Plot --- plt.figure(0) plt.plot(cnn.history['acc'], 'r') plt.plot(cnn.history['val_acc'], 'g') plt.xticks(np.arange(0, nb_epochs+1, 2.0)) plt.rcParams['figure.figsize'] = (8, 6) plt.xlabel("Num of Epochs") plt.ylabel("Accuracy") plt.title("Training Accuracy vs Validation Accuracy") plt.legend(['Training Accuracy', 'Validation Accuracy']) # Save the plot (supports png, pdf, svg, etc.) plt.savefig('accuracy_plot.png', dpi=300, bbox_inches='tight') plt.close() # Free up memory by closing the figure # --- Training vs Validation Loss Plot (Optional) --- plt.figure(1) plt.plot(cnn.history['loss'], 'r') plt.plot(cnn.history['val_loss'], 'g') plt.xticks(np.arange(0, nb_epochs+1, 2.0)) plt.rcParams['figure.figsize'] = (8, 6) plt.xlabel("Num of Epochs") plt.ylabel("Loss") plt.title("Training Loss vs Validation Loss") plt.legend(['Training Loss', 'Validation Loss']) plt.savefig('loss_plot.png', dpi=300, bbox_inches='tight') plt.close()
dpi=300ensures high-resolution images (perfect for reports or papers)bbox_inches='tight'prevents plot elements like labels from getting cut offplt.close()avoids memory leaks when generating multiple plots
内容的提问来源于stack exchange,提问作者Charlie Parker

