FlowResampler / utils /plot_utils.py
Android12138's picture
Upload folder using huggingface_hub (part 2)
01dee4e verified
Raw History Blame Contribute Delete
5.45 kB
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)