File size: 3,823 Bytes
c162d1f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
757e891
c162d1f
 
 
 
 
 
 
 
 
 
 
 
 
757e891
 
c162d1f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
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