Download src/plotter.py from EmbodiedCity/EmbodiedCity-Leaderboard: direct link, hf CLI and curl.
- Browser
- Download file 3.82 kB
-
https://huggingface.co/spaces/EmbodiedCity/EmbodiedCity-Leaderboard/resolve/main/src/plotter.py
- Command line
-
hf download hf://spaces/EmbodiedCity/EmbodiedCity-Leaderboard/src/plotter.py
-
curl -L -o plotter.py https://huggingface.co/spaces/EmbodiedCity/EmbodiedCity-Leaderboard/resolve/main/src/plotter.py
3.82 kB
| from __future__ import annotations | |
| from typing import Optional | |
| import matplotlib.pyplot as plt | |
| import numpy as np | |
| import pandas as pd | |
| class Plotter: | |
| def __init__(self, data_loader): | |
| self.data_loader = data_loader | |
| def create_comparison_plot( | |
| self, | |
| model_filter: str, | |
| entry_type_filter: str, | |
| selected_plot_metric: str, | |
| plot_sort_mode: str, | |
| display_metric_name: Optional[str] = None, | |
| ) -> plt.Figure: | |
| metric_display_name = display_metric_name or selected_plot_metric | |
| df = self.data_loader.df_all | |
| if df is None or df.empty: | |
| return self._empty_plot("No data available for plotting") | |
| if model_filter and model_filter.strip(): | |
| df = df[df["Model"].str.contains(model_filter, case=False, na=False)] | |
| if entry_type_filter and entry_type_filter != "All": | |
| df = df[df["entry_type"] == entry_type_filter] | |
| if df.empty: | |
| return self._empty_plot("No models match the filter criteria") | |
| if selected_plot_metric not in df.columns: | |
| return self._empty_plot(f"Metric '{metric_display_name}' not found") | |
| plot_df = df[["Model", selected_plot_metric]].dropna(subset=[selected_plot_metric]).copy() | |
| if plot_df.empty: | |
| return self._empty_plot(f"No data for metric '{metric_display_name}'") | |
| plot_df.columns = ["Model", "Score"] | |
| plot_df["Score"] = pd.to_numeric(plot_df["Score"], errors="coerce") | |
| plot_df = plot_df.dropna(subset=["Score"]) | |
| if plot_df.empty: | |
| return self._empty_plot(f"No valid data for metric '{metric_display_name}'") | |
| plot_df["Score"] = plot_df["Score"].round(2) | |
| ascending = plot_sort_mode == "Ascending (low → high)" | |
| plot_df = plot_df.sort_values("Score", ascending=ascending) | |
| plt.rcParams["font.sans-serif"] = ["Arial", "Microsoft YaHei", "SimHei", "DejaVu Sans"] | |
| plt.rcParams["axes.unicode_minus"] = False | |
| model_count = len(plot_df) | |
| figure_height = max(8, min(20, 0.48 * model_count + 4)) | |
| ytick_fontsize = max(11, min(22, 26 - model_count * 0.3)) | |
| fig, ax = plt.subplots(figsize=(16, figure_height), dpi=100) | |
| colors = plt.get_cmap("coolwarm_r")(np.linspace(0.12, 0.9, len(plot_df))) | |
| background_target = max(100, plot_df["Score"].max() * 1.12) | |
| ax.barh(plot_df["Model"], [background_target] * len(plot_df), color="#FAFAFA", edgecolor="none", height=0.7) | |
| bars = ax.barh(plot_df["Model"], plot_df["Score"], color=colors, edgecolor="none", height=0.7) | |
| for bar in bars: | |
| width = bar.get_width() | |
| ax.text(width + background_target * 0.015, bar.get_y() + bar.get_height() / 2, f"{width:.2f}", ha="left", va="center", fontsize=max(10, ytick_fontsize - 1), fontweight="bold", color="#444444") | |
| for spine in ax.spines.values(): | |
| spine.set_visible(False) | |
| ax.tick_params(axis="both", which="both", length=0) | |
| ax.set_xticks([]) | |
| ax.set_xlim(0, background_target * 1.05) | |
| ax.set_xlabel(metric_display_name, fontsize=24, fontweight="bold", labelpad=5, x=0.32, horizontalalignment="center") | |
| ax.set_ylabel("Model", fontsize=24, fontweight="bold") | |
| plt.yticks(fontsize=ytick_fontsize, fontweight="bold") | |
| plt.subplots_adjust(left=0.28 if model_count <= 16 else 0.42) | |
| ax.set_title(f"{metric_display_name} Leaderboard", fontsize=20, fontweight="bold", pad=30, x=0.32, horizontalalignment="center", y=1.05) | |
| plt.tight_layout() | |
| return fig | |
| def _empty_plot(self, message: str) -> plt.Figure: | |
| fig, ax = plt.subplots(figsize=(8, 6)) | |
| ax.text(0.5, 0.5, message, ha="center", va="center", fontsize=14) | |
| ax.axis("off") | |
| return fig | |