如何批量保存H2O模型explain()方法生成的所有输出图像?
explain() Method in a Script I get it—saving individual plots like SHAP summaries is straightforward with save_plot_path, but explain() dumps a whole suite of visualizations without a built-in save option, especially when you're running a script instead of a Jupyter Notebook. Here's a reliable way to capture all those plots in one go:
Step-by-Step Solution
The key is to leverage the H2OExplainability object returned by explain(), which contains references to every generated plot. We'll use Matplotlib to save each plot to a directory, with non-interactive backend setup to avoid GUI popups.
Full Working Code
import h2o from h2o.automl import H2OAutoML import matplotlib # Set non-interactive backend for script execution (no GUI windows) matplotlib.use('Agg') import matplotlib.pyplot as plt import os # Initialize H2O and load data h2o.init() df = h2o.import_file("https://h2o-public-test-data.s3.amazonaws.com/smalldata/wine/winequality-redwhite-no-BOM.csv") response = "quality" predictors = [ "fixed acidity", "volatile acidity", "citric acid", "residual sugar", "chlorides", "free sulfur dioxide", "total sulfur dioxide", "density", "pH", "sulphates", "alcohol", "type" ] train, test = df.split_frame(seed=1) # Train AutoML model aml = H2OAutoML(max_runtime_secs=120, seed=1) aml.train(x=predictors, y=response, training_frame=train) leader_model = aml.leader # Generate explanations and capture the explainability object explanation = leader_model.explain(test) # Create a directory to save plots (if it doesn't exist) save_dir = "h2o_explain_plots" os.makedirs(save_dir, exist_ok=True) # Iterate through all explanation modules and save each plot for idx, exp_module in enumerate(explanation._explanations): # Clean up the plot title to use as a filename plot_title = exp_module.get("title", f"plot_{idx}").replace(" ", "_").replace("/", "_") plot_path = f"{save_dir}/{plot_title}.png" # Get the Matplotlib figure object and save it fig = exp_module.get("figure") if fig is not None: plt.figure(fig.number) # Switch to the target figure plt.savefig(plot_path, bbox_inches='tight', dpi=300) # Save with high resolution plt.close(fig) # Free up memory by closing the figure print(f"Successfully saved: {plot_path}")
Key Details Explained
- Non-Interactive Backend:
matplotlib.use('Agg')is critical for script mode—it disables GUI windows, so your script can run headless without user input. - Explainability Object:
explanation._explanationsis a list where each item represents one explanation module (e.g., SHAP summary, partial dependence plots, variable importance). Each module includes the plot's title and the MatplotlibFigureobject. - Filename Sanitization: We replace spaces and slashes in plot titles to avoid filesystem errors when saving.
- Memory Management: Closing each figure with
plt.close(fig)prevents memory leaks, especially if you're working with large datasets or multiple models.
This method will save all plots generated by explain() into a dedicated directory, so you don't have to manually capture or recreate each one.
内容的提问来源于stack exchange,提问作者Shubham Goel

