import os import matplotlib.pyplot as plt from tensorboard.backend.event_processing import event_accumulator def parse_events_and_save_plot( event_file_path: str, output_image_path: str, ): """Parse TensorBoard event file and save scalar plots to an image file. Args: event_file_path (str): Path to the TensorBoard event file. output_image_path (str): Path to save the output image file. """ # 1. Initialize the EventAccumulator. # size_guidance is a dictionary that tells the accumulator how much data to load. # Setting SCALARS to 0 loads all scalar data. ea = event_accumulator.EventAccumulator( event_file_path, size_guidance={event_accumulator.SCALARS: 0} ) # 2. Load the events from the file. ea.Reload() # 3. Get all available scalar tags (e.g., 'Loss/train', 'Accuracy/validation'). scalar_tags = ea.Tags()['scalars'] # 4. Create subplots. # The number of rows will match the number of tags to plot each on its own subplot. num_tags = len(scalar_tags) fig, axes = plt.subplots(num_tags, 1, figsize=(12, 6 * num_tags), squeeze=False) # Using squeeze=False ensures `axes` is always a 2D array, which simplifies indexing. axes = axes.flatten() # Flatten the 2D array to 1D for easy iteration. # 5. Iterate through each tag, extract its data, and plot it. for i, tag in enumerate(scalar_tags): # Extract scalar events for the current tag. scalar_events = ea.Scalars(tag) # Extract the step and value for each event point. steps = [event.step for event in scalar_events] values = [event.value for event in scalar_events] # Plot the data on the corresponding subplot. ax = axes[i] ax.plot(steps, values, label=tag, color=f'C{i}') ax.set_title(tag, fontsize=16) ax.set_xlabel("Step") ax.set_ylabel("Value") ax.grid(True, linestyle="--", alpha=0.6) ax.legend() # 6. Adjust layout and save the figure. fig.tight_layout(pad=3.0) # Adjust subplot params for a tight layout. # Save the figure to the specified path. # `bbox_inches='tight'` helps prevent labels from being cut off. plt.savefig(output_image_path, dpi=150, bbox_inches="tight") # Close the figure to free up memory. plt.close(fig) # Font definitions FONT1 = {"family": "serif", "color": "Black", "weight": "bold", "size": 16} # global figure FONT2 = {"family": "serif", "color": "Black", "weight": "bold", "size": 13} # plot-level title FONT3 = {"family": "serif", "color": "Black", "weight": "bold", "size": 8} # axis labels, ticks, legend def set_ax_style( ax, *, # force keyword arguments title: str | None = None, xlabel: str | None = None, ylabel: str | None = None, show_legend: bool = False, legend_kwargs: dict | None = None, grid_axis: str | None = "y", ): """Apply unified axis, tick, label, legend, and grid styles. Args: ax: Matplotlib axis object to style. title (str | None): Title of the plot. xlabel (str | None): Label for the x-axis. ylabel (str | None): Label for the y-axis. show_legend (bool): Whether to display the legend. legend_kwargs (dict | None): Additional keyword arguments for the legend. grid_axis (str | None): Axis for grid lines ('x', 'y', or None). """ # 1) Tick and border settings ax.tick_params(which="major", axis="x", length=2, width=0.6) ax.tick_params(which="major", axis="y", length=2, width=0.6) ax.tick_params(which="minor", axis="x", length=1, width=0.6) ax.tick_params(which="minor", axis="y", length=1, width=0.6) for side in ["bottom", "left", "right", "top"]: ax.spines[side].set_linewidth(1) # 2) Tick label font labels = ax.get_xticklabels() + ax.get_yticklabels() for label in labels: label.set_fontname(FONT3["family"]) label.set_color(FONT3["color"]) label.set_fontweight(FONT3["weight"]) label.set_fontsize(FONT3["size"]) # 3) Axis labels and title if xlabel is not None: ax.set_xlabel(xlabel, fontdict=FONT3) if ylabel is not None: ax.set_ylabel(ylabel, fontdict=FONT3) if title is not None: ax.set_title(title, fontdict=FONT2) # 4) Legend if show_legend: lg_kwargs = {"prop": {"family": FONT3["family"], "size": FONT3["size"], "weight": FONT3["weight"]}} if legend_kwargs: lg_kwargs.update(legend_kwargs) leg = ax.legend(**lg_kwargs) if leg is not None: leg.get_frame().set_linewidth(0.8) # 5) Grid if grid_axis is not None: ax.grid(axis=grid_axis, linestyle="--", alpha=0.5) def save_fig( fig, output_dir: str, sub_dir: str, file_name: str ): """Apply tight layout, save figure, and close it. Args: fig: Matplotlib figure object to save. output_dir (str): Base output directory. sub_dir (str): sub_directory within the output directory. file_name (str): file_name (without extension) for the saved figure. """ os.makedirs(os.path.join(output_dir, sub_dir), exist_ok=True) fig.tight_layout() fig.savefig(f"{output_dir}/{sub_dir}/{file_name}.pdf", dpi=600, bbox_inches="tight", pad_inches=0.02,) plt.close(fig)