Download utils/plot_utils.py from Android12138/FlowResampler: direct link, hf CLI and curl.
- Browser
- Download file 5.45 kB
-
https://huggingface.co/Android12138/FlowResampler/resolve/main/utils/plot_utils.py
- Command line
-
hf download hf://Android12138/FlowResampler/utils/plot_utils.py
-
curl -L -o plot_utils.py https://huggingface.co/Android12138/FlowResampler/resolve/main/utils/plot_utils.py
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) | |