serendipityAc2Win's picture
Replace tag filters with entry type filter
757e891 verified
Raw
History Blame Contribute Delete
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